FMHA SM100: Add TMEM+correction epilogue kernel (Priority 2)

New file: fmha_epilogue_sm100.cuh
- TMEM alloc/dealloc/load/store via tcgen05 PTX
- One-way correction epilogue: TMEM→regs→normalize→BF16→GMEM
- D1.5 fix: O rescale in REGISTERS (TMEM→regs→multiply→TMEM)
- Same pattern as MoE epilogue but with normalize instead of SwiGLU
- Unblocks D2 multi-CTA and NVFP4-1.2 (register slot for FP4 pack)

Test: hd=64 + hd=128, reference vs TMEM kernels
This commit is contained in:
2026-05-28 06:27:56 +00:00
parent 8eb735618f
commit bcc5d0b6cb
3 changed files with 511 additions and 127 deletions

View File

@@ -1,172 +1,156 @@
/**
* Standalone CUDA test for FMHA SM100 decode kernel.
* Launches the kernel directly via CUDA runtime, compares against CPU reference.
* No PyTorch or pybind11 needed — just nvcc + CUDA runtime.
* 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_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: simple attention
// CPU reference
void attention_ref_cpu(
const float* q, const float* k, const float* v,
float* o, float* lse,
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;
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;
// S = Q @ K^T * scale
float* s = (float*)malloc(sk * sizeof(float));
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];
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]);
}
// Softmax
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] = expf(s[c] - s_max); sum += s[c]; }
for (int c = 0; c < sk; c++) s[c] /= sum;
// O = S @ V
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];
}
for (int c = 0; c < sk; c++) oh[d] += s[c] * vb[d*sk+c];
}
if (lse) lse[b * H + h] = logf(sum) + s_max;
free(s);
}
}
}
// BF16 conversion helpers for CPU
uint16_t f32_to_bf16_cpu(float f) {
uint32_t u;
memcpy(&u, &f, 4);
uint16_t h = (uint16_t)(u >> 16);
return h;
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;
}
float bf16_to_f32_cpu(uint16_t h) {
uint32_t u = ((uint32_t)h) << 16;
float f;
memcpy(&f, &u, 4);
return f;
}
int main() {
printf("=== FMHA SM100 Decode Kernel Test ===\n");
const int B = 1, H = 1, HD = 64, sk = 128;
const float scale = 1.0f / sqrtf((float)HD);
const int smem = 128 * HD * 2 * sizeof(uint16_t) + 1024; // K + V + slack
// Allocate host memory
float *hq = (float*)malloc(B * H * HD * sizeof(float));
float *hk = (float*)malloc(B * sk * HD * sizeof(float));
float *hv = (float*)malloc(B * HD * sk * sizeof(float));
float *ho_ref = (float*)malloc(B * H * HD * sizeof(float));
// Init with random data
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;
// CPU reference
attention_ref_cpu(hq, hk, hv, ho_ref, NULL, B, H, sk, HD, scale);
// Convert to BF16
uint16_t *hqb = (uint16_t*)malloc(B * H * HD * sizeof(uint16_t));
uint16_t *hkb = (uint16_t*)malloc(B * sk * HD * sizeof(uint16_t));
uint16_t *hvb = (uint16_t*)malloc(B * HD * sk * sizeof(uint16_t));
uint16_t *hob = (uint16_t*)malloc(B * H * HD * sizeof(uint16_t));
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]);
// Allocate GPU memory
uint16_t *dq, *dk, *dv, *do_;
float *d_lse;
cudaMalloc(&dq, B * H * HD * sizeof(uint16_t));
cudaMalloc(&dk, B * sk * HD * sizeof(uint16_t));
cudaMalloc(&dv, B * HD * sk * sizeof(uint16_t));
cudaMalloc(&do_, B * H * HD * sizeof(uint16_t));
cudaMalloc(&d_lse, B * H * sizeof(float));
// Copy to GPU
cudaMemcpy(dq, hqb, B * H * HD * sizeof(uint16_t), cudaMemcpyHostToDevice);
cudaMemcpy(dk, hkb, B * sk * HD * sizeof(uint16_t), cudaMemcpyHostToDevice);
cudaMemcpy(dv, hvb, B * HD * sk * sizeof(uint16_t), cudaMemcpyHostToDevice);
cudaMemset(do_, 0, B * H * HD * sizeof(uint16_t));
// Launch kernel
int test_kernel(const char* name, int HD, 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);
int smem = (HD * sizeof(float)) + 128 + 1024; // Q + row_sums + slack
printf("Launching fmha_decode_ref<%d> <<<(%d,%d,%d), %d>>>...\n", HD, grid.x, grid.y, grid.z, block.x);
cudaMemset(do_gpu, 0, B*H*HD*sizeof(uint16_t));
fmha_decode_ref<HD><<<grid, block, smem>>>(
dq, dk, dv, do_,
H * HD, sk * HD, H * HD,
sk, 0, 0, scale, NULL, d_lse
);
if (strcmp(name, "reference") == 0) {
fmha_decode_ref<HD><<<grid, block, smem>>>(
dq, dk, dv, do_gpu,
H*HD, sk*HD, H*HD,
sk, 0, 0, scale, NULL, d_lse);
} else {
fmha_decode_tmem<HD><<<grid, block, smem>>>(
dq, dk, dv, do_gpu,
H*HD, sk*HD, H*HD,
sk, 0, 0, scale, NULL, d_lse);
}
cudaError_t err = cudaDeviceSynchronize();
if (err != cudaSuccess) {
printf("❌ Kernel launch failed: %s\n", cudaGetErrorString(err));
return 1;
}
printf("✅ Kernel launched successfully!\n");
// Copy result back
cudaMemcpy(hob, do_, B * H * HD * sizeof(uint16_t), cudaMemcpyDeviceToHost);
// Compare with reference
float cos_sim = 0.0f, norm_a = 0.0f, norm_b = 0.0f;
for (int i = 0; i < B * H * HD; i++) {
float gpu_val = bf16_to_f32_cpu(hob[i]);
float ref_val = ho_ref[i];
cos_sim += gpu_val * ref_val;
norm_a += gpu_val * gpu_val;
norm_b += ref_val * ref_val;
}
float denom = sqrtf(norm_a) * sqrtf(norm_b);
if (denom > 0) cos_sim /= denom;
printf("\nhd=%d, s_k=%d: cos %.6f %s\n", HD, sk, cos_sim, cos_sim > 0.999f ? "✅ PASS" : "❌ FAIL");
if (cos_sim < 0.999f) {
printf("First 8 values (GPU vs Ref):\n");
for (int i = 0; i < 8; i++) {
printf(" [%d] GPU=%f Ref=%f\n", i, bf16_to_f32_cpu(hob[i]), ho_ref[i]);
}
printf(" ❌ %s: kernel failed: %s\n", name, cudaGetErrorString(err));
return 0;
}
// Cleanup
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); free(hob);
// Copy result and compare
uint16_t* hob = (uint16_t*)malloc(B*H*HD*sizeof(uint16_t));
cudaMemcpy(hob, do_gpu, B*H*HD*sizeof(uint16_t), cudaMemcpyDeviceToHost);
return cos_sim > 0.999f ? 0 : 1;
float* ho_gpu = (float*)malloc(B*H*HD*sizeof(float));
for (int i = 0; i < B*H*HD; i++) ho_gpu[i] = bf16_to_f32_cpu(hob[i]);
float cos = cosine_sim(ho_gpu, ho_ref, B*H*HD);
int pass = cos > 0.999f;
printf(" %s hd=%d s_k=%d: cos %.6f %s\n", name, HD, sk, cos, 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;
}