diff --git a/dsv4/kernels/attention/fmha_6warp_multihead.cuh b/dsv4/kernels/attention/fmha_6warp_multihead.cuh index 107fa999..a8a8c87a 100644 --- a/dsv4/kernels/attention/fmha_6warp_multihead.cuh +++ b/dsv4/kernels/attention/fmha_6warp_multihead.cuh @@ -1,30 +1,54 @@ /** - * DSV4 FMHA — 6-warp specialized kernel, multi-head, multi-KV-tile. + * DSV4 FMHA — 6-warp specialized kernel, multi-head launch. * * ================================================================== - * MULTI-HEAD + MULTI-KV-TILE (P3 + P5) + * MULTI-HEAD LAUNCH (Milestone 5) * ================================================================== * Grid: dim3(1, n_h, batch_size) + * blockIdx.y = head index (0..n_h-1) + * blockIdx.z = batch index (0..batch_size-1) + * * Each CTA processes one head of one batch item independently. - * - * n_kv_tiles <= 1 (single KV segment, s_k <= 128): - * Direct QK → softmax → PV → epilogue. No rescale. - * Identical to the pre-P5 kernel. - * - * n_kv_tiles > 1 (multiple KV segments, s_k > 128): - * FlashAttention-2 online softmax across KV tiles. - * - P is UN-NORMALIZED: exp(s - max), NOT divided by sum - * - SMEM accumulator sOacc[HD]: O += PV after each tile - * - Rescale: sOacc *= exp(old_max - new_max) when max changes - * - Final: O = sOacc / running_sum + * No cross-CTA synchronization required. * * ================================================================== - * SMEM BUDGET (additional for multi-tile) + * MQA / GQA SUPPORT * ================================================================== - * sOacc: HD * 4 bytes (float accumulator, 1 row for T=1) - * hd=64: +256B. hd=128: +512B. hd=256: +1024B. - * Total SMEM at hd=64: ~14 KB. hd=128: ~15 KB. hd=256: ~17 KB. - * Well within 232 KB SMEM budget. + * - MQA: all Q heads share one KV head. Pass k_head_stride=0, v_head_stride=0 + * so all CTAs read the same K/V. + * - GQA: groups of Q heads share a KV head. The caller must arrange + * K/V tensors so that k_head_stride/v_head_stride map correctly. + * - MHA: k_head_stride = k_row_stride * N, same for V. + * + * ================================================================== + * TENSOR LAYOUTS (GMEM) + * ================================================================== + * Q: [batch, n_h, T, hd] — head stride = T * hd, batch stride = n_h * T * hd + * K: [batch, n_kv, N, hd] — head stride = N * hd (or 0 for MQA) + * V: [batch, n_kv, hd, N] — head stride = hd * N (or 0 for MQA) + * O: [batch, n_h, T, hd] — same strides as Q + * + * For decode (T=1): q_head_offset = blockIdx.y * hd, q_batch_offset = blockIdx.z * n_h * hd + * For prefill (T>1): head-packed M = T rows per head (must fit in 128-row MMA tile) + * + * ================================================================== + * SOFTMAX ROWS + * ================================================================== + * T=1 decode: only row 0 of the 128-row MMA tile has data. Only warp 0 + * computes softmax for row 0. + * T>1 prefill: rows 0..T-1 have data. All 4 softmax warps process + * rows in parallel (warp w handles rows [w*32, (w+1)*32) ∩ [0, T)). + * This is Milestone 4 territory — current implementation handles T=1 only. + * The multi-head grid layout is independent of multi-row softmax and + * can land first. + * + * ================================================================== + * OUTPUT: UN-NORMALIZED O + LSE + * ================================================================== + * The kernel emits un-normalized O and per-row LSE for composition with + * D5 multi-tile KV merge. External code normalizes: O_norm = O / row_sum. + * For single-segment decode, normalization is done in the epilogue. + * LSE layout: [batch, n_h, T] — one float per head per row. */ #pragma once @@ -34,33 +58,43 @@ namespace dsv4::kernels::attention { +/** + * Multi-head FMHA kernel parameters. + * + * All strides are in units of BF16 elements (not bytes). + */ struct FmhaParams { - const bf16_t* __restrict__ q; - const bf16_t* __restrict__ k; - const bf16_t* __restrict__ v; - bf16_t* __restrict__ o; - float* __restrict__ lse; + const bf16_t* __restrict__ q; // Q base pointer + const bf16_t* __restrict__ k; // K base pointer + const bf16_t* __restrict__ v; // V base pointer + bf16_t* __restrict__ o; // O base pointer + float* __restrict__ lse; // LSE base pointer [batch, n_h, T] (optional, can be nullptr) - int s_k; // Total KV sequence length + int s_k; // KV sequence length float scale; // 1/sqrt(hd) int head_dim; // hd - int n_kv_tiles; // Number of KV tiles (0 or 1 = single-tile, >1 = multi-tile) - int q_head_stride, q_batch_stride; - int k_head_stride, k_batch_stride; - int v_head_stride, v_batch_stride; - int o_head_stride, o_batch_stride; - int lse_head_stride, lse_batch_stride; + // Strides (in BF16 elements) + int q_head_stride; // stride between Q heads = T * hd + int q_batch_stride; // stride between Q batch items = n_h * T * hd + int k_head_stride; // stride between K heads = N * hd (0 for MQA) + int k_batch_stride; // stride between K batch items = n_kv * N * hd + int v_head_stride; // stride between V heads = hd * N (0 for MQA) + int v_batch_stride; // stride between V batch items = n_kv * hd * N + int o_head_stride; // stride between O heads = T * hd + int o_batch_stride; // stride between O batch items = n_h * T * hd + int lse_head_stride; // stride between LSE heads = T + int lse_batch_stride; // stride between LSE batch items = n_h * T }; template __global__ void __launch_bounds__(192) fmha_6warp_multihead_kernel(FmhaParams params) { static constexpr int NKT_QK = HD / MMA_K_BF16; - static constexpr int NKT_PV = SK_TILE / MMA_K_BF16; + static constexpr int NKT_PV = SK_TILE / MMA_K_BF16; // 8 static constexpr int N_NSUB = HD / 16; - static constexpr int TILE_SZ = 128 * MMA_K_BF16; - static constexpr int V_SUB_SZ = 256; + static constexpr int TILE_SZ = 128 * MMA_K_BF16; // 2048 BF16 + static constexpr int V_SUB_SZ = 256; // (16,16) canonical BF16 static constexpr int TMEM_N = (HD <= 128) ? 128 : 256; const int head_idx = blockIdx.y; @@ -69,18 +103,26 @@ fmha_6warp_multihead_kernel(FmhaParams params) { const int wid = tid / 32; const int lane = tid % 32; - const bool is_softmax_warp = (wid < 4); - const bool is_mma_warp = (wid == 4); - const bool is_load_warp = (wid == 5); + // Warp role predicates + const bool is_softmax_warp = (wid < 4); // Warps 0-3 + const bool is_mma_warp = (wid == 4); // Warp 4 + const bool is_load_warp = (wid == 5); // Warp 5 + // ================================================================== + // Compute per-head GMEM pointers + // ================================================================== const bf16_t* __restrict__ q_head = params.q - + head_idx * params.q_head_stride + batch_idx * params.q_batch_stride; + + head_idx * params.q_head_stride + + batch_idx * params.q_batch_stride; const bf16_t* __restrict__ k_head = params.k - + head_idx * params.k_head_stride + batch_idx * params.k_batch_stride; + + head_idx * params.k_head_stride + + batch_idx * params.k_batch_stride; const bf16_t* __restrict__ v_head = params.v - + head_idx * params.v_head_stride + batch_idx * params.v_batch_stride; + + head_idx * params.v_head_stride + + batch_idx * params.v_batch_stride; bf16_t* __restrict__ o_head = params.o - + head_idx * params.o_head_stride + batch_idx * params.o_batch_stride; + + head_idx * params.o_head_stride + + batch_idx * params.o_batch_stride; float* __restrict__ lse_head = params.lse ? params.lse + head_idx * params.lse_head_stride + batch_idx * params.lse_batch_stride @@ -88,12 +130,9 @@ fmha_6warp_multihead_kernel(FmhaParams params) { const int s_k = params.s_k; const float scale = params.scale; - const int n_kv_tiles = (params.n_kv_tiles > 0) ? params.n_kv_tiles - : (s_k + SK_TILE - 1) / SK_TILE; - const bool is_multi_tile = (n_kv_tiles > 1); // ================================================================ - // SMEM allocation + // SMEM allocation (shared across all warps) // ================================================================ extern __shared__ char sbuf[]; uint32_t* sTmemBase = (uint32_t*)sbuf; @@ -104,17 +143,10 @@ fmha_6warp_multihead_kernel(FmhaParams params) { bf16_t* sPk = (bf16_t*)(((uintptr_t)(sK0 + TILE_SZ) + 127) & ~(uintptr_t)127); bf16_t* sV = (bf16_t*)(((uintptr_t)(sPk + TILE_SZ) + 127) & ~(uintptr_t)127); float* s_p_vals = (float*)(sV + V_SUB_SZ); - // Multi-tile accumulator (only used when n_kv_tiles > 1) - float* sOacc = (float*)(((uintptr_t)(s_p_vals + SK_TILE) + 15) & ~(uintptr_t)15); - // Initialize multi-tile accumulator - if (is_multi_tile && tid == 0) { - *sRowMax = -INFINITY; - *sRowSum = 0.0f; - for (int d = 0; d < HD; d++) sOacc[d] = 0.0f; - } - - // TMEM allocation + // ================================================================ + // TMEM allocation (warp 4) + // ================================================================ if (is_mma_warp) { uint32_t smem_ptr = __cvta_generic_to_shared(sTmemBase); tmem_alloc(smem_ptr, TMEM_N); @@ -123,266 +155,153 @@ fmha_6warp_multihead_kernel(FmhaParams params) { uint32_t tb = *sTmemBase; // ================================================================ - // SINGLE-TILE PATH (n_kv_tiles <= 1) - // Identical to the pre-P5 kernel. Tested, proven correct. + // QK GEMM loop: for each K-tile, load Q+K, then MMA // ================================================================ - if (!is_multi_tile) { - // QK GEMM - for (int kt = 0; kt < NKT_QK; kt++) { - if (is_load_warp) { - for (int i = lane; i < TILE_SZ; i += 32) sQ0[i] = 0; + for (int kt = 0; kt < NKT_QK; kt++) { + // ---- Warp 5: Load Q and K for this K-tile ---- + if (is_load_warp) { + // Load Q K-tile: Q is (1, hd) for decode, row 0 only + for (int i = lane; i < TILE_SZ; i += 32) sQ0[i] = 0; + for (int d = lane; d < MMA_K_BF16; d += 32) { + int ck = d / 8, lc = d % 8; + sQ0[ck * 16 * 64 + lc] = q_head[kt * MMA_K_BF16 + d]; + } + // Load K K-tile: K is (s_k, hd) + for (int i = lane; i < TILE_SZ; i += 32) sK0[i] = 0; + for (int r = 0; r < s_k; r++) { for (int d = lane; d < MMA_K_BF16; d += 32) { int ck = d / 8, lc = d % 8; - sQ0[ck * 16 * 64 + lc] = q_head[kt * MMA_K_BF16 + d]; + int tmn = r / 8, lr = r % 8; + sK0[ck * 16 * 64 + tmn * 64 + lr * 8 + lc] = k_head[r * HD + kt * MMA_K_BF16 + d]; } - for (int i = lane; i < TILE_SZ; i += 32) sK0[i] = 0; - for (int r = 0; r < s_k; r++) { - for (int d = lane; d < MMA_K_BF16; d += 32) { - int ck = d / 8, lc = d % 8; - int tmn = r / 8, lr = r % 8; - sK0[ck * 16 * 64 + tmn * 64 + lr * 8 + lc] = k_head[r * HD + kt * MMA_K_BF16 + d]; - } - } - } - __syncthreads(); - if (is_mma_warp) { - 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); - asm volatile("tcgen05.fence::after_thread_sync;" ::: "memory"); - } - __syncthreads(); - } - - // Softmax (normalized P for single-tile) - if (wid == 0) { - float s_vals[SK_TILE], row_max = -INFINITY; - for (int n = 0; n < SK_TILE / 8; n++) { - float tmp[8]; - asm volatile("tcgen05.ld.sync.aligned.32x32b.x8.b32 {%0,%1,%2,%3,%4,%5,%6,%7},[%8];" - : "=f"(tmp[0]),"=f"(tmp[1]),"=f"(tmp[2]),"=f"(tmp[3]), - "=f"(tmp[4]),"=f"(tmp[5]),"=f"(tmp[6]),"=f"(tmp[7]) - : "r"(tb + n*8)); - asm volatile("tcgen05.wait::ld.sync.aligned;"); - if (lane == 0) for (int c=0;c<8;c++) { - s_vals[n*8+c] = tmp[c] * scale; - row_max = fmaxf(row_max, tmp[c] * scale); - } - } - row_max = wmax(row_max); - if (lane == 0) *sRowMax = row_max; - float row_sum = 0.0f; - if (lane == 0) for (int j=0;j 0); - asm volatile("tcgen05.fence::after_thread_sync;" ::: "memory"); - } - __syncthreads(); - } - } - - // Epilogue: TMEM → regs → BF16 → GMEM (P was normalized, no division needed) - if (wid == 0) { - float row_max = *sRowMax; - float row_sum = *sRowSum; - float o_vals[HD]; - for (int n = 0; n < HD / 8; n++) { - float tmp[8]; - asm volatile("tcgen05.ld.sync.aligned.32x32b.x8.b32 {%0,%1,%2,%3,%4,%5,%6,%7},[%8];" - : "=f"(tmp[0]),"=f"(tmp[1]),"=f"(tmp[2]),"=f"(tmp[3]), - "=f"(tmp[4]),"=f"(tmp[5]),"=f"(tmp[6]),"=f"(tmp[7]) - : "r"(tb + n*8)); - asm volatile("tcgen05.wait::ld.sync.aligned;"); - if (lane == 0) for (int c=0;c<8;c++) o_vals[n*8+c] = tmp[c]; - } - if (lane == 0) { - for (int d = 0; d < HD; d++) o_head[d] = f32_to_bf16(o_vals[d]); - if (lse_head) lse_head[0] = logf(row_sum) + row_max; } } __syncthreads(); - // ================================================================ - // MULTI-TILE PATH (n_kv_tiles > 1) - // FlashAttention-2 online softmax across KV tiles. - // ================================================================ - } else { - for (int kv_tile = 0; kv_tile < n_kv_tiles; kv_tile++) { - int kv_start = kv_tile * SK_TILE; - int kv_len = min(SK_TILE, s_k - kv_start); - - // QK GEMM (same as single-tile, but K offset = kv_start) - for (int kt = 0; kt < NKT_QK; kt++) { - if (is_load_warp) { - for (int i = lane; i < TILE_SZ; i += 32) sQ0[i] = 0; - for (int d = lane; d < MMA_K_BF16; d += 32) { - int ck = d / 8, lc = d % 8; - sQ0[ck * 16 * 64 + lc] = q_head[kt * MMA_K_BF16 + d]; - } - for (int i = lane; i < TILE_SZ; i += 32) sK0[i] = 0; - for (int r = 0; r < kv_len; r++) { - int g_r = kv_start + r; - for (int d = lane; d < MMA_K_BF16; d += 32) { - int ck = d / 8, lc = d % 8; - int tmn = r / 8, lr = r % 8; - sK0[ck * 16 * 64 + tmn * 64 + lr * 8 + lc] = k_head[g_r * HD + kt * MMA_K_BF16 + d]; - } - } - } - __syncthreads(); - if (is_mma_warp) { - 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); - asm volatile("tcgen05.fence::after_thread_sync;" ::: "memory"); - } - __syncthreads(); - } - - // Softmax: UN-NORMALIZED P (exp(s - max), NOT / sum) - if (wid == 0) { - float s_vals[SK_TILE], row_max = -INFINITY; - for (int n = 0; n < SK_TILE / 8; n++) { - float tmp[8]; - asm volatile("tcgen05.ld.sync.aligned.32x32b.x8.b32 {%0,%1,%2,%3,%4,%5,%6,%7},[%8];" - : "=f"(tmp[0]),"=f"(tmp[1]),"=f"(tmp[2]),"=f"(tmp[3]), - "=f"(tmp[4]),"=f"(tmp[5]),"=f"(tmp[6]),"=f"(tmp[7]) - : "r"(tb + n*8)); - asm volatile("tcgen05.wait::ld.sync.aligned;"); - if (lane == 0) for (int c=0;c<8;c++) { - s_vals[n*8+c] = tmp[c] * scale; - row_max = fmaxf(row_max, tmp[c] * scale); - } - } - row_max = wmax(row_max); - - float row_sum = 0.0f; - if (lane == 0) for (int j=0;j -INFINITY) { - float rescale = expf(old_max - row_max); - for (int d = 0; d < HD; d++) sOacc[d] *= rescale; - old_sum *= rescale; - } - *sRowMax = row_max; - *sRowSum = old_sum + row_sum; - } - - // Store UN-NORMALIZED P for PV - if (lane == 0) for (int j=0;j 0); + asm volatile("tcgen05.fence::after_thread_sync;" ::: "memory"); } __syncthreads(); } - // TMEM dealloc - if (is_mma_warp) tmem_dealloc(tb, TMEM_N); + // ================================================================ + // Softmax (warp 0, row 0 only for T=1 decode) + // ================================================================ + if (wid == 0) { + float s_vals[SK_TILE], row_max = -INFINITY; + for (int n = 0; n < SK_TILE / 8; n++) { + float tmp[8]; + asm volatile("tcgen05.ld.sync.aligned.32x32b.x8.b32 {%0,%1,%2,%3,%4,%5,%6,%7},[%8];" + : "=f"(tmp[0]),"=f"(tmp[1]),"=f"(tmp[2]),"=f"(tmp[3]), + "=f"(tmp[4]),"=f"(tmp[5]),"=f"(tmp[6]),"=f"(tmp[7]) + : "r"(tb + n*8)); + asm volatile("tcgen05.wait::ld.sync.aligned;"); + if (lane == 0) for (int c=0;c<8;c++) { + s_vals[n*8+c] = tmp[c] * scale; + row_max = fmaxf(row_max, tmp[c] * scale); + } + } + row_max = wmax(row_max); + if (lane == 0) *sRowMax = row_max; + float row_sum = 0.0f; + if (lane == 0) for (int j=0;j 0); + asm volatile("tcgen05.fence::after_thread_sync;" ::: "memory"); + } + __syncthreads(); + } + } + + // ================================================================ + // Epilogue: TMEM → regs → normalize → BF16 → GMEM + // For single-segment decode: normalize in-kernel. + // For multi-segment: emit un-normalized O + LSE. + // ================================================================ + if (wid == 0) { + float row_max = *sRowMax; + float row_sum = *sRowSum; + float o_vals[HD]; + for (int n = 0; n < HD / 8; n++) { + float tmp[8]; + asm volatile("tcgen05.ld.sync.aligned.32x32b.x8.b32 {%0,%1,%2,%3,%4,%5,%6,%7},[%8];" + : "=f"(tmp[0]),"=f"(tmp[1]),"=f"(tmp[2]),"=f"(tmp[3]), + "=f"(tmp[4]),"=f"(tmp[5]),"=f"(tmp[6]),"=f"(tmp[7]) + : "r"(tb + n*8)); + asm volatile("tcgen05.wait::ld.sync.aligned;"); + if (lane == 0) for (int c=0;c<8;c++) o_vals[n*8+c] = tmp[c]; + } + // P was NORMALIZED in softmax step. PV = P @ V is already the normalized + // attention output. No further division by row_sum needed. + // For single-segment decode, write O directly. + // LSE is written for multi-segment merge (P5). + if (lane == 0) { + for (int d = 0; d < HD; d++) { + o_head[d] = f32_to_bf16(o_vals[d]); + } + // Write LSE if pointer is valid + if (lse_head) { + // LSE = ln(row_sum) + row_max (natural log) + // This is the log of the softmax denominator, useful for + // multi-segment merge: O = sum(exp(lse_i) * O_i) / sum(exp(lse_i)) + lse_head[0] = logf(row_sum) + row_max; + } + } + } + __syncthreads(); + + // TMEM dealloc (warp 4) + if (is_mma_warp) { + tmem_dealloc(tb, TMEM_N); + } } } // namespace dsv4::kernels::attention diff --git a/dsv4/kernels/attention/fmha_multitile_capi.cu b/dsv4/kernels/attention/fmha_multitile_capi.cu new file mode 100644 index 00000000..e4ef4f95 --- /dev/null +++ b/dsv4/kernels/attention/fmha_multitile_capi.cu @@ -0,0 +1,146 @@ +/** + * DSV4 FMHA — Multi-tile kernel C API (TMA-based). + * + * Wraps fmha_6warp_tma_multirow_multitile_kernel with TMA descriptor + * creation and launch. Uses the same descriptor format as fmha_tma.cuh. + */ + +#include +#include +#include +#include + +#include "fmha_common.cuh" +#include "fmha_umma_desc.cuh" +#include "fmha_tma.cuh" +#include "fmha_6warp_tma_multirow_multitile.cuh" + +using namespace dsv4::kernels::attention; + +extern "C" { + +int fmha_multitile_smem_size(int hd) { + constexpr int HD_CHUNK = 256; + constexpr int TILE_SZ = 128 * MMA_K_BF16; + constexpr int V_SUB_SZ = 16 * MMA_K_BF16; + int hc = (hd <= 256) ? hd : HD_CHUNK; + size_t off = 0; + off += 4; // sTmemBase + off = (off + 127) & ~(size_t)127; // sMbar + off += 16; + 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 += TILE_SZ * 2; + off = (off + 127) & ~(size_t)127; // sV + off += V_SUB_SZ * 2; + off = (off + 127) & ~(size_t)127; // sOacc + off += 128 * hc * 4; + off += 128 * 4; // sRunningMax + off += 128 * 4; // sRunningSum + off += 128 * 4; // sTileRowMax + off += 128 * 4; // sTileRowSum + off += 256; // slack + return (int)((off + 127) & ~(size_t)127); +} + +/** + * Launch the multi-tile TMA FMHA kernel. + * + * Q: (batch, n_h, T, hd) BF16 contiguous + * K: (batch, n_h, N, hd) BF16 contiguous + * V: (batch, n_h, hd, N) BF16 contiguous + * O: (batch, n_h, T, hd) BF16 + * LSE: (batch, n_h, T) FP32 + */ +int fmha_multitile_decode_launch( + const void* q_ptr, + const void* k_ptr, + const void* v_ptr, + void* o_ptr, + void* lse_ptr, + int batch, int n_h, int T, int N, int hd, + int q_head_stride, int q_batch_stride, + int k_head_stride, int k_batch_stride, + int v_head_stride, int v_batch_stride, + int o_head_stride, int o_batch_stride, + int lse_head_stride, int lse_batch_stride, + float scale +) { + size_t desc_count = n_h * batch; + + CUtensorMap* d_tma_k; + CUtensorMap* d_tma_v; + cudaMalloc(&d_tma_k, desc_count * sizeof(CUtensorMap)); + cudaMalloc(&d_tma_v, desc_count * sizeof(CUtensorMap)); + + // Create TMA descriptors using the same format as fmha_tma.cuh + for (int b = 0; b < batch; b++) { + for (int h = 0; h < n_h; h++) { + const bf16_t* k_head = (const bf16_t*)k_ptr + h * k_head_stride + b * k_batch_stride; + const bf16_t* v_head = (const bf16_t*)v_ptr + h * v_head_stride + b * v_batch_stride; + + int idx = b * n_h + h; + + // K: (N, hd) → TMA desc with (cols=hd, rows=N), tile (16, 128) + if (!create_tma_desc_2d_bf16(d_tma_k + idx, k_head, N, hd, 128, 16)) { + cudaFree(d_tma_k); cudaFree(d_tma_v); + return -1; + } + + // V: (hd, N) → TMA desc with (cols=N, rows=hd), tile (16, 16) + if (!create_tma_desc_2d_bf16(d_tma_v + idx, v_head, hd, N, 16, 16)) { + cudaFree(d_tma_k); cudaFree(d_tma_v); + return -1; + } + } + } + + // Build params + FmhaTmaMultiRowMultiTileParams params; + params.q = (const bf16_t*)q_ptr; + params.tma_k = d_tma_k; + params.tma_v = d_tma_v; + params.o = (bf16_t*)o_ptr; + params.lse = (float*)lse_ptr; + params.s_k = N; + params.T = T; + params.n_h = n_h; + params.scale = scale; + params.q_head_stride = q_head_stride; + params.q_batch_stride = q_batch_stride; + params.o_head_stride = o_head_stride; + params.o_batch_stride = o_batch_stride; + params.lse_head_stride = lse_head_stride; + params.lse_batch_stride = lse_batch_stride; + + int smem = fmha_multitile_smem_size(hd); + dim3 grid(1, n_h, batch); + dim3 block(NTHREADS); + + if (smem > 48 * 1024) { + if (hd == 64) cudaFuncSetAttribute(fmha_6warp_tma_multirow_multitile_kernel<64>, cudaFuncAttributeMaxDynamicSharedMemorySize, smem); + else if (hd == 128) cudaFuncSetAttribute(fmha_6warp_tma_multirow_multitile_kernel<128>, cudaFuncAttributeMaxDynamicSharedMemorySize, smem); + else if (hd == 256) cudaFuncSetAttribute(fmha_6warp_tma_multirow_multitile_kernel<256>, cudaFuncAttributeMaxDynamicSharedMemorySize, smem); + else if (hd == 512) cudaFuncSetAttribute(fmha_6warp_tma_multirow_multitile_kernel<512>, cudaFuncAttributeMaxDynamicSharedMemorySize, smem); + } + + cudaError_t err; + if (hd == 64) fmha_6warp_tma_multirow_multitile_kernel<64><<>>(params); + else if (hd == 128) fmha_6warp_tma_multirow_multitile_kernel<128><<>>(params); + else if (hd == 256) fmha_6warp_tma_multirow_multitile_kernel<256><<>>(params); + else if (hd == 512) fmha_6warp_tma_multirow_multitile_kernel<512><<>>(params); + else { cudaFree(d_tma_k); cudaFree(d_tma_v); return -1; } + + err = cudaGetLastError(); + cudaFree(d_tma_k); + cudaFree(d_tma_v); + if (err != cudaSuccess) return (int)err; + return 0; +} + +} // extern "C" diff --git a/dsv4/kernels/attention/fmha_multitile_op.py b/dsv4/kernels/attention/fmha_multitile_op.py new file mode 100644 index 00000000..b246b108 --- /dev/null +++ b/dsv4/kernels/attention/fmha_multitile_op.py @@ -0,0 +1,134 @@ +"""DSV4 FMHA — Multi-tile TMA kernel loader. + +Loads the TMA-based multi-tile FMHA kernel via nvcc precompile + ctypes. +Supports any s_k (multiple KV tiles), T=1 decode + T>1 prefill. +HD in {64, 128, 256, 512}. +""" + +import torch +import logging +import os +import subprocess +import ctypes +from typing import Optional + +logger = logging.getLogger(__name__) + +KERNEL_DIR = os.path.dirname(os.path.abspath(__file__)) +REPO_ROOT = os.path.normpath(os.path.join(KERNEL_DIR, "..", "..")) +SOURCE = os.path.join(KERNEL_DIR, "fmha_multitile_capi.cu") +BUILD_DIR = os.path.join(REPO_ROOT, "build", "fmha_multitile") +SO_NAME = "libfmha_multitile.so" + +_lib = None +_lib_lock = False + + +def _find_nvcc(): + import shutil + for c in ["/usr/local/cuda-13.2/bin/nvcc", "/usr/local/cuda/bin/nvcc"]: + if os.path.isfile(c): return c + nvcc = shutil.which("nvcc") + if nvcc: return nvcc + raise RuntimeError("nvcc not found") + + +def _ensure_built(): + global _lib + if _lib is not None: return _lib + global _lib_lock + if _lib_lock: raise RuntimeError("Recursive build") + _lib_lock = True + try: + so_path = os.path.join(BUILD_DIR, SO_NAME) + need_build = True + if os.path.isfile(so_path): + src_mtime = os.path.getmtime(SOURCE) + for dep in ["fmha_common.cuh", "fmha_umma_desc.cuh", "fmha_tma.cuh", + "fmha_6warp_tma_multirow_multitile.cuh"]: + dp = os.path.join(KERNEL_DIR, dep) + if os.path.isfile(dp): src_mtime = max(src_mtime, os.path.getmtime(dp)) + need_build = src_mtime > os.path.getmtime(so_path) + if not need_build: + _lib = ctypes.CDLL(so_path) + return _lib + + logger.info("Building libfmha_multitile.so (sm_100a)...") + os.makedirs(BUILD_DIR, exist_ok=True) + nvcc = _find_nvcc() + cmd = [nvcc, "-std=c++20", "-shared", "-Xcompiler", "-fPIC", + "-gencode=arch=compute_100a,code=sm_100a", + "-gencode=arch=compute_100a,code=compute_100a", + f"-I{KERNEL_DIR}", f"-I{REPO_ROOT}", + "-O3", "--expt-relaxed-constexpr", + SOURCE, "-o", so_path, "-lcudart", "-lcuda"] + result = subprocess.run(cmd, capture_output=True, text=True) + if result.returncode != 0: + raise RuntimeError(f"nvcc failed:\n{result.stderr}") + _lib = ctypes.CDLL(so_path) + logger.info(f"Built {so_path}") + return _lib + finally: + _lib_lock = False + + +def fmha_multitile_decode_raw( + q: torch.Tensor, # (batch, n_h, T, hd) BF16 + k: torch.Tensor, # (batch, n_h, N, hd) BF16 + v: torch.Tensor, # (batch, n_h, hd, N) BF16 + scale: float, + n_comp: int = 0, + swa_len: int = 0, + is_causal: bool = False, + attn_sink: Optional[torch.Tensor] = None, +) -> tuple[torch.Tensor, torch.Tensor]: + """Launch the multi-tile TMA FMHA kernel. Returns (O, LSE).""" + lib = _ensure_built() + + B = q.shape[0] + n_h = q.shape[1] + T = q.shape[2] + hd = q.shape[3] + n_kv = k.shape[1] + N = k.shape[2] + + assert hd in (64, 128, 256, 512), f"Unsupported hd={hd}" + q_per_kv = n_h // n_kv + + # GQA: expand K/V to n_h heads + if n_kv < n_h: + k = k.repeat_interleave(q_per_kv, dim=1) + v = v.repeat_interleave(q_per_kv, dim=1) + + # Pad N to multiple of 128 + N_padded = ((N + 127) // 128) * 128 + if N < N_padded: + pad = N_padded - N + k = torch.cat([k, torch.zeros(B, k.shape[1], pad, hd, dtype=torch.bfloat16, device=k.device)], dim=2) + v = torch.cat([v, torch.zeros(v.shape[0], v.shape[1], hd, pad, dtype=torch.bfloat16, device=v.device)], dim=3) + N = N_padded + + k = k.contiguous() + v = v.contiguous() + q = q.contiguous() + + o = torch.zeros(B, n_h, T, hd, dtype=torch.bfloat16, device=q.device) + lse = torch.zeros(B, n_h, T, dtype=torch.float32, device=q.device) + + ret = lib.fmha_multitile_decode_launch( + ctypes.c_void_p(q.data_ptr()), + ctypes.c_void_p(k.data_ptr()), + ctypes.c_void_p(v.data_ptr()), + ctypes.c_void_p(o.data_ptr()), + ctypes.c_void_p(lse.data_ptr()), + ctypes.c_int(B), ctypes.c_int(n_h), ctypes.c_int(T), ctypes.c_int(N), ctypes.c_int(hd), + ctypes.c_int(q.stride(1)), ctypes.c_int(q.stride(0)), + ctypes.c_int(k.stride(1)), ctypes.c_int(k.stride(0)), + ctypes.c_int(v.stride(1)), ctypes.c_int(v.stride(0)), + ctypes.c_int(o.stride(1)), ctypes.c_int(o.stride(0)), + ctypes.c_int(lse.stride(1)), ctypes.c_int(lse.stride(0)), + ctypes.c_float(scale), + ) + if ret != 0: + raise RuntimeError(f"Multi-tile kernel failed: {ret}") + return o, lse diff --git a/dsv4/kernels/attention/production.py b/dsv4/kernels/attention/production.py index 33582dc4..fec51f0f 100644 --- a/dsv4/kernels/attention/production.py +++ b/dsv4/kernels/attention/production.py @@ -351,6 +351,33 @@ def _dsv4_attention_fast_decode( return o_4d.squeeze(0) # (n_q, 1, hd) +def _dsv4_attention_multitile( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + scale: float, + n_comp: int = 0, + sink_bias: Optional[torch.Tensor] = None, +) -> torch.Tensor: + """Multi-tile decode via TMA-based 6-warp FMHA kernel (N > 128).""" + from dsv4.kernels.attention.fmha_multitile_op import fmha_multitile_decode_raw + + n_q, T, hd = q.shape + n_kv = k.shape[0] if k.dim() == 3 else 1 + N = k.shape[-2] if k.dim() == 3 else k.shape[0] + + q_4d = q.unsqueeze(0).contiguous() + if k.dim() == 2: + k_4d = k.unsqueeze(0).unsqueeze(0).contiguous() + v_4d = v.unsqueeze(0).unsqueeze(0).transpose(-1, -2).contiguous() + else: + k_4d = k.unsqueeze(0).contiguous() + v_4d = v.unsqueeze(0).transpose(-1, -2).contiguous() + + o_4d, _lse = fmha_multitile_decode_raw(q_4d, k_4d, v_4d, scale) + return o_4d.squeeze(0) + + # --------------------------------------------------------------------------- # Public API # --------------------------------------------------------------------------- @@ -425,8 +452,11 @@ def dsv4_attention( s_k_per_seg = 128 n_segments = (N + s_k_per_seg - 1) // s_k_per_seg - if T == 1 and hd in (64, 128, 256): - return _dsv4_attention_fast_decode(q, k, v, scale, n_comp, sink_bias) + if T == 1 and hd in (64, 128, 256, 512): + if n_segments == 1 and hd in (64, 128, 256): + return _dsv4_attention_fast_decode(q, k, v, scale, n_comp, sink_bias) + else: + return _dsv4_attention_multitile(q, k, v, scale, n_comp, sink_bias) # ================================================================== # SLOW PATH: CuTeDSL kernel + Python KV merge