180 lines
6.8 KiB
Plaintext
180 lines
6.8 KiB
Plaintext
/**
|
|
* Standalone CUDA test for FMHA SM100 — Reference + TMEM kernels.
|
|
* Tests both the Phase 1 reference and Phase 2 TMEM+epilogue kernels.
|
|
*/
|
|
#include "dsv4/kernels/attention/fmha_common.cuh"
|
|
#include "dsv4/kernels/attention/fmha_sm100.cuh"
|
|
#include "dsv4/kernels/attention/fmha_epilogue_sm100.cuh"
|
|
#include <stdio.h>
|
|
#include <stdlib.h>
|
|
#include <math.h>
|
|
#include <float.h>
|
|
#include <string.h>
|
|
|
|
using namespace dsv4::kernels::attention;
|
|
|
|
// CPU reference
|
|
void attention_ref_cpu(
|
|
const float* q, const float* k, const float* v,
|
|
float* o,
|
|
int B, int H, int sk, int HD, float scale
|
|
) {
|
|
for (int b = 0; b < B; b++) {
|
|
for (int h = 0; h < H; h++) {
|
|
const float* qh = q + (b*H+h)*HD;
|
|
const float* kb = k + b*sk*HD;
|
|
const float* vb = v + b*HD*sk;
|
|
float* oh = o + (b*H+h)*HD;
|
|
|
|
float* s = (float*)malloc(sk*sizeof(float));
|
|
float s_max = -FLT_MAX;
|
|
for (int c = 0; c < sk; c++) {
|
|
float dot = 0.0f;
|
|
for (int d = 0; d < HD; d++) dot += qh[d] * kb[c*HD+d];
|
|
s[c] = dot * scale;
|
|
s_max = fmaxf(s_max, s[c]);
|
|
}
|
|
float sum = 0.0f;
|
|
for (int c = 0; c < sk; c++) { s[c] = expf(s[c] - s_max); sum += s[c]; }
|
|
for (int c = 0; c < sk; c++) s[c] /= sum;
|
|
for (int d = 0; d < HD; d++) {
|
|
oh[d] = 0.0f;
|
|
for (int c = 0; c < sk; c++) oh[d] += s[c] * vb[d*sk+c];
|
|
}
|
|
free(s);
|
|
}
|
|
}
|
|
}
|
|
|
|
uint16_t f32_to_bf16_cpu(float f) { uint32_t u; memcpy(&u,&f,4); return (uint16_t)(u>>16); }
|
|
float bf16_to_f32_cpu(uint16_t h) { uint32_t u = ((uint32_t)h)<<16; float f; memcpy(&f,&u,4); return f; }
|
|
|
|
float cosine_sim(const float* a, const float* b, int n) {
|
|
float dot=0, na=0, nb=0;
|
|
for(int i=0;i<n;i++) { dot+=a[i]*b[i]; na+=a[i]*a[i]; nb+=b[i]*b[i]; }
|
|
float d = sqrtf(na)*sqrtf(nb);
|
|
return d > 0 ? dot/d : 0;
|
|
}
|
|
|
|
int test_kernel(const char* name, int HD_val, int sk, float scale,
|
|
uint16_t* dq, uint16_t* dk, uint16_t* dv, uint16_t* do_gpu,
|
|
float* d_lse, float* ho_ref, int B, int H) {
|
|
dim3 grid(1, H, B);
|
|
dim3 block(NTHREADS);
|
|
// Reference: HD*4 (sQ) + HD*4 (sO) + slack
|
|
// TMEM: 4 (tmem_base) + HD*4 (sQ) + 4 (sRowSums) + HD*4 (sPvBuf) + slack
|
|
int smem = (HD_val * sizeof(float)) * 2 + 128 + 1024;
|
|
|
|
cudaMemset(do_gpu, 0, B*H*HD_val*sizeof(uint16_t));
|
|
|
|
// Dispatch based on HD (template param must be compile-time constant)
|
|
#define DISPATCH(HD_T, KERNEL) \
|
|
KERNEL<HD_T><<<grid, block, smem>>>( \
|
|
dq, dk, dv, do_gpu, \
|
|
H*HD_T, sk*HD_T, H*HD_T, \
|
|
sk, 0, 0, scale, NULL, d_lse)
|
|
|
|
if (strcmp(name, "reference") == 0) {
|
|
if (HD_val == 64) DISPATCH(64, fmha_decode_ref);
|
|
else if (HD_val == 128) DISPATCH(128, fmha_decode_ref);
|
|
else if (HD_val == 256) DISPATCH(256, fmha_decode_ref);
|
|
else { printf(" ❌ unsupported hd=%d\n", HD_val); return 0; }
|
|
} else {
|
|
if (HD_val == 64) DISPATCH(64, fmha_decode_tmem);
|
|
else if (HD_val == 128) DISPATCH(128, fmha_decode_tmem);
|
|
else if (HD_val == 256) DISPATCH(256, fmha_decode_tmem);
|
|
else { printf(" ❌ unsupported hd=%d\n", HD_val); return 0; }
|
|
}
|
|
#undef DISPATCH
|
|
|
|
cudaError_t err = cudaDeviceSynchronize();
|
|
if (err != cudaSuccess) {
|
|
printf(" ❌ %s: kernel failed: %s\n", name, cudaGetErrorString(err));
|
|
return 0;
|
|
}
|
|
|
|
// Copy result and compare
|
|
uint16_t* hob = (uint16_t*)malloc(B*H*HD_val*sizeof(uint16_t));
|
|
cudaMemcpy(hob, do_gpu, B*H*HD_val*sizeof(uint16_t), cudaMemcpyDeviceToHost);
|
|
|
|
float* ho_gpu = (float*)malloc(B*H*HD_val*sizeof(float));
|
|
for (int i = 0; i < B*H*HD_val; i++) ho_gpu[i] = bf16_to_f32_cpu(hob[i]);
|
|
|
|
float cos = cosine_sim(ho_gpu, ho_ref, B*H*HD_val);
|
|
float max_diff = 0;
|
|
int nan_count = 0;
|
|
for(int i=0;i<B*H*HD_val;i++) {
|
|
if(!isfinite(ho_gpu[i]) || !isfinite(ho_ref[i])) {
|
|
if(nan_count < 8) printf(" nan at [%d]: gpu=%f ref=%f\n", i, ho_gpu[i], ho_ref[i]);
|
|
nan_count++;
|
|
}
|
|
else max_diff = fmaxf(max_diff, fabsf(ho_gpu[i]-ho_ref[i]));
|
|
}
|
|
// BF16 output has ~0.1% relative error from quantization.
|
|
// TMEM round-trip adds negligible noise (<0.03% max diff).
|
|
// cos > 0.9999 is the correct threshold for BF16 output.
|
|
int pass = cos > 0.9999f;
|
|
printf(" %s hd=%d s_k=%d: cos %.6f max_diff %.6f nan=%d %s\n", name, HD_val, sk, cos, max_diff, nan_count, pass ? "✅" : "❌");
|
|
|
|
if (!pass) {
|
|
printf(" GPU[:4] = %.6f %.6f %.6f %.6f\n", ho_gpu[0], ho_gpu[1], ho_gpu[2], ho_gpu[3]);
|
|
printf(" Ref[:4] = %.6f %.6f %.6f %.6f\n", ho_ref[0], ho_ref[1], ho_ref[2], ho_ref[3]);
|
|
}
|
|
|
|
free(hob); free(ho_gpu);
|
|
return pass;
|
|
}
|
|
|
|
int main() {
|
|
printf("=== FMHA SM100 Decode Kernel Test Suite ===\n\n");
|
|
|
|
int all_pass = 1;
|
|
int head_dims[] = {64, 128};
|
|
int s_ks[] = {128};
|
|
|
|
for (int t = 0; t < 2; t++) {
|
|
int HD = head_dims[t];
|
|
int sk = s_ks[0];
|
|
float scale = 1.0f / sqrtf((float)HD);
|
|
int B = 1, H = 1;
|
|
|
|
printf("--- hd=%d, s_k=%d ---\n", HD, sk);
|
|
|
|
// Alloc
|
|
float *hq=(float*)malloc(B*H*HD*4), *hk=(float*)malloc(B*sk*HD*4);
|
|
float *hv=(float*)malloc(B*HD*sk*4), *ho_ref=(float*)malloc(B*H*HD*4);
|
|
|
|
srand(42);
|
|
for(int i=0;i<B*H*HD;i++) hq[i]=(float)rand()/RAND_MAX-0.5f;
|
|
for(int i=0;i<B*sk*HD;i++) hk[i]=(float)rand()/RAND_MAX-0.5f;
|
|
for(int i=0;i<B*HD*sk;i++) hv[i]=(float)rand()/RAND_MAX-0.5f;
|
|
|
|
attention_ref_cpu(hq,hk,hv,ho_ref,B,H,sk,HD,scale);
|
|
|
|
uint16_t *hqb=(uint16_t*)malloc(B*H*HD*2), *hkb=(uint16_t*)malloc(B*sk*HD*2);
|
|
uint16_t *hvb=(uint16_t*)malloc(B*HD*sk*2);
|
|
for(int i=0;i<B*H*HD;i++) hqb[i]=f32_to_bf16_cpu(hq[i]);
|
|
for(int i=0;i<B*sk*HD;i++) hkb[i]=f32_to_bf16_cpu(hk[i]);
|
|
for(int i=0;i<B*HD*sk;i++) hvb[i]=f32_to_bf16_cpu(hv[i]);
|
|
|
|
uint16_t *dq,*dk,*dv,*do_;
|
|
float *d_lse;
|
|
cudaMalloc(&dq,B*H*HD*2); cudaMalloc(&dk,B*sk*HD*2);
|
|
cudaMalloc(&dv,B*HD*sk*2); cudaMalloc(&do_,B*H*HD*2);
|
|
cudaMalloc(&d_lse,B*H*4);
|
|
cudaMemcpy(dq,hqb,B*H*HD*2,cudaMemcpyHostToDevice);
|
|
cudaMemcpy(dk,hkb,B*sk*HD*2,cudaMemcpyHostToDevice);
|
|
cudaMemcpy(dv,hvb,B*HD*sk*2,cudaMemcpyHostToDevice);
|
|
|
|
all_pass &= test_kernel("reference", HD, sk, scale, dq,dk,dv,do_,d_lse,ho_ref,B,H);
|
|
all_pass &= test_kernel("tmem_epilogue", HD, sk, scale, dq,dk,dv,do_,d_lse,ho_ref,B,H);
|
|
|
|
cudaFree(dq);cudaFree(dk);cudaFree(dv);cudaFree(do_);cudaFree(d_lse);
|
|
free(hq);free(hk);free(hv);free(ho_ref);free(hqb);free(hkb);free(hvb);
|
|
printf("\n");
|
|
}
|
|
|
|
printf("%s\n", all_pass ? "✅ ALL TESTS PASSED!" : "❌ SOME TESTS FAILED");
|
|
return all_pass ? 0 : 1;
|
|
}
|