159 lines
6.2 KiB
Plaintext
159 lines
6.2 KiB
Plaintext
/**
|
|
* Test 6-warp TMA FMHA multi-tile KV kernel (s_k > 128).
|
|
* Tests in-kernel online softmax rescale across KV tiles.
|
|
*/
|
|
|
|
#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;
|
|
|
|
#include "dsv4/kernels/attention/fmha_6warp_tma_multitile.cuh"
|
|
|
|
static size_t compute_smem() {
|
|
size_t off = 0;
|
|
off += 4; off = (off+127)&~(size_t)127;
|
|
off += 16; off = (off+127)&~(size_t)127;
|
|
off += TILE_SZ * 2; off = (off+127)&~(size_t)127; // sTmaBuf
|
|
off += TILE_SZ * 2; off = (off+127)&~(size_t)127; // sQ0
|
|
off += TILE_SZ * 2; off = (off+127)&~(size_t)127; // sK0
|
|
off += TILE_SZ * 2; off = (off+127)&~(size_t)127; // sPk
|
|
off += 16 * MY_MMA_K * 2; // sV
|
|
off += SK * 4; // s_p_vals
|
|
off += 4; // sTileMax
|
|
off += 4; // sTileSum
|
|
return off;
|
|
}
|
|
|
|
static void reference_attention(
|
|
const bf16_t* q, const bf16_t* k, const bf16_t* v,
|
|
float* o_ref, float* lse_ref,
|
|
int hd, int s_k, float scale
|
|
) {
|
|
float s[4096];
|
|
for (int j = 0; j < s_k; j++) {
|
|
float dot = 0.0f;
|
|
for (int d = 0; d < hd; d++) dot += bf16_to_f32_host(q[d]) * bf16_to_f32_host(k[j*hd+d]);
|
|
s[j] = dot * scale;
|
|
}
|
|
float mx = -INFINITY;
|
|
for (int j = 0; j < s_k; j++) mx = fmaxf(mx, s[j]);
|
|
float sm = 0.0f;
|
|
for (int j = 0; j < s_k; j++) { s[j] = expf(s[j] - mx); sm += s[j]; }
|
|
for (int j = 0; j < s_k; j++) s[j] /= sm;
|
|
for (int d = 0; d < hd; d++) {
|
|
float ov = 0.0f;
|
|
for (int j = 0; j < s_k; j++) ov += s[j] * bf16_to_f32_host(v[d*s_k+j]);
|
|
o_ref[d] = ov;
|
|
}
|
|
if (lse_ref) *lse_ref = logf(sm) + mx;
|
|
}
|
|
|
|
int main() {
|
|
printf("=== 6-warp TMA FMHA multi-tile HD=%d ===\n", HD);
|
|
const float SCALE = 1.0f / sqrtf((float)HD);
|
|
|
|
int total_fail = 0;
|
|
|
|
for (int s_k : {128, 256, 384, 512}) {
|
|
printf("\n--- s_k=%d (%d KV tiles) ---\n", s_k, (s_k + 127) / 128);
|
|
|
|
bf16_t* h_q = (bf16_t*)calloc(HD, sizeof(bf16_t));
|
|
bf16_t* h_k = (bf16_t*)calloc(s_k * HD, sizeof(bf16_t));
|
|
bf16_t* h_v = (bf16_t*)calloc(HD * s_k, 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 i=0;i<HD;i++) h_q[i] = f32_to_bf16_host((float)(rand()%100)/100.0f-0.5f);
|
|
for (int i=0;i<s_k*HD;i++) h_k[i] = f32_to_bf16_host((float)(rand()%100)/100.0f-0.5f);
|
|
for (int i=0;i<HD*s_k;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, s_k*HD*sizeof(bf16_t));
|
|
cudaMalloc(&d_v, HD*s_k*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, s_k*HD*sizeof(bf16_t), cudaMemcpyHostToDevice);
|
|
cudaMemcpy(d_v, h_v, HD*s_k*sizeof(bf16_t), cudaMemcpyHostToDevice);
|
|
|
|
CUtensorMap tma_k; CUtensorMap* d_tma_k;
|
|
create_tma_desc_2d_bf16(&tma_k, d_k, s_k, HD, 128, 16);
|
|
cudaMalloc(&d_tma_k, sizeof(CUtensorMap));
|
|
cudaMemcpy(d_tma_k, &tma_k, sizeof(CUtensorMap), cudaMemcpyHostToDevice);
|
|
|
|
CUtensorMap tma_v; CUtensorMap* d_tma_v;
|
|
create_tma_desc_2d_bf16(&tma_v, d_v, HD, s_k, 16, 16);
|
|
cudaMalloc(&d_tma_v, sizeof(CUtensorMap));
|
|
cudaMemcpy(d_tma_v, &tma_v, sizeof(CUtensorMap), cudaMemcpyHostToDevice);
|
|
|
|
FmhaTmaMultiTileParams params;
|
|
params.q = d_q; params.tma_k = d_tma_k; params.tma_v = d_tma_v; params.v = d_v;
|
|
params.o = d_o; params.lse = d_lse;
|
|
params.s_k = s_k; params.n_h = 1; params.scale = SCALE;
|
|
params.q_head_stride = HD; params.q_batch_stride = HD;
|
|
params.v_head_stride = HD*s_k; params.v_batch_stride = HD*s_k;
|
|
params.o_head_stride = HD; params.o_batch_stride = HD;
|
|
params.lse_head_stride = 1; params.lse_batch_stride = 1;
|
|
|
|
size_t smem = compute_smem();
|
|
if (smem > 48*1024)
|
|
cudaFuncSetAttribute(fmha_6warp_tma_multitile_kernel<HD>, cudaFuncAttributeMaxDynamicSharedMemorySize, (int)smem);
|
|
|
|
fmha_6warp_tma_multitile_kernel<HD><<<1, 192, smem>>>(params);
|
|
|
|
cudaError_t err = cudaDeviceSynchronize();
|
|
if (err != cudaSuccess) {
|
|
printf(" CUDA ERROR: %s\n", cudaGetErrorString(err));
|
|
total_fail++; continue;
|
|
}
|
|
|
|
cudaMemcpy(h_o, d_o, HD*sizeof(bf16_t), cudaMemcpyDeviceToHost);
|
|
|
|
float* o_ref = (float*)calloc(HD, sizeof(float));
|
|
reference_attention(h_q, h_k, h_v, o_ref, nullptr, HD, s_k, SCALE);
|
|
|
|
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(" cosine=%.8f %s\n", cs, cs>0.999f?"PASS":"FAIL");
|
|
if (cs < 0.999f) total_fail++;
|
|
|
|
if (HD <= 64 && total_fail == 0) {
|
|
printf(" O[0..3]: "); for(int d=0;d<4;d++) printf("%.6f ", bf16_to_f32_host(h_o[d])); printf("\n");
|
|
printf(" ref[0..3]: "); for(int d=0;d<4;d++) printf("%.6f ", o_ref[d]); printf("\n");
|
|
}
|
|
|
|
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); free(o_ref);
|
|
}
|
|
|
|
printf("\nOverall: %s\n", total_fail==0?"ALL PASSED":"SOME FAILED");
|
|
return total_fail == 0 ? 0 : 1;
|
|
}
|