From 88c72a887e36ae7498c9439f04a3c03f2af6acf2 Mon Sep 17 00:00:00 2001 From: biondizzle Date: Fri, 29 May 2026 22:51:24 +0000 Subject: [PATCH] feat: V TMA loads in multi-row kernel --- .../attention/fmha_6warp_tma_multirow.cuh | 31 +++++++++++-------- tests/unit/test_fmha_6warp_tma_multirow.cu | 15 +++++++-- 2 files changed, 30 insertions(+), 16 deletions(-) diff --git a/dsv4/kernels/attention/fmha_6warp_tma_multirow.cuh b/dsv4/kernels/attention/fmha_6warp_tma_multirow.cuh index 617152d5..7d66627b 100644 --- a/dsv4/kernels/attention/fmha_6warp_tma_multirow.cuh +++ b/dsv4/kernels/attention/fmha_6warp_tma_multirow.cuh @@ -35,7 +35,8 @@ namespace dsv4::kernels::attention { struct FmhaTmaMultiRowParams { const bf16_t* __restrict__ q; CUtensorMap* __restrict__ tma_k; // Array of [n_h] TMA descriptors for K - const bf16_t* __restrict__ v; // V: direct GMEM (HD, s_k) + CUtensorMap* __restrict__ tma_v; // Array of [n_h] TMA descriptors for V + const bf16_t* __restrict__ v; // V: direct fallback (HD, s_k) bf16_t* __restrict__ o; float* __restrict__ lse; int s_k, T, n_h; @@ -113,6 +114,7 @@ fmha_6warp_tma_multirow_kernel(FmhaTmaMultiRowParams params) { int phase = 0; CUtensorMap* __restrict__ my_tma_k = params.tma_k + batch_idx * params.n_h + head_idx; + CUtensorMap* __restrict__ my_tma_v = params.tma_v + batch_idx * params.n_h + head_idx; const bool my_warp_active = (T <= 32) ? (wid == 0) : is_softmax_warp; const int my_row = my_warp_active ? (wid * 32 + lane) : 0; const bool my_row_active = my_warp_active && (my_row < T); @@ -243,18 +245,21 @@ fmha_6warp_tma_multirow_kernel(FmhaTmaMultiRowParams params) { } __syncthreads(); - // Load V sub-tile: direct from GMEM - if (is_load_warp) { - for (int i = lane; i < V_SUB_SZ; i += 32) sV[i] = 0; - for (int dd = lane; dd < 16; dd += 32) { - for (int lr = 0; lr < MMA_K_BF16; lr++) { - int r = col_start + lr; - if (r < s_k && (d_base + dd) < HD) { - int g_mn = dd/8, g_k = lr/8, llr = dd%8, lc = lr%8; - sV[g_k*2*64 + g_mn*64 + llr*8 + lc] = v_head[(d_base+dd)*s_k + r]; - } - } - } + // Load V sub-tile via TMA + if (is_load_warp && lane == 0) { + tma_load_2d((uint32_t)__cvta_generic_to_shared(sTmaBuf), (uint64_t)my_tma_v, + mbar_addr, col_start, d_base); + tma_mbarrier_arrive_expect_tx(mbar_addr, V_SUB_SZ * sizeof(bf16_t)); + } + tma_mbarrier_wait(mbar_addr, phase); phase ^= 1; + __syncthreads(); + + // Convert sTmaBuf → canonical sV + for (int i = tid; i < V_SUB_SZ; i += 192) sV[i] = 0; + for (int i = tid; i < 16 * MMA_K_BF16; i += 192) { + int dd = i / MMA_K_BF16, lr = i % MMA_K_BF16; + int g_mn = dd/8, g_k = lr/8, llr = dd%8, lc = lr%8; + sV[g_k*2*64 + g_mn*64 + llr*8 + lc] = sTmaBuf[i]; } __syncthreads(); diff --git a/tests/unit/test_fmha_6warp_tma_multirow.cu b/tests/unit/test_fmha_6warp_tma_multirow.cu index bc589ab0..8189bdf7 100644 --- a/tests/unit/test_fmha_6warp_tma_multirow.cu +++ b/tests/unit/test_fmha_6warp_tma_multirow.cu @@ -109,8 +109,17 @@ static int test_single(int T, int n_h = 1, int batch = 1) { } cudaMemcpy(d_tma_k, tma_k_arr, total_heads * sizeof(CUtensorMap), cudaMemcpyHostToDevice); + // TMA descriptors for V: one per head + CUtensorMap* tma_v_arr = (CUtensorMap*)malloc(total_heads * sizeof(CUtensorMap)); + CUtensorMap* d_tma_v; + cudaMalloc(&d_tma_v, total_heads * sizeof(CUtensorMap)); + for (int h = 0; h < total_heads; h++) { + create_tma_desc_2d_bf16(&tma_v_arr[h], d_v + h*HD*SK, HD, SK, 16, 16); + } + cudaMemcpy(d_tma_v, tma_v_arr, total_heads * sizeof(CUtensorMap), cudaMemcpyHostToDevice); + FmhaTmaMultiRowParams params; - params.q = d_q; params.tma_k = d_tma_k; params.v = d_v; + 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 = SK; params.T = T; params.scale = SCALE; params.n_h = n_h; params.head_dim = HD; @@ -159,8 +168,8 @@ static int test_single(int T, int n_h = 1, int batch = 1) { } printf(" min_cos=%.8f %s\n", min_cos, min_cos>0.999f?"PASS":"FAIL"); - cudaFree(d_q); cudaFree(d_k); cudaFree(d_v); cudaFree(d_o); cudaFree(d_lse); cudaFree(d_tma_k); - free(h_q); free(h_k); free(h_v); free(h_o); free(h_lse); free(tma_k_arr); + 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(tma_k_arr); free(tma_v_arr); return min_cos > 0.999f ? 0 : 1; }