Files
nvfp4-megamoe-kernel/tests/unit/test_fmha_6warp_tma.cu
2026-05-30 03:46:38 +00:00

154 lines
5.8 KiB
Plaintext

/**
* Test 6-warp TMA FMHA kernel for HD=64/128/256.
* TMA loads K, direct loads Q and V.
*/
#include <cuda_runtime.h>
#include <cuda.h>
#include <cstdio>
#include <cmath>
#include <cstdlib>
#include <cstring>
#ifndef HD_VAL
#define HD_VAL 64
#endif
#include "dsv4/kernels/attention/fmha_common.cuh"
#include "dsv4/kernels/attention/fmha_umma_desc.cuh"
#include "dsv4/kernels/attention/fmha_tma.cuh"
using namespace dsv4::kernels::attention;
static bf16_t f32_to_bf16_host(float f) { uint32_t u; memcpy(&u,&f,4); return (uint16_t)(u>>16); }
static float bf16_to_f32_host(bf16_t h) { uint32_t u=(uint32_t)h<<16; float f; memcpy(&f,&u,4); return f; }
constexpr int HD = HD_VAL;
constexpr int SK = 128;
constexpr int MY_MMA_K = 16;
constexpr int TILE_SZ = 128 * MY_MMA_K;
constexpr int V_SUB_SZ = 16 * MY_MMA_K;
constexpr int TMEM_N = (HD <= 128) ? 128 : (HD <= 256) ? 256 : 512;
#include "dsv4/kernels/attention/fmha_6warp_tma.cuh"
static size_t compute_smem() {
size_t off = 0;
off += 4; // sTmemBase
off = (off + 127) & ~(size_t)127;
off += 16; // sMbar
off = (off + 127) & ~(size_t)127;
off += TILE_SZ * sizeof(bf16_t); // sTmaBuf
off = (off + 127) & ~(size_t)127;
off += TILE_SZ * sizeof(bf16_t); // sQ0
off = (off + 127) & ~(size_t)127;
off += TILE_SZ * sizeof(bf16_t); // sK0
off = (off + 127) & ~(size_t)127;
off += TILE_SZ * sizeof(bf16_t); // sK1 (double buffer)
off = (off + 127) & ~(size_t)127;
off += TILE_SZ * sizeof(bf16_t); // sPk
off = (off + 127) & ~(size_t)127;
off += V_SUB_SZ * sizeof(bf16_t); // sV
off += 128 * sizeof(float); // sRowMax
off += 128 * sizeof(float); // sRowSum
off += SK * sizeof(float); // s_p_vals
return off;
}
int main() {
printf("=== 6-warp TMA FMHA HD=%d ===\n", HD);
const float SCALE = 1.0f / sqrtf((float)HD);
bf16_t* h_q = (bf16_t*)malloc(HD*sizeof(bf16_t));
bf16_t* h_k = (bf16_t*)malloc(SK*HD*sizeof(bf16_t));
bf16_t* h_v = (bf16_t*)malloc(HD*SK*sizeof(bf16_t));
bf16_t* h_o = (bf16_t*)calloc(HD, sizeof(bf16_t));
float* h_lse = (float*)calloc(1, sizeof(float));
srand(42);
for (int d=0;d<HD;d++) h_q[d] = f32_to_bf16_host((float)(rand()%100)/100.0f-0.5f);
for (int i=0;i<SK*HD;i++) h_k[i] = f32_to_bf16_host((float)(rand()%100)/100.0f-0.5f);
for (int i=0;i<HD*SK;i++) h_v[i] = f32_to_bf16_host((float)(rand()%100)/100.0f-0.5f);
bf16_t *d_q,*d_k,*d_v,*d_o;
float *d_lse;
cudaMalloc(&d_q, HD*sizeof(bf16_t));
cudaMalloc(&d_k, SK*HD*sizeof(bf16_t));
cudaMalloc(&d_v, HD*SK*sizeof(bf16_t));
cudaMalloc(&d_o, HD*sizeof(bf16_t));
cudaMalloc(&d_lse, sizeof(float));
cudaMemcpy(d_q, h_q, HD*sizeof(bf16_t), cudaMemcpyHostToDevice);
cudaMemcpy(d_k, h_k, SK*HD*sizeof(bf16_t), cudaMemcpyHostToDevice);
cudaMemcpy(d_v, h_v, HD*SK*sizeof(bf16_t), cudaMemcpyHostToDevice);
// TMA descriptor for K: (SK, HD) with tile (128, 16)
CUtensorMap tma_k; CUtensorMap* d_tma_k;
if (!create_tma_desc_2d_bf16(&tma_k, d_k, SK, HD, 128, 16)) {
printf("TMA K desc FAILED\n"); return 1;
}
cudaMalloc(&d_tma_k, sizeof(CUtensorMap));
cudaMemcpy(d_tma_k, &tma_k, sizeof(CUtensorMap), cudaMemcpyHostToDevice);
// TMA descriptor for V: (HD, SK) with tile (16, 16)
CUtensorMap tma_v; CUtensorMap* d_tma_v;
if (!create_tma_desc_2d_bf16(&tma_v, d_v, HD, SK, 16, 16)) {
printf("TMA V desc FAILED\n"); return 1;
}
cudaMalloc(&d_tma_v, sizeof(CUtensorMap));
cudaMemcpy(d_tma_v, &tma_v, sizeof(CUtensorMap), cudaMemcpyHostToDevice);
// Compute reference
float o_ref[HD];
{
float s[SK];
for (int j=0;j<SK;j++) {
float dot = 0.0f;
for (int d=0;d<HD;d++) dot += bf16_to_f32_host(h_q[d]) * bf16_to_f32_host(h_k[j*HD+d]);
s[j] = dot * SCALE;
}
float mx = -INFINITY;
for (int j=0;j<SK;j++) mx = fmaxf(mx, s[j]);
float sm = 0.0f;
for (int j=0;j<SK;j++) { s[j] = expf(s[j]-mx); sm += s[j]; }
for (int j=0;j<SK;j++) s[j] /= sm;
for (int d=0;d<HD;d++) {
float ov = 0.0f;
for (int j=0;j<SK;j++) ov += s[j] * bf16_to_f32_host(h_v[d*SK+j]);
o_ref[d] = ov;
}
}
size_t smem = compute_smem();
printf("SMEM: %zu bytes (%.1f KB), 6 warps (192 threads), TMEM: %d cols\n", smem, smem/1024.0, TMEM_N);
if (smem > 48 * 1024) {
cudaFuncSetAttribute(fmha_6warp_tma_kernel<HD>, cudaFuncAttributeMaxDynamicSharedMemorySize, (int)smem);
}
fmha_6warp_tma_kernel<HD><<<1, 192, smem>>>(d_q, d_tma_k, d_tma_v, d_o, d_lse, SK, SCALE);
cudaError_t launch_err = cudaGetLastError();
if (launch_err != cudaSuccess) { printf("LAUNCH ERROR: %s\n", cudaGetErrorString(launch_err)); return 1; }
cudaError_t err = cudaDeviceSynchronize();
if (err != cudaSuccess) { printf("CUDA ERROR: %s\n", cudaGetErrorString(err)); return 1; }
cudaMemcpy(h_o, d_o, HD*sizeof(bf16_t), cudaMemcpyDeviceToHost);
cudaMemcpy(h_lse, d_lse, sizeof(float), cudaMemcpyDeviceToHost);
printf("O[0..7] MMA: "); for(int d=0;d<min(8,HD);d++) printf("%.6f ",bf16_to_f32_host(h_o[d])); printf("\n");
printf("O[0..7] ref: "); for(int d=0;d<min(8,HD);d++) printf("%.6f ",o_ref[d]); printf("\n");
float cs=0,na=0,nb=0;
for (int d=0;d<HD;d++) {
float a=bf16_to_f32_host(h_o[d]),b=o_ref[d];
if(fabsf(b)>1e-4f) { cs+=a*b; na+=a*a; nb+=b*b; }
}
cs /= (sqrtf(na)*sqrtf(nb)+1e-10f);
printf("Filtered cosine: %.8f\n", cs);
printf("Test %s\n", cs > 0.999f ? "PASSED" : "FAILED");
cudaFree(d_q); cudaFree(d_k); cudaFree(d_v); cudaFree(d_o); cudaFree(d_lse); cudaFree(d_tma_k); cudaFree(d_tma_v);
free(h_q); free(h_k); free(h_v); free(h_o); free(h_lse);
return cs > 0.999f ? 0 : 1;
}