test: minimal QK — 128 threads, tid==0 MMA, match working gen kernel pattern
This commit is contained in:
@@ -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; }
|
||||
|
||||
|
||||
Reference in New Issue
Block a user