test: minimal QK — 128 threads, tid==0 MMA, match working gen kernel pattern

This commit is contained in:
2026-05-29 18:40:11 +00:00
parent 39aef1284f
commit bd1309ba88

View File

@@ -27,7 +27,7 @@ constexpr int SK = 128;
constexpr int NKT = HD / MMA_K_BF16;
constexpr int CORES_MN = 16; // 128/8
__global__ void __launch_bounds__(192)
__global__ void __launch_bounds__(128)
test_qk_minimal_kernel(float* __restrict__ out_s,
const bf16_t* __restrict__ q, const bf16_t* __restrict__ k, int T, int s_k)
{
@@ -36,9 +36,9 @@ test_qk_minimal_kernel(float* __restrict__ out_s,
static constexpr int NUM_READS = SK / 8;
const int tid = threadIdx.x;
const int wid = tid / 32;
const int lane = tid % 32;
const bool is_mma_warp = (wid == 4);
// Simple 4-warp: warp 0 = load+softmax+MMA, warps 1-3 = softmax
// ALL 128 threads participate in loading and MMA is called by tid==0
// SMEM: sQ0 and sK0 are (128, 16) each
extern __shared__ __align__(128) char sbuf[];
@@ -49,44 +49,39 @@ test_qk_minimal_kernel(float* __restrict__ out_s,
off = (off + 127) & ~(size_t)127;
bf16_t* sK0 = (bf16_t*)(sbuf + off); off += TILE_SZ * sizeof(bf16_t);
if (is_mma_warp) tmem_alloc(__cvta_generic_to_shared(sTmemBase), TMEM_N);
if (tid < 32) tmem_alloc(__cvta_generic_to_shared(sTmemBase), TMEM_N);
__syncthreads();
uint32_t tb = *sTmemBase;
for (int kt = 0; kt < NKT; kt++) {
// Load Q sub-tile
for (int i = tid; i < TILE_SZ; i += NTHREADS) sQ0[i] = 0;
for (int r = tid / 32; r < T; r += 6) { // one row per warp
if (r < T) {
for (int d = lane; d < MMA_K_BF16; d += 32) {
int full_d = kt * MMA_K_BF16 + d;
if (full_d < HD) {
int ck = d/8, lc = d%8, cm = r/8, lr = r%8;
sQ0[ck*CORES_MN*64 + cm*64 + lr*8 + lc] = q[r * HD + full_d];
}
}
// Load Q sub-tile — all 128 threads participate
for (int i = tid; i < TILE_SZ; i += 128) sQ0[i] = 0;
for (int d = tid; d < MMA_K_BF16; d += 128) {
int full_d = kt * MMA_K_BF16 + d;
if (full_d < HD) {
// Q has 1 row (T=1), row 0: cm=0, lr=0
int ck = d/8, lc = d%8;
sQ0[ck*CORES_MN*64 + 0*64 + 0*8 + lc] = q[full_d];
}
}
__syncthreads();
// Load K sub-tile
for (int i = tid; i < TILE_SZ; i += NTHREADS) sK0[i] = 0;
for (int r = lane; r < s_k; r += 32) {
// Load K sub-tile — all 128 threads
for (int i = tid; i < TILE_SZ; i += 128) sK0[i] = 0;
for (int r = tid; r < s_k; r += 128) {
for (int d = 0; d < MMA_K_BF16; d++) {
int full_d = kt * MMA_K_BF16 + d;
if (full_d < HD) {
int ck = d/8, lc = d%8, cm = r/8, lr = r%8;
sK0[ck*CORES_MN*64 + cm*64 + lr*8 + lc] = k[r * HD + full_d];
}
int ck = d/8, lc = d%8, cm = r/8, lr = r%8;
sK0[ck*CORES_MN*64 + cm*64 + lr*8 + lc] = k[r * HD + full_d];
}
}
__syncthreads();
if (is_mma_warp) {
if (tid == 0) {
uint32_t idesc = make_idesc(128, 128);
uint64_t dq = make_umma_desc_kmajor_none(__cvta_generic_to_shared(sQ0), 128);
uint64_t dk = make_umma_desc_kmajor_none(__cvta_generic_to_shared(sK0), 128);
if (tid == 128) umma_ss_f16(tb, dq, dk, idesc, kt > 0);
umma_ss_f16(tb, dq, dk, idesc, kt > 0);
asm volatile("tcgen05.fence::after_thread_sync;" ::: "memory");
}
__syncthreads();
@@ -134,7 +129,7 @@ int main() {
cudaMemcpy(d_k, h_k, SK * HD * sizeof(bf16_t), cudaMemcpyHostToDevice);
int smem = 4 + 128 + 128*16*2 + 128*16*2 + 4096;
test_qk_minimal_kernel<<<1, 192, smem>>>(d_out, d_q, d_k, T, SK);
test_qk_minimal_kernel<<<1, 128, smem>>>(d_out, d_q, d_k, T, SK);
cudaError_t err = cudaDeviceSynchronize();
if (err != cudaSuccess) { printf("CUDA ERROR: %s\n", cudaGetErrorString(err)); return 1; }