From d5c008673797b961ba3ca3d21a0e68f03b7769f5 Mon Sep 17 00:00:00 2001 From: biondizzle Date: Sat, 30 May 2026 08:19:16 +0000 Subject: [PATCH] P3: fix SMEM computation, pad K/V to 128, remove stale files - fmha_multihead_capi.cu: SMEM formula matches standalone test Added cudaFuncSetAttribute for dynamic SMEM > 48KB - fmha_multihead_op.py: pad K/V to N=128 when N<128 (kernel softmax loop is hardcoded to SK_TILE=128) - Removed fmha_multihead_launch.cu (ATen approach, didn't work) - Removed test_p3_ctypes_minimal.py (superseded by main test) --- dsv4/kernels/attention/fmha_multihead_capi.cu | 71 +++++--- .../attention/fmha_multihead_launch.cu | 167 ------------------ dsv4/kernels/attention/fmha_multihead_op.py | 22 ++- tests/unit/test_p3_ctypes_minimal.py | 130 -------------- 4 files changed, 62 insertions(+), 328 deletions(-) delete mode 100644 dsv4/kernels/attention/fmha_multihead_launch.cu delete mode 100644 tests/unit/test_p3_ctypes_minimal.py diff --git a/dsv4/kernels/attention/fmha_multihead_capi.cu b/dsv4/kernels/attention/fmha_multihead_capi.cu index fff23852..e643506f 100644 --- a/dsv4/kernels/attention/fmha_multihead_capi.cu +++ b/dsv4/kernels/attention/fmha_multihead_capi.cu @@ -1,41 +1,40 @@ /** * DSV4 FMHA Multi-Head — C API for ctypes loading. * - * This provides a pure C API (no pybind11, no ATen) so we can compile - * with nvcc -arch=sm_100a and load via Python ctypes. PyTorch tensors - * are passed as raw device pointers + shapes + strides. - * - * The Python side handles tensor → (ptr, shape, stride) conversion - * and wraps the result back into torch tensors. + * Pure C API (no pybind11, no ATen) so we can compile with nvcc -arch=sm_100a + * and load via Python ctypes. PyTorch tensors are passed as raw device + * pointers + shapes + strides. */ #include #include -// Forward declaration of the kernel (from fmha_6warp_multihead.cuh) -// We need the template instantiations - #include "fmha_common.cuh" #include "fmha_umma_desc.cuh" #include "fmha_6warp_multihead.cuh" extern "C" { -/** Compute dynamic SMEM size for the 6-warp multi-head kernel. */ +/** + * Compute dynamic SMEM size for the 6-warp multi-head kernel. + * Matches the standalone test's formula exactly. + */ int fmha_compute_smem(int hd) { using namespace dsv4::kernels::attention; - int base = 4 + 4 + 4; - base = (base + 15) & ~15; - int sQ0_sz = 128 * hd * 2; - int sK0_sz = 128 * hd * 2; - int sPk_offset = ((base + sQ0_sz + sK0_sz) + 127) & ~127; - int sPk_sz = 128 * 16 * 2; - int sV_offset = ((sPk_offset + sPk_sz) + 127) & ~127; - int sV_sz = 16 * 16 * 2; - int sp_vals_offset = sV_offset + sV_sz; - int sp_vals_sz = 128 * 4; - int total = sp_vals_offset + sp_vals_sz; - return (total + 127) & ~127; + constexpr int SK = 128; + constexpr int TILE_SZ = 128 * MMA_K_BF16; // 2048 BF16 + constexpr int V_SUB_SZ = 256; // (16,16) canonical BF16 + // 4: tmem_base + // 8: row_max + row_sum + // 16: alignment padding + // TILE_SZ*2: sQ (BF16 bytes) + // TILE_SZ*2: sK (BF16 bytes) + // TILE_SZ*2: sPk (BF16 bytes) — always 128*16 in the kernel + // V_SUB_SZ*2: sV (BF16 bytes) + // SK*4: s_p_vals (float bytes) + // 256: extra alignment + slack + int smem = (4 + 8 + 16 + TILE_SZ*2 + TILE_SZ*2 + TILE_SZ*2 + V_SUB_SZ*2 + SK*4 + 256 + 127) & ~127; + return smem; } /** @@ -46,11 +45,11 @@ int fmha_compute_smem(int hd) { * Returns 0 on success, non-zero on error. */ int fmha_multihead_decode_launch( - const void* q_ptr, // Q base: (batch, n_h, 1, hd) BF16 - const void* k_ptr, // K base: (batch, n_kv, N, hd) BF16 - const void* v_ptr, // V base: (batch, n_kv, hd, N) BF16 - void* o_ptr, // O base: (batch, n_h, 1, hd) BF16 - void* lse_ptr, // LSE base: (batch, n_h, 1) FP32 (can be NULL) + const void* q_ptr, + const void* k_ptr, + const void* v_ptr, + void* o_ptr, + void* lse_ptr, int batch, int n_h, int n_kv, int N, int hd, int q_head_stride, int q_batch_stride, int k_head_stride, int k_batch_stride, @@ -85,6 +84,22 @@ int fmha_multihead_decode_launch( dim3 grid(1, n_h, batch); dim3 block(NTHREADS); + // Increase SMEM limit if needed + if (smem > 48 * 1024) { + cudaError_t err; + if (hd == 64) { + err = cudaFuncSetAttribute(fmha_6warp_multihead_kernel<64, 128>, + cudaFuncAttributeMaxDynamicSharedMemorySize, smem); + } else if (hd == 128) { + err = cudaFuncSetAttribute(fmha_6warp_multihead_kernel<128, 128>, + cudaFuncAttributeMaxDynamicSharedMemorySize, smem); + } else if (hd == 256) { + err = cudaFuncSetAttribute(fmha_6warp_multihead_kernel<256, 128>, + cudaFuncAttributeMaxDynamicSharedMemorySize, smem); + } + if (err != cudaSuccess) return (int)err; + } + cudaError_t err; if (hd == 64) { fmha_6warp_multihead_kernel<64, 128><<>>(params); @@ -93,7 +108,7 @@ int fmha_multihead_decode_launch( } else if (hd == 256) { fmha_6warp_multihead_kernel<256, 128><<>>(params); } else { - return -1; // unsupported hd + return -1; } err = cudaGetLastError(); diff --git a/dsv4/kernels/attention/fmha_multihead_launch.cu b/dsv4/kernels/attention/fmha_multihead_launch.cu deleted file mode 100644 index ebd68c42..00000000 --- a/dsv4/kernels/attention/fmha_multihead_launch.cu +++ /dev/null @@ -1,167 +0,0 @@ -/** - * DSV4 FMHA Multi-Head — PyTorch launch wrapper. - * - * Bridges the raw CUDA 6-warp multi-head FMHA kernel to PyTorch tensors. - * Uses uint16_t (bf16_t) inside the kernel, c10::BFloat16 at the API boundary. - * Casts are zero-cost reinterpret_casts (same bit layout). - * - * Grid: dim3(1, n_h, batch_size) - * Each CTA processes one (head, batch) pair independently. - * Decode-only: T=1, single KV segment (s_k <= 128). - */ - -#include -#include -#include - -// Pull in the kernel (uses uint16_t bf16_t internally, no ATen deps) -#include "fmha_common.cuh" -#include "fmha_umma_desc.cuh" -#include "fmha_6warp_multihead.cuh" - -namespace dsv4::kernels::attention { - -/** Compute dynamic SMEM size for the 6-warp multi-head kernel. */ -static int compute_smem_multihead(int hd) { - // From fmha_6warp_multihead.cuh SMEM layout: - // sTmemBase(4) + sRowMax(4) + sRowSum(4) + padding to 16B - // + sQ0: 128 * hd * 2 - // + sK0: 128 * hd * 2 (128B aligned) - // + sPk: 128 * 16 * 2 = 4096 (128B aligned, always TILE_SZ for the P buffer) - // + sV: 16 * 16 * 2 = 512 (128B aligned, V_SUB_SZ) - // + s_p_vals: 128 * 4 = 512 - // Over-allocate: sum the worst case - int base = 4 + 4 + 4; // tmem_base, row_max, row_sum - base = (base + 15) & ~15; // align to 16B - int sQ0_offset = base; - int sQ0_sz = 128 * hd * 2; - int sK0_offset = sQ0_offset + sQ0_sz; - int sK0_sz = 128 * hd * 2; - // Pk needs 128B alignment - int sPk_offset = ((sK0_offset + sK0_sz) + 127) & ~127; - int sPk_sz = 128 * 16 * 2; // TILE_SZ in BF16 - // V needs 128B alignment - int sV_offset = ((sPk_offset + sPk_sz) + 127) & ~127; - int sV_sz = 16 * 16 * 2; // V_SUB_SZ in BF16 - int sp_vals_offset = sV_offset + sV_sz; - int sp_vals_sz = 128 * 4; // SK_TILE * sizeof(float) - int total = sp_vals_offset + sp_vals_sz; - // Pad to 128B alignment - total = (total + 127) & ~127; - return total; -} - -/** - * Launch the 6-warp multi-head FMHA kernel for decode (T=1). - * - * Q: (batch, n_h, 1, hd) contiguous BF16 - * K: (batch, n_kv, N, hd) contiguous BF16 - * V: (batch, n_kv, hd, N) contiguous BF16 — NOTE: V is (hd, N) per head! - * Returns: (O, LSE) where O is (batch, n_h, 1, hd) BF16, LSE is (batch, n_h, 1) FP32 - * - * For MQA: n_kv = 1, k/v head stride = 0 (broadcast) - * For GQA: n_kv < n_h, k/v head stride reflects the KV group layout - */ -std::tuple fmha_multihead_decode_cuda( - torch::Tensor q, // (batch, n_h, 1, hd) BF16 - torch::Tensor k, // (batch, n_kv, N, hd) BF16 - torch::Tensor v, // (batch, n_kv, hd, N) BF16 - double scale, // 1/sqrt(hd) - int64_t n_comp, // compressed KV length (0 = no compression) - int64_t swa_len, // sliding window length (0 = no SWA mask) - bool is_causal, // causal mask on SWA - c10::optional attn_sink // (batch, n_h) FP32 — per-head sink bias -) { - TORCH_CHECK(q.dim() == 4, "Q must be 4D (batch, n_h, 1, hd)"); - TORCH_CHECK(q.size(2) == 1, "Decode mode requires T=1"); - TORCH_CHECK(k.dim() == 4, "K must be 4D (batch, n_kv, N, hd)"); - TORCH_CHECK(v.dim() == 4, "V must be 4D (batch, n_kv, hd, N)"); - - const int B = q.size(0); - const int n_h = q.size(1); - const int hd = q.size(3); - const int n_kv = k.size(1); - const int N = k.size(2); - - TORCH_CHECK(N <= 128, "Decode fast path requires N <= 128 (single KV tile)"); - TORCH_CHECK(hd == 64 || hd == 128 || hd == 256, - "Unsupported head_dim: ", hd, " (decode fast path: 64/128/256)"); - - auto opts_bf16 = q.options().dtype(torch::kBFloat16); - auto opts_f32 = q.options().dtype(torch::kFloat32); - auto o = torch::zeros({B, n_h, 1, hd}, opts_bf16); - auto lse = torch::zeros({B, n_h, 1}, opts_f32); - - // Build FmhaParams - FmhaParams params; - params.q = reinterpret_cast(q.data_ptr()); - params.k = reinterpret_cast(k.data_ptr()); - params.v = reinterpret_cast(v.data_ptr()); - params.o = reinterpret_cast(o.data_ptr()); - params.lse = lse.data_ptr(); - params.s_k = N; - params.scale = static_cast(scale); - params.head_dim = hd; - - // Strides in BF16 elements - params.q_head_stride = q.stride(1); // stride between heads in the batch - params.q_batch_stride = q.stride(0); // stride between batch items - params.k_head_stride = k.stride(1); - params.k_batch_stride = k.stride(0); - params.v_head_stride = v.stride(1); - params.v_batch_stride = v.stride(0); - params.o_head_stride = o.stride(1); - params.o_batch_stride = o.stride(0); - params.lse_head_stride = lse.stride(1); - params.lse_batch_stride = lse.stride(0); - - // Zero out head strides for MQA (n_kv == 1 means all Q heads share same K/V) - // Actually: the kernel uses head_idx * k_head_stride, which for MQA should give - // the same address for all heads. If k has n_kv=1, then stride(1) may be N*hd - // (the single head's total). For MQA we want k_head_stride=0 so all heads point - // to the same K. The caller should ensure K has shape (batch, 1, N, hd) for MQA, - // and we set k_head_stride=0 when n_kv==1. - if (n_kv == 1) { - params.k_head_stride = 0; - params.v_head_stride = 0; - } - - // For n_comp=0, the kernel doesn't need sink bias, but we pass nullptr - // The kernel's epilogue only writes LSE if lse_head != nullptr - - // Copy params to device (constant memory or managed) - // Simpler: just pass by value through kernel parameter struct - // CUDA kernels can take structs by value - int smem = compute_smem_multihead(hd); - - dim3 grid(1, n_h, B); - dim3 block(NTHREADS); - - // Launch with template instantiation - #define LAUNCH(D) \ - fmha_6warp_multihead_kernel<<>>(params) - - cudaError_t err; - if (hd == 64) { - LAUNCH(64); - } else if (hd == 128) { - LAUNCH(128); - } else if (hd == 256) { - LAUNCH(256); - } else { - TORCH_CHECK(false, "Unsupported head_dim: ", hd); - } - #undef LAUNCH - - err = cudaGetLastError(); - TORCH_CHECK(err == cudaSuccess, "Kernel launch failed: ", cudaGetErrorString(err)); - - return std::make_tuple(o, lse); -} - -} // namespace dsv4::kernels::attention - -PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { - m.def("fmha_multihead_decode", &dsv4::kernels::attention::fmha_multihead_decode_cuda, - "DSV4 FMHA multi-head decode kernel (6-warp, T=1, single KV tile)"); -} diff --git a/dsv4/kernels/attention/fmha_multihead_op.py b/dsv4/kernels/attention/fmha_multihead_op.py index a0695f62..255e1615 100644 --- a/dsv4/kernels/attention/fmha_multihead_op.py +++ b/dsv4/kernels/attention/fmha_multihead_op.py @@ -140,9 +140,25 @@ def fmha_multihead_decode_raw( assert q.shape[2] == 1, f"Decode requires T=1, got T={q.shape[2]}" assert hd in (64, 128, 256), f"Unsupported hd={hd}" assert N <= 128, f"Decode fast path requires N<=128, got N={N}" - assert q.is_contiguous() - assert k.is_contiguous() - assert v.is_contiguous() + + # The kernel template has SK_TILE=128 hardcoded in the softmax loop. + # When N < 128, pad K and V to 128 so the kernel processes zeros for + # the extra positions (correctly gets zero attention weight after softmax). + if N < 128: + pad_len = 128 - N + k_padded = torch.cat([k, + torch.zeros(B, n_kv, pad_len, hd, dtype=torch.bfloat16, device=k.device)], dim=2) + # V is (B, n_kv, hd, N) — pad the last dim + v_padded = torch.cat([v, + torch.zeros(B, n_kv, hd, pad_len, dtype=torch.bfloat16, device=v.device)], dim=3) + k = k_padded.contiguous() + v = v_padded.contiguous() + N = 128 # Tell kernel we have 128 positions + else: + k = k.contiguous() + v = v.contiguous() + + q = q.contiguous() o = torch.zeros(B, n_h, 1, hd, dtype=torch.bfloat16, device=q.device) lse = torch.zeros(B, n_h, 1, dtype=torch.float32, device=q.device) diff --git a/tests/unit/test_p3_ctypes_minimal.py b/tests/unit/test_p3_ctypes_minimal.py deleted file mode 100644 index 3e3bd7cd..00000000 --- a/tests/unit/test_p3_ctypes_minimal.py +++ /dev/null @@ -1,130 +0,0 @@ -""" -Minimal ctypes test: exact same setup as standalone test_fmha_6warp_multihead_hd64.cu -Uses raw CUDA memory, not PyTorch tensors, to isolate kernel correctness. -""" -import torch -import ctypes -import math -import os -import sys -import subprocess - -sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) - -from dsv4.kernels.attention.fmha_multihead_op import _find_nvcc, _ensure_built, BUILD_DIR, SO_NAME - -def f32_to_bf16_bits(f): - """Convert float to BF16 bit pattern (uint16).""" - import struct - u = struct.unpack('I', struct.pack('f', f))[0] - return (u >> 16) & 0xFFFF - -def bf16_bits_to_f32(h): - """Convert BF16 bit pattern (uint16) to float.""" - import struct - u = h << 16 - return struct.unpack('f', struct.pack('I', u))[0] - -def test_minimal(): - # Build the .so - lib = _ensure_built() - - hd = 64 - n_h = 4 - N = 128 # SK - batch = 1 - scale = 1.0 / math.sqrt(hd) - - # Create data on GPU using PyTorch (easier than raw CUDA malloc for setup) - torch.manual_seed(42) - - # Q: (batch, n_h, 1, hd) — each head has 1 row of hd elements - q_data = torch.randn(n_h, hd, dtype=torch.bfloat16, device='cuda') - q_4d = q_data.unsqueeze(0).unsqueeze(2).contiguous() # (1, n_h, 1, hd) - - # K: (batch, n_h, N, hd) — each head has N rows of hd elements - k_data = torch.randn(n_h, N, hd, dtype=torch.bfloat16, device='cuda') - k_4d = k_data.unsqueeze(0).contiguous() # (1, n_h, N, hd) - - # V: (batch, n_h, hd, N) — transposed - v_data = torch.randn(n_h, hd, N, dtype=torch.bfloat16, device='cuda') - v_4d = v_data.unsqueeze(0).contiguous() # (1, n_h, hd, N) - - # Output - o_4d = torch.zeros(1, n_h, 1, hd, dtype=torch.bfloat16, device='cuda') - lse_4d = torch.zeros(1, n_h, 1, dtype=torch.float32, device='cuda') - - # Strides - q_hs = q_4d.stride(1) # hd - q_bs = q_4d.stride(0) # n_h * 1 * hd = n_h * hd - k_hs = k_4d.stride(1) # N * hd - k_bs = k_4d.stride(0) # n_h * N * hd - v_hs = v_4d.stride(1) # hd * N - v_bs = v_4d.stride(0) # n_h * hd * N - o_hs = o_4d.stride(1) - o_bs = o_4d.stride(0) - lse_hs = lse_4d.stride(1) - lse_bs = lse_4d.stride(0) - - print(f"Q shape: {q_4d.shape}, strides: {q_4d.stride()}") - print(f"K shape: {k_4d.shape}, strides: {k_4d.stride()}") - print(f"V shape: {v_4d.shape}, strides: {v_4d.stride()}") - print(f"O shape: {o_4d.shape}, strides: {o_4d.stride()}") - print(f"LSE shape: {lse_4d.shape}, strides: {lse_4d.stride()}") - print(f"q_hs={q_hs}, q_bs={q_bs}, k_hs={k_hs}, k_bs={k_bs}") - print(f"v_hs={v_hs}, v_bs={v_bs}, o_hs={o_hs}, o_bs={o_bs}") - print(f"lse_hs={lse_hs}, lse_bs={lse_bs}") - - ret = lib.fmha_multihead_decode_launch( - ctypes.c_void_p(q_4d.data_ptr()), - ctypes.c_void_p(k_4d.data_ptr()), - ctypes.c_void_p(v_4d.data_ptr()), - ctypes.c_void_p(o_4d.data_ptr()), - ctypes.c_void_p(lse_4d.data_ptr()), - ctypes.c_int(batch), - ctypes.c_int(n_h), - ctypes.c_int(n_h), # n_kv = n_h for MHA - ctypes.c_int(N), - ctypes.c_int(hd), - ctypes.c_int(q_hs), - ctypes.c_int(q_bs), - ctypes.c_int(k_hs), - ctypes.c_int(k_bs), - ctypes.c_int(v_hs), - ctypes.c_int(v_bs), - ctypes.c_int(o_hs), - ctypes.c_int(o_bs), - ctypes.c_int(lse_hs), - ctypes.c_int(lse_bs), - ctypes.c_float(scale), - ) - print(f"Kernel return: {ret}") - - # Reference: pure PyTorch - o_ref = torch.zeros(n_h, 1, hd, dtype=torch.bfloat16, device='cuda') - for h in range(n_h): - q_h = q_data[h:h+1] # (1, hd) - k_h = k_data[h] # (N, hd) - v_h = v_data[h].T # (N, hd) — V is (hd, N), transpose to (N, hd) - s = torch.matmul(q_h.float(), k_h.float().T) * scale # (1, N) - s = torch.softmax(s, dim=-1) - o = torch.matmul(s, v_h.float()) # (1, hd) - o_ref[h] = o.bfloat16() - - # Compare - o_kernel = o_4d.squeeze(0).squeeze(1) # (n_h, hd) - o_ref_flat = o_ref.squeeze(1) # (n_h, hd) - - for h in range(n_h): - cos = torch.nn.functional.cosine_similarity( - o_kernel[h].float().unsqueeze(0), - o_ref_flat[h].float().unsqueeze(0) - ).item() - print(f" Head {h}: cos={cos:.6f}") - - torch.cuda.synchronize() - print("Done") - - -if __name__ == "__main__": - test_minimal()