From a01e8bbe984a230d5405b461394ea5616a67c8a3 Mon Sep 17 00:00:00 2001 From: ViperEkura <3081035982@qq.com> Date: Tue, 21 Jul 2026 21:52:46 +0800 Subject: [PATCH] refactor: adopt FA2-style KernelTraits + compile-time causal/mask dispatch - Introduce KernelTraits compile-time config bundle, replacing scattered template params - Template all MMA and scalar kernels on IsCausal/HasMask bools to eliminate inner-loop runtime branches - Dispatch to 4-path IsCausal/HasMask kernel variants at entry points based on p.causal_offset and p.use_mask - Update standalone test files with new kernel signatures, add causal test cases - Fix duplicate using bf16 in MMA kernels that include attn_mma_utils.cuh --- csrc/kernels/attn_decode.cu | 61 ++++-- csrc/kernels/attn_decode_split_kv.cuh | 30 +-- csrc/kernels/attn_decode_split_kv_mma.cuh | 127 +++++------- csrc/kernels/attn_mma_utils.cuh | 152 +++++++------- csrc/kernels/attn_paged_decode.cu | 54 +++-- csrc/kernels/attn_paged_decode_split_kv.cuh | 24 ++- .../attn_paged_decode_split_kv_mma.cuh | 115 +++++------ csrc/kernels/attn_prefill.cu | 56 ++++-- csrc/kernels/attn_prefill_split_q.cuh | 37 ++-- csrc/kernels/attn_prefill_split_q_mma.cuh | 186 +++++++---------- csrc/tests/attn_decode_test.cu | 189 ++++++++++-------- csrc/tests/attn_paged_decode_test.cu | 99 +++++---- csrc/tests/attn_prefill_test.cu | 163 +++++++++------ 13 files changed, 682 insertions(+), 611 deletions(-) diff --git a/csrc/kernels/attn_decode.cu b/csrc/kernels/attn_decode.cu index 9d7fcc8..3a8d802 100644 --- a/csrc/kernels/attn_decode.cu +++ b/csrc/kernels/attn_decode.cu @@ -3,9 +3,26 @@ #ifndef ASTRAI_NO_MMA #include "attn_decode_split_kv_mma.cuh" + +template +static void launch_mma_decode_impl(AttentionParams& p) { + using Traits = KernelTraits; + int tiles_total = (p.kv_len + BC - 1) / BC; + p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total); + alloc_split_partials(p); + + attn_decode_split_kv_mma_kernel<<>>(p); + attn_decode_combine_kernel<<>>(p); +} + +template +static void launch_mma_decode(AttentionParams& p) { + constexpr int STAGES = (HEAD_DIM <= 128) ? 2 : 1; + launch_mma_decode_impl(p); +} #endif -// Scalar fallback: one warp per query head, split-KV across grid.z. +template static void launch_scalar_decode(AttentionParams& p) { int group_size = p.q_head / p.kv_head; int chunks_total = (p.kv_len + DC_CHUNK - 1) / DC_CHUNK; @@ -13,37 +30,38 @@ static void launch_scalar_decode(AttentionParams& p) { alloc_split_partials(p); size_t smem = DC_CHUNK * p.head_dim * sizeof(bf16); - attn_decode_split_kv_kernel<<>>(p); + + dim3 grid(p.batch * p.kv_head, 1, p.num_splits); + dim3 block(32, group_size); + attn_decode_split_kv_kernel<<>>(p); attn_decode_combine_kernel<<>>(p); } -#ifndef ASTRAI_NO_MMA -// MMA head-packing requires G <= 16 (BR=16 rows). sm_80+ tensor-core -// + cp.async wins even at G=1 (decode is memory-bound, not compute-bound). -// STAGES=2 (double-buffer) for D<=128 (smem 16 KB); STAGES=1 for D=256 -// (double-buffer would be 32 KB, near the 48 KB static cap — keep single -// to preserve occupancy). -template -static void launch_mma_decode(AttentionParams& p) { - int tiles_total = (p.kv_len + BC - 1) / BC; - p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total); - alloc_split_partials(p); - - attn_decode_split_kv_mma_kernel<<>>(p); - attn_decode_combine_kernel<<>>(p); -} -#endif - template static void dispatch_decode(AttentionParams& p) { + bool is_causal = (p.causal_offset >= 0); + bool has_mask = (p.use_mask && p.mask); + #ifndef ASTRAI_NO_MMA int G = p.q_head / p.kv_head; if (G >= 1 && G <= 16) { - launch_mma_decode(p); + if (is_causal) { + if (has_mask) launch_mma_decode(p); + else launch_mma_decode(p); + } else { + if (has_mask) launch_mma_decode(p); + else launch_mma_decode(p); + } return; } #endif - launch_scalar_decode(p); + if (is_causal) { + if (has_mask) launch_scalar_decode(p); + else launch_scalar_decode(p); + } else { + if (has_mask) launch_scalar_decode(p); + else launch_scalar_decode(p); + } } torch::Tensor attn_decode( @@ -60,7 +78,6 @@ torch::Tensor attn_decode( TORCH_CHECK(p.q_len == 1, "Q seq_len must be 1"); TORCH_CHECK(p.head_dim % 32 == 0, "head_dim must be multiple of 32"); - // O matches Q's original layout auto O = torch::empty_strided(q.sizes(), q.strides(), q.options()); auto O_view = (layout == 1) ? O.transpose(1, 2) : O; p.o = (bf16*)O_view.data_ptr(); diff --git a/csrc/kernels/attn_decode_split_kv.cuh b/csrc/kernels/attn_decode_split_kv.cuh index b0ed646..6fc39dc 100644 --- a/csrc/kernels/attn_decode_split_kv.cuh +++ b/csrc/kernels/attn_decode_split_kv.cuh @@ -12,6 +12,7 @@ __device__ inline float warp_reduce_sum(float val) { return val; } +template __global__ void attn_decode_split_kv_kernel(AttentionParams p) { int batch = blockIdx.x / p.kv_head; int kv_head = blockIdx.x % p.kv_head; @@ -48,7 +49,8 @@ __global__ void attn_decode_split_kv_kernel(AttentionParams p) { // Load K into shared memory (gather from strided global) int total = this_chunk * p.head_dim; - for (int i = threadIdx.y * 32 + lane; i < total; i += blockDim.x * blockDim.y) { + for (int i = threadIdx.y * 32 + lane; i < total; + i += blockDim.x * blockDim.y) { int s = i / p.head_dim; int d_dim = i % p.head_dim; int kv_idx = chunk_start + s; @@ -60,24 +62,30 @@ __global__ void attn_decode_split_kv_kernel(AttentionParams p) { for (int s = 0; s < this_chunk; s++) { float partial = 0.0f; for (int i = 0; i < hd_per_thread; i++) - partial += q_reg[i] * __bfloat162float(k_smem[s * p.head_dim + lane * hd_per_thread + i]); + partial += q_reg[i] * __bfloat162float( + k_smem[s * p.head_dim + lane * hd_per_thread + i]); partial = warp_reduce_sum(partial) * p.scale; int kv_idx = chunk_start + s; - if (p.use_mask && p.mask && !p.mask[mask_base + kv_idx]) - partial = -FLT_MAX; - if (p.causal_offset >= 0 && kv_idx > p.causal_offset) - partial = -FLT_MAX; + if constexpr (HasMask) { + if (!p.mask[mask_base + kv_idx]) + partial = -FLT_MAX; + } + if constexpr (IsCausal) { + if (kv_idx > p.causal_offset) + partial = -FLT_MAX; + } float new_m = fmaxf(m, partial); float alpha = expf(m - new_m); float beta = expf(partial - new_m); d = d * alpha + beta; - // V: stride-based read - int v_off = kv_base + kv_idx * p.kv_stride_l + lane * hd_per_thread * p.kv_stride_d; + int v_off = kv_base + kv_idx * p.kv_stride_l + + lane * hd_per_thread * p.kv_stride_d; for (int i = 0; i < hd_per_thread; i++) - acc_reg[i] = acc_reg[i] * alpha + __bfloat162float(p.v[v_off + i * p.kv_stride_d]) * beta; + acc_reg[i] = acc_reg[i] * alpha + + __bfloat162float(p.v[v_off + i * p.kv_stride_d]) * beta; m = new_m; } __syncthreads(); @@ -97,9 +105,6 @@ __global__ void attn_decode_split_kv_kernel(AttentionParams p) { } } -// Reduce split-K partials into the final bf16 output. One block per (batch, -// q_head); each thread folds across all splits with a single-pass -// online-rescale reduction (expf + FMA counts halved vs 3-pass original). __global__ void attn_decode_combine_kernel(AttentionParams p) { int bh = blockIdx.x; int d = threadIdx.x; @@ -126,7 +131,6 @@ __global__ void attn_decode_combine_kernel(AttentionParams p) { } float inv = (l > 1e-20f) ? (1.0f / l) : 0.0f; - // Stride-based output write (q_len=1 for decode, so stride_l not needed) int o_off = batch * p.q_stride_b + q_head * p.q_stride_h + d * p.q_stride_d; p.o[o_off] = __float2bfloat16(acc * inv); } diff --git a/csrc/kernels/attn_decode_split_kv_mma.cuh b/csrc/kernels/attn_decode_split_kv_mma.cuh index a78a3e1..b88daef 100644 --- a/csrc/kernels/attn_decode_split_kv_mma.cuh +++ b/csrc/kernels/attn_decode_split_kv_mma.cuh @@ -4,34 +4,17 @@ #include "attn_common.h" #include "attn_mma_utils.cuh" -using bf16 = __nv_bfloat16; - // Split-K (FlashDecoding) tensor-core decode via GQA head-packing. +// Decode has q_len == 1, so we pack G = q_head/kv_head query heads into the +// M=16 rows of mma.sync.m16n8k16, turning G independent GEMVs into a single +// GEMM that reuses each loaded K/V tile across all G heads. // -// Decode has q_len == 1, so S = q @ K^T is a GEMV per head — no tensor-core -// work on its own. But GQA gives us G = q_head / kv_head query heads that all -// share one kv_head. We pack those G heads into the M=16 rows of -// mma.sync.m16n8k16, turning G independent GEMVs into a single GEMM that -// reuses each loaded K/V tile across all G heads (K/V load is the decode -// bottleneck, so the reuse is the win, not the flops). The KV sequence is -// partitioned across gridDim.z blocks so that a decode with only -// batch*kv_head independent tasks can fill all SMs. Each (batch, kv_head, -// split) block computes an UN-normalised partial (Oacc, m, l) over its KV -// slice; the combine kernel below reduces across splits. Fixes the "grid too -// small" bottleneck (0.04 waves/SM → many blocks) for long-context, -// small-batch decode. - -template +// IsCausal and HasMask are compile-time bools — no runtime branch in the +// inner compute loop. +// +// Traits = KernelTraits>. +template __global__ void attn_decode_split_kv_mma_kernel(AttentionParams p) { - constexpr int KD = HEAD_DIM / 16; - constexpr int NC8 = BC / 8; - constexpr int KT2 = BC / 16; - constexpr int DN8 = HEAD_DIM / 8; - constexpr int LD = HEAD_DIM; - constexpr int SWIZ_MASK = (HEAD_DIM >= 64) ? 7 : (HEAD_DIM / 8 - 1); - constexpr int VEC = 8; - constexpr int TOTAL = BC * HEAD_DIM; - const int lane = threadIdx.x; const int gid = lane >> 2; const int tid4 = lane & 3; @@ -42,46 +25,44 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams p) { const int G = p.q_head / p.kv_head; const int q_head0 = kv_head * G; - // Double-buffered shared memory for K/V (no sQ needed — Q goes direct - // from global to registers). - __shared__ __align__(16) bf16 sK[STAGES * BC * LD]; - __shared__ __align__(16) bf16 sV[STAGES * BC * LD]; + // Double-buffered shared memory for K/V (no sQ needed) + __shared__ __align__(16) bf16 sK[Traits::STAGES * Traits::BC * Traits::LD]; + __shared__ __align__(16) bf16 sV[Traits::STAGES * Traits::BC * Traits::LD]; - // ---- Load Q directly from global into mma A-operand registers ---- + // Load Q directly from global into mma A-operand registers. + // stride_row = p.q_stride_h for decode (q_len=1). const int q_base = batch * p.q_stride_b + q_head0 * p.q_stride_h; const int qra = gid; const int qrb = gid + 8; const bool va = qra < G, vb = qrb < G; - unsigned Qa[KD][4]; - load_q_mma_frags(p.q + q_base, p.q_stride_h, p.q_stride_d, - qra, qrb, va, vb, tid4, Qa); + unsigned Qa[Traits::KD][4]; + load_q_mma_frags(p.q + q_base, p.q_stride_h, p.q_stride_d, + qra, qrb, va, vb, tid4, Qa); - float Oacc[DN8][4]; -#pragma unroll - for (int j = 0; j < DN8; j++) + float Oacc[Traits::DN8][4]; + #pragma unroll + for (int j = 0; j < Traits::DN8; j++) Oacc[j][0] = Oacc[j][1] = Oacc[j][2] = Oacc[j][3] = 0.0f; float m0 = -FLT_MAX, m1 = -FLT_MAX, l0 = 0.0f, l1 = 0.0f; - // KV: stride-based base — [batch, kv_head, kv_len, head_dim] const int kv_base = batch * p.kv_stride_b + kv_head * p.kv_stride_h; - const int tiles_total = (p.kv_len + BC - 1) / BC; + const int tiles_total = (p.kv_len + Traits::BC - 1) / Traits::BC; const int tiles_per_split = (tiles_total + p.num_splits - 1) / p.num_splits; const int ti_begin = split * tiles_per_split; const int ti_end = min(tiles_total, ti_begin + tiles_per_split); - const int has_mask = p.use_mask && p.mask; - // ---- Load tile lambda: predicated cp.async, unified full/partial ---- + // ---- Load tile lambda: predicated cp.async ---- auto load_tile = [&](int ti, int buf) { - int kv0 = ti * BC; - bf16* dK = sK + buf * BC * LD; - bf16* dV = sV + buf * BC * LD; -#pragma unroll - for (int i = lane * VEC; i < TOTAL; i += 32 * VEC) { - int r = i / HEAD_DIM, d = i % HEAD_DIM; + int kv0 = ti * Traits::BC; + bf16* dK = sK + buf * Traits::BC * Traits::LD; + bf16* dV = sV + buf * Traits::BC * Traits::LD; + #pragma unroll + for (int i = lane * Traits::VEC; i < Traits::TOTAL; + i += Traits::NUM_THREADS * Traits::VEC) { + int r = i / Traits::HEAD_DIM, d = i % Traits::HEAD_DIM; int kc = kv0 + r; bool valid = kc < p.kv_len; - int off = r * LD + swiz_col(d, r, SWIZ_MASK); - // KV stride-based: contiguous within head_dim (stride_d == 1 typically) + int off = r * Traits::LD + swiz_col(d, r, Traits::SWIZ_MASK); int g_off = kv_base + kc * p.kv_stride_l + d * p.kv_stride_d; cp_async_16_pred(&dK[off], &p.k[g_off], valid); cp_async_16_pred(&dV[off], &p.v[g_off], valid); @@ -89,50 +70,48 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams p) { cp_async_commit(); }; - // ---- Prologue: issue first tile load ---- + constexpr int BUF_MASK = (Traits::STAGES > 1) ? (Traits::STAGES - 1) : 0; + + // Prologue if (ti_begin < ti_end) { load_tile(ti_begin, 0); } for (int ti = ti_begin; ti < ti_end; ti++) { - constexpr int BUF_MASK = (STAGES > 1) ? (STAGES - 1) : 0; int buf = (ti - ti_begin) & BUF_MASK; - // Wait for current tile, then issue next tile's prefetch (overlaps - // with this tile's compute). Single syncwarp covers both hazards. - // When STAGES==1, no prefetch — load happens at end of prior iter. cp_async_wait_group<0>(); __syncwarp(); - if constexpr (STAGES > 1) { + if constexpr (Traits::STAGES > 1) { if (ti + 1 < ti_end) load_tile(ti + 1, (ti + 1 - ti_begin) & BUF_MASK); } - const bf16* bK = sK + buf * BC * LD; - const bf16* bV = sV + buf * BC * LD; - int kv0 = ti * BC; + const bf16* bK = sK + buf * Traits::BC * Traits::LD; + const bf16* bV = sV + buf * Traits::BC * Traits::LD; + int kv0 = ti * Traits::BC; - float Sacc[NC8][4]; - mma_compute_scores(Qa, bK, LD, SWIZ_MASK, lane, Sacc); + float Sacc[Traits::NC8][4]; + mma_compute_scores(Qa, bK, lane, Sacc); #pragma unroll - for (int n8 = 0; n8 < NC8; n8++) + for (int n8 = 0; n8 < Traits::NC8; n8++) Sacc[n8][0] *= p.scale, Sacc[n8][1] *= p.scale, Sacc[n8][2] *= p.scale, Sacc[n8][3] *= p.scale; - // Decode: q_len=1, so qrow0=qrow1=0, mask_q_stride irrelevant - int maxc = (p.causal_offset >= 0) ? min(p.kv_len, p.causal_offset + 1) : p.kv_len; - mma_softmax_tile(kv0, maxc, maxc, - 0, 0, - p.mask_b_stride, 0, - batch, - p.mask, has_mask, - Sacc, Oacc, m0, m1, l0, l1, lane); + // Decode: q_len=1, so qrow0=qrow1=0 + int maxc = IsCausal ? min(p.kv_len, p.causal_offset + 1) : p.kv_len; + mma_softmax_tile(kv0, maxc, maxc, + 0, 0, + p.mask_b_stride, 0, + batch, + p.mask, + Sacc, Oacc, m0, m1, l0, l1, lane); - mma_pv_accumulate(Sacc, bV, LD, SWIZ_MASK, lane, Oacc); + mma_pv_accumulate(Sacc, bV, lane, Oacc); __syncwarp(); - if constexpr (STAGES == 1) { + if constexpr (Traits::STAGES == 1) { if (ti + 1 < ti_end) load_tile(ti + 1, 0); } @@ -143,19 +122,19 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams p) { size_t bh = (size_t)batch * p.q_head + h; return bh * p.num_splits + split; }; -#pragma unroll - for (int dn8 = 0; dn8 < DN8; dn8++) { + #pragma unroll + for (int dn8 = 0; dn8 < Traits::DN8; dn8++) { int d = dn8 * 8 + 2 * tid4; int r0 = gid, r1 = gid + 8; if (r0 < G) { int h = q_head0 + r0; - float* op = p.o_part + split_slot(h) * HEAD_DIM; + float* op = p.o_part + split_slot(h) * Traits::HEAD_DIM; op[d] = Oacc[dn8][0]; op[d + 1] = Oacc[dn8][1]; } if (r1 < G) { int h = q_head0 + r1; - float* op = p.o_part + split_slot(h) * HEAD_DIM; + float* op = p.o_part + split_slot(h) * Traits::HEAD_DIM; op[d] = Oacc[dn8][2]; op[d + 1] = Oacc[dn8][3]; } diff --git a/csrc/kernels/attn_mma_utils.cuh b/csrc/kernels/attn_mma_utils.cuh index 4baefa6..c29bb0d 100644 --- a/csrc/kernels/attn_mma_utils.cuh +++ b/csrc/kernels/attn_mma_utils.cuh @@ -3,10 +3,41 @@ #include #include -// Shared MMA utilities for tensor-core GQA kernels. -// mma.sync.m16n8k16 PTX wrappers, ldmatrix helpers, and bf16 packing. +// ============================================================================ +// KernelTraits — FlashAttention-v2 style compile-time configuration bundle. +// +// Bundles all dimension-dependent constants so device functions only need a +// single Traits template parameter rather than scattered . +// ============================================================================ +template +struct KernelTraits { + static constexpr int HEAD_DIM = HEAD_DIM_; + static constexpr int BC = BC_; // K/V tile size along seq dim + static constexpr int WARPS = WARPS_; // warps per block + static constexpr int STAGES = STAGES_; // double-buffer stages (1 or 2) + + static constexpr int BR = 16; // Q rows per warp (mma M=16) + + // Derived: mma.sync.m16n8k16 tile counts + static constexpr int KD = HEAD_DIM / 16; // Q/K k-slides + static constexpr int NC8 = BC / 8; // S n-tiles (N=8) + static constexpr int KT2 = BC / 16; // P k-tiles (K=16) + static constexpr int DN8 = HEAD_DIM / 8; // O n-tiles (N=8) + + static constexpr int LD = HEAD_DIM; // smem leading dim + + // XOR swizzle chunk bits for ldmatrix bank-conflict avoidance. + // mask = log2(LD/8) bits, clamped to stay within LD. + static constexpr int SWIZ_MASK = (HEAD_DIM >= 64) ? 7 : (HEAD_DIM / 8 - 1); + + static constexpr int NUM_THREADS = WARPS * 32; + static constexpr int VEC = 8; // bf16 per cp.async unit (16 bytes) + static constexpr int TOTAL = BC * HEAD_DIM; // total elements per tile +}; + +// ---- PTX wrappers ---- +using bf16 = __nv_bfloat16; -// mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32 __device__ __forceinline__ void mma16816(float* d, const unsigned* a, const unsigned* b, const float* c) { asm volatile( @@ -37,9 +68,7 @@ __device__ __forceinline__ unsigned pkb(bf16 a, bf16 b) { } // ldmatrix: cooperatively load mma fragments from smem (one instruction per -// 16x16 / 16x8 tile) with the exact register layout mma expects — replaces the -// scalar per-thread fragment packing, cutting shared-load instructions and bank -// conflicts. Each lane supplies the shared address of one 8-wide row. +// 16x16 / 16x8 tile) with the exact register layout mma expects. __device__ __forceinline__ void ldmatrix_x4(unsigned* r, const bf16* p) { unsigned a = __cvta_generic_to_shared(p); asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0,%1,%2,%3}, [%4];" @@ -60,29 +89,19 @@ __device__ __forceinline__ void ldmatrix_x2_trans(unsigned* r, const bf16* p) { } // XOR swizzle for shared-memory column at 8-bf16 chunk granularity. -// Eliminates ldmatrix bank conflicts without LD padding: consecutive rows -// land in distinct bank groups. swiz_col(d, r, mask) = ((d>>3)^(r&mask))<<3 | (d&7). -// mask must cover log2(HEAD_DIM/8) chunk bits but stay within LD: use 7 for -// HEAD_DIM>=64 (8+ chunks), 3 for HEAD_DIM=32 (4 chunks). Default 7 keeps -// existing HEAD_DIM>=64 call sites working unchanged. __device__ __forceinline__ int swiz_col(int d, int r, int mask = 7) { return ((d >> 3) ^ (r & mask)) << 3 | (d & 7); } -// cp.async: copy 16 bytes (8 bf16) from global to shared memory directly, -// bypassing registers. Eliminates shared-store bank conflicts and cuts -// load-loop instruction count in half (1 cp.async vs 1 LDG + 1 STS). -// Requires sm_80+. +// cp.async: copy 16 bytes (8 bf16) from global to shared memory directly. __device__ __forceinline__ void cp_async_16(bf16* smem_ptr, const void* gmem_ptr) { unsigned smem_addr = __cvta_generic_to_shared(smem_ptr); asm volatile("cp.async.ca.shared.global [%0], [%1], 16;" :: "r"(smem_addr), "l"(gmem_ptr)); } -// Predicated cp.async: copy 16 bytes when `pred`, otherwise zero-fill the -// destination (src-size operand = 0 → no bytes read from src, so an -// out-of-bounds src address is never dereferenced). Lets full and partial -// tiles share one uniform async load path — no scalar fallback branch. +// Predicated cp.async: copy 16 bytes when `pred`, otherwise zero-fill. +// src_size=0 → no bytes read from src, so out-of-bounds src address is safe. __device__ __forceinline__ void cp_async_16_pred(bf16* smem_ptr, const void* gmem_ptr, bool pred) { @@ -100,9 +119,6 @@ __device__ __forceinline__ void cp_async_wait_all() { asm volatile("cp.async.wait_all;"); } -// Wait until at most N commit groups are still in flight. Used for -// double-buffered pipelining: wait_group<1> lets the next tile's cp.async -// continue while ensuring the current tile's data is ready. template __device__ __forceinline__ void cp_async_wait_group() { asm volatile("cp.async.wait_group %0;" :: "n"(N)); @@ -139,78 +155,65 @@ __device__ inline void load_q_mma_frags( } // --------------------------------------------------------------------------- -// Shared MMA compute functions — used by both decode and prefill MMA kernels. -// Extracted because S=Q@K^T, online softmax, and P@V are structurally identical -// between the two kernels; only the per-row causal/mask bounds differ. -// --------------------------------------------------------------------------- - // S = Q @ K^T (Qa pre-loaded by the caller; scale applied post-mma in the // caller to avoid bf16 precision loss). -// LD and SWIZ_MASK are constexpr in the calling kernel — passing them as -// runtime ints lets the compiler fold them while keeping the signature clean. -template +// Traits provides KD, NC8, LD, and SWIZ_MASK. +// --------------------------------------------------------------------------- +template __device__ inline void mma_compute_scores( - const unsigned Qa[KD][4], + const unsigned Qa[Traits::KD][4], const bf16* __restrict__ sK, - int LD, - int SWIZ_MASK, int lane, - float Sacc[NC8][4]) + float Sacc[Traits::NC8][4]) { #pragma unroll - for (int n8 = 0; n8 < NC8; n8++) { + for (int n8 = 0; n8 < Traits::NC8; n8++) { Sacc[n8][0] = Sacc[n8][1] = Sacc[n8][2] = Sacc[n8][3] = 0.0f; int krow_l = n8 * 8 + (lane & 7); int kcol_h = (lane & 8) ? 8 : 0; #pragma unroll - for (int kt = 0; kt < KD; kt++) { + for (int kt = 0; kt < Traits::KD; kt++) { unsigned b[2]; - ldmatrix_x2(b, &sK[krow_l * LD + swiz_col(kt * 16 + kcol_h, krow_l, SWIZ_MASK)]); + ldmatrix_x2(b, &sK[krow_l * Traits::LD + + swiz_col(kt * 16 + kcol_h, krow_l, Traits::SWIZ_MASK)]); mma16816(Sacc[n8], Qa[kt], b, Sacc[n8]); } } } +// --------------------------------------------------------------------------- // Online softmax + Oacc rescale for one K/V tile. -// maxc0/maxc1: per-row KV column bounds (prefill: per-query-row causal limits; -// decode: same value for both rows since q_len==1). -// qrow0/qrow1: query row indices (for 3D mask indexing; decode passes 0). -// mask_b_stride/mask_q_stride: mask layout (2D: mask_q_stride=0; 3D: =kv_len). -// Reads Sacc (Q@K^T scores), applies causal/mask, computes P = exp(S - nm), -// rescales Oacc by exp(m_old - nm), and updates m/l — all in place. -template +// +// HasMask is a compile-time template bool: when false, the mask branch is +// entirely dead-code-eliminated from the inner unrolled loop. +// --------------------------------------------------------------------------- +template __device__ inline void mma_softmax_tile( int kv0, - int maxc0, - int maxc1, - int qrow0, - int qrow1, - int mask_b_stride, - int mask_q_stride, + int maxc0, int maxc1, + int qrow0, int qrow1, + int mask_b_stride, int mask_q_stride, int mask_batch, const bool* __restrict__ mask, - bool has_mask, - float Sacc[NC8][4], - float Oacc[DN8][4], + float Sacc[Traits::NC8][4], + float Oacc[Traits::DN8][4], float& m0, float& m1, float& l0, float& l1, int lane) { int tid4 = lane & 3; - // Mask out-of-bounds / masked columns: set -FLT_MAX so expf → 0 downstream - // without per-element sentinel checks. Compute tile-local row maxima. float rmax0 = -FLT_MAX, rmax1 = -FLT_MAX; int mask_base0 = mask_batch * mask_b_stride + qrow0 * mask_q_stride; int mask_base1 = mask_batch * mask_b_stride + qrow1 * mask_q_stride; #pragma unroll - for (int n8 = 0; n8 < NC8; n8++) { + for (int n8 = 0; n8 < Traits::NC8; n8++) { int cc = kv0 + n8 * 8 + 2 * tid4; int c1 = cc + 1; - bool b0 = (cc >= maxc0) || (has_mask && !mask[mask_base0 + cc]); - bool b1 = (c1 >= maxc0) || (has_mask && !mask[mask_base0 + c1]); - bool b2 = (cc >= maxc1) || (has_mask && !mask[mask_base1 + cc]); - bool b3 = (c1 >= maxc1) || (has_mask && !mask[mask_base1 + c1]); + bool b0 = (cc >= maxc0) || (HasMask && !mask[mask_base0 + cc]); + bool b1 = (c1 >= maxc0) || (HasMask && !mask[mask_base0 + c1]); + bool b2 = (cc >= maxc1) || (HasMask && !mask[mask_base1 + cc]); + bool b3 = (c1 >= maxc1) || (HasMask && !mask[mask_base1 + c1]); float s0 = b0 ? -FLT_MAX : Sacc[n8][0]; float s1 = b1 ? -FLT_MAX : Sacc[n8][1]; float s2 = b2 ? -FLT_MAX : Sacc[n8][2]; @@ -220,29 +223,20 @@ __device__ inline void mma_softmax_tile( rmax0 = fmaxf(rmax0, fmaxf(s0, s1)); rmax1 = fmaxf(rmax1, fmaxf(s2, s3)); } - // Warp-reduce row maxima across the 4-lane thread group (xor 1, xor 2). rmax0 = fmaxf(rmax0, __shfl_xor_sync(0xFFFFFFFF, rmax0, 1)); rmax0 = fmaxf(rmax0, __shfl_xor_sync(0xFFFFFFFF, rmax0, 2)); rmax1 = fmaxf(rmax1, __shfl_xor_sync(0xFFFFFFFF, rmax1, 1)); rmax1 = fmaxf(rmax1, __shfl_xor_sync(0xFFFFFFFF, rmax1, 2)); - // nm = max(running max m, tile-local max rmax) — updated running maximum. float nm0 = fmaxf(m0, rmax0), nm1 = fmaxf(m1, rmax1); - // corr rescales Oacc and l by exp(m_old - nm). When all-masked (m == nm == - // -FLT_MAX), exp(0) = 1 — correct, no guard needed. float corr0 = __expf(m0 - nm0); float corr1 = __expf(m1 - nm1); - // pn guards only the all-masked-row edge: if nm == -FLT_MAX, exp(S - nm) - // gives 1 not 0 for masked entries. Two scalar masks replace 4*NC8 - // per-element comparisons. float pn0 = (nm0 == -FLT_MAX) ? 0.0f : 1.0f; float pn1 = (nm1 == -FLT_MAX) ? 0.0f : 1.0f; - // P = exp(S - nm) for each element. Masked entries (Sacc = -FLT_MAX) give - // exp(-inf) ≈ 0 naturally; pn zero-fills the all-masked-row edge. float rsum0 = 0.0f, rsum1 = 0.0f; #pragma unroll - for (int n8 = 0; n8 < NC8; n8++) { + for (int n8 = 0; n8 < Traits::NC8; n8++) { float p0 = pn0 * __expf(Sacc[n8][0] - nm0); float p1 = pn0 * __expf(Sacc[n8][1] - nm0); float p2 = pn1 * __expf(Sacc[n8][2] - nm1); @@ -261,22 +255,25 @@ __device__ inline void mma_softmax_tile( m0 = nm0; m1 = nm1; #pragma unroll - for (int j = 0; j < DN8; j++) { + for (int j = 0; j < Traits::DN8; j++) { Oacc[j][0] *= corr0; Oacc[j][1] *= corr0; Oacc[j][2] *= corr1; Oacc[j][3] *= corr1; } } +// --------------------------------------------------------------------------- // O += P @ V (Sacc must contain P = attention weights after softmax). -template +// Traits provides DN8, KT2, LD, and SWIZ_MASK. +// --------------------------------------------------------------------------- +template __device__ inline void mma_pv_accumulate( float Sacc[][4], const bf16* __restrict__ sV, - int LD, int SWIZ_MASK, int lane, - float Oacc[DN8][4]) + int lane, + float Oacc[Traits::DN8][4]) { #pragma unroll - for (int kt2 = 0; kt2 < KT2; kt2++) { + for (int kt2 = 0; kt2 < Traits::KT2; kt2++) { unsigned Pa[4]; Pa[0] = pk2(Sacc[kt2 * 2][0], Sacc[kt2 * 2][1]); Pa[1] = pk2(Sacc[kt2 * 2][2], Sacc[kt2 * 2][3]); @@ -284,9 +281,10 @@ __device__ inline void mma_pv_accumulate( Pa[3] = pk2(Sacc[kt2 * 2 + 1][2], Sacc[kt2 * 2 + 1][3]); int vrow_l = kt2 * 16 + (lane & 15); #pragma unroll - for (int dn8 = 0; dn8 < DN8; dn8++) { + for (int dn8 = 0; dn8 < Traits::DN8; dn8++) { unsigned b[2]; - ldmatrix_x2_trans(b, &sV[vrow_l * LD + swiz_col(dn8 * 8, vrow_l, SWIZ_MASK)]); + ldmatrix_x2_trans(b, &sV[vrow_l * Traits::LD + + swiz_col(dn8 * 8, vrow_l, Traits::SWIZ_MASK)]); mma16816(Oacc[dn8], Pa, b, Oacc[dn8]); } } diff --git a/csrc/kernels/attn_paged_decode.cu b/csrc/kernels/attn_paged_decode.cu index 71ebe62..b68d0e7 100644 --- a/csrc/kernels/attn_paged_decode.cu +++ b/csrc/kernels/attn_paged_decode.cu @@ -5,6 +5,27 @@ #include "attn_entry_utils.cuh" +#ifndef ASTRAI_NO_MMA +template +static void launch_paged_mma_decode_impl(PagedAttentionParams& p) { + using Traits = KernelTraits; + int tiles_total = (p.kv_len + BC - 1) / BC; + p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total); + alloc_split_partials(p); + + paged_attn_decode_split_kv_mma_kernel + <<>>(p); + paged_attn_decode_combine_kernel<<>>(p); +} + +template +static void launch_paged_mma_decode(PagedAttentionParams& p) { + constexpr int STAGES = (HEAD_DIM <= 128) ? 2 : 1; + launch_paged_mma_decode_impl(p); +} +#endif + +template static void launch_paged_scalar_decode(PagedAttentionParams& p) { int group_size = p.q_head / p.kv_head; int chunks_total = (p.kv_len + PDC_CHUNK - 1) / PDC_CHUNK; @@ -14,32 +35,35 @@ static void launch_paged_scalar_decode(PagedAttentionParams& p) { size_t smem = PDC_CHUNK * p.head_dim * sizeof(bf16); dim3 grid = dim3(p.batch * p.kv_head, 1, p.num_splits); dim3 block = dim3(32, group_size); - paged_attn_decode_split_kv_kernel<<>>(p); + paged_attn_decode_split_kv_kernel<<>>(p); paged_attn_decode_combine_kernel<<>>(p); } -#ifndef ASTRAI_NO_MMA -template -static void launch_paged_mma_decode(PagedAttentionParams& p) { - int tiles_total = (p.kv_len + BC - 1) / BC; - p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total); - alloc_split_partials(p); - - paged_attn_decode_split_kv_mma_kernel<<>>(p); - paged_attn_decode_combine_kernel<<>>(p); -} -#endif - template static void dispatch_paged_decode(PagedAttentionParams& p) { + bool is_causal = (p.causal_offset >= 0); + bool has_mask = (p.use_mask && p.mask); + #ifndef ASTRAI_NO_MMA int G = p.q_head / p.kv_head; if (G >= 1 && G <= 16 && p.page_size >= 32) { - launch_paged_mma_decode(p); + if (is_causal) { + if (has_mask) launch_paged_mma_decode(p); + else launch_paged_mma_decode(p); + } else { + if (has_mask) launch_paged_mma_decode(p); + else launch_paged_mma_decode(p); + } return; } #endif - launch_paged_scalar_decode(p); + if (is_causal) { + if (has_mask) launch_paged_scalar_decode(p); + else launch_paged_scalar_decode(p); + } else { + if (has_mask) launch_paged_scalar_decode(p); + else launch_paged_scalar_decode(p); + } } torch::Tensor attn_paged_decode( diff --git a/csrc/kernels/attn_paged_decode_split_kv.cuh b/csrc/kernels/attn_paged_decode_split_kv.cuh index 7a66fc0..14492cc 100644 --- a/csrc/kernels/attn_paged_decode_split_kv.cuh +++ b/csrc/kernels/attn_paged_decode_split_kv.cuh @@ -12,7 +12,7 @@ __device__ inline float paged_warp_reduce_sum(float val) { return val; } -// Split-KV scalar decode: one warp per query head, grid.z partitions KV. +template __global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams p) { int batch = blockIdx.x / p.kv_head; int kv_head = blockIdx.x % p.kv_head; @@ -22,7 +22,6 @@ __global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams p) int lane = threadIdx.x; int hd_per_thread = p.head_dim / 32; - // Q: stride-based [batch, q_head, q_len=1, head_dim] float q_reg[8]; int q_off = batch * p.q_stride_b + q_head * p.q_stride_h + lane * hd_per_thread * p.q_stride_d; @@ -46,7 +45,8 @@ __global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams p) int this_chunk = min(PDC_CHUNK, p.kv_len - chunk_start); int total = this_chunk * p.head_dim; - for (int i = threadIdx.y * 32 + lane; i < total; i += blockDim.x * blockDim.y) { + for (int i = threadIdx.y * 32 + lane; i < total; + i += blockDim.x * blockDim.y) { int s = i / p.head_dim; int d_dim = i % p.head_dim; int pos = chunk_start + s; @@ -69,14 +69,19 @@ __global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams p) float partial = 0.0f; #pragma unroll for (int i = 0; i < hd_per_thread; i++) - partial += q_reg[i] * __bfloat162float(k_smem[s * p.head_dim + lane * hd_per_thread + i]); + partial += q_reg[i] * __bfloat162float( + k_smem[s * p.head_dim + lane * hd_per_thread + i]); partial = paged_warp_reduce_sum(partial) * p.scale; int kv_idx = chunk_start + s; - if (p.use_mask && p.mask && !p.mask[mask_base + kv_idx]) - partial = -FLT_MAX; - if (p.causal_offset >= 0 && kv_idx > p.causal_offset) - partial = -FLT_MAX; + if constexpr (HasMask) { + if (!p.mask[mask_base + kv_idx]) + partial = -FLT_MAX; + } + if constexpr (IsCausal) { + if (kv_idx > p.causal_offset) + partial = -FLT_MAX; + } float new_m = fmaxf(m, partial); float alpha = expf(m - new_m); @@ -93,7 +98,8 @@ __global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams p) + (int64_t)kv_head * p.head_dim; #pragma unroll for (int i = 0; i < hd_per_thread; i++) - acc_reg[i] = acc_reg[i] * alpha + __bfloat162float(p.v_cache[v_base + lane * hd_per_thread + i]) * beta; + acc_reg[i] = acc_reg[i] * alpha + + __bfloat162float(p.v_cache[v_base + lane * hd_per_thread + i]) * beta; } else { #pragma unroll for (int i = 0; i < hd_per_thread; i++) diff --git a/csrc/kernels/attn_paged_decode_split_kv_mma.cuh b/csrc/kernels/attn_paged_decode_split_kv_mma.cuh index 99c0d21..8baa8f7 100644 --- a/csrc/kernels/attn_paged_decode_split_kv_mma.cuh +++ b/csrc/kernels/attn_paged_decode_split_kv_mma.cuh @@ -4,25 +4,14 @@ #include "attn_common.h" #include "attn_mma_utils.cuh" -using bf16 = __nv_bfloat16; - // Paged split-KV tensor-core decode via GQA head-packing. -// Identical algorithm to attn_decode_split_kv_mma_kernel but reads K/V -// directly from the page pool through a page table, eliminating the gather -// copy. Each tile (BC=32) fits within a single page (page_size >= 32), so -// the page-table lookup happens once per tile for cp.async. - -template +// Reads K/V directly from the page pool through a page table — one tile +// (BC=32) fits within a single page (page_size >= 32), so the page-table +// lookup happens once per tile for cp.async. +// +// IsCausal and HasMask are compile-time bools. +template __global__ void paged_attn_decode_split_kv_mma_kernel(PagedAttentionParams p) { - constexpr int KD = HEAD_DIM / 16; - constexpr int NC8 = BC / 8; - constexpr int KT2 = BC / 16; - constexpr int DN8 = HEAD_DIM / 8; - constexpr int LD = HEAD_DIM; - constexpr int SWIZ_MASK = (HEAD_DIM >= 64) ? 7 : (HEAD_DIM / 8 - 1); - constexpr int VEC = 8; - constexpr int TOTAL = BC * HEAD_DIM; - const int lane = threadIdx.x; const int gid = lane >> 2; const int tid4 = lane & 3; @@ -33,123 +22,119 @@ __global__ void paged_attn_decode_split_kv_mma_kernel(PagedAttentionParams const int G = p.q_head / p.kv_head; const int q_head0 = kv_head * G; - __shared__ __align__(16) bf16 sK[STAGES * BC * LD]; - __shared__ __align__(16) bf16 sV[STAGES * BC * LD]; + __shared__ __align__(16) bf16 sK[Traits::STAGES * Traits::BC * Traits::LD]; + __shared__ __align__(16) bf16 sV[Traits::STAGES * Traits::BC * Traits::LD]; - // ---- Load Q directly from global into mma A-operand registers ---- const int q_base = batch * p.q_stride_b + q_head0 * p.q_stride_h; const int qra = gid; const int qrb = gid + 8; const bool va = qra < G, vb = qrb < G; - unsigned Qa[KD][4]; - load_q_mma_frags(p.q + q_base, p.q_stride_h, p.q_stride_d, - qra, qrb, va, vb, tid4, Qa); + unsigned Qa[Traits::KD][4]; + load_q_mma_frags(p.q + q_base, p.q_stride_h, p.q_stride_d, + qra, qrb, va, vb, tid4, Qa); - float Oacc[DN8][4]; -#pragma unroll - for (int j = 0; j < DN8; j++) + float Oacc[Traits::DN8][4]; + #pragma unroll + for (int j = 0; j < Traits::DN8; j++) Oacc[j][0] = Oacc[j][1] = Oacc[j][2] = Oacc[j][3] = 0.0f; float m0 = -FLT_MAX, m1 = -FLT_MAX, l0 = 0.0f, l1 = 0.0f; - const int tiles_total = (p.kv_len + BC - 1) / BC; + const int tiles_total = (p.kv_len + Traits::BC - 1) / Traits::BC; const int tiles_per_split = (tiles_total + p.num_splits - 1) / p.num_splits; const int ti_begin = split * tiles_per_split; const int ti_end = min(tiles_total, ti_begin + tiles_per_split); - const int has_mask = p.use_mask && p.mask; - // Paged strides (constant for the block) - const int64_t page_stride = (int64_t)p.page_size * p.kv_head * HEAD_DIM; - const int64_t pos_stride = (int64_t)p.kv_head * HEAD_DIM; - const int64_t head_off = (int64_t)kv_head * HEAD_DIM; + const int64_t page_stride = (int64_t)p.page_size * p.kv_head * Traits::HEAD_DIM; + const int64_t pos_stride = (int64_t)p.kv_head * Traits::HEAD_DIM; + const int64_t head_off = (int64_t)kv_head * Traits::HEAD_DIM; - // ---- Load tile lambda: predicated cp.async, paged addressing ---- + // ---- Load tile lambda: paged addressing ---- auto load_tile = [&](int ti, int buf) { - int kv0 = ti * BC; - bf16* dK = sK + buf * BC * LD; - bf16* dV = sV + buf * BC * LD; + int kv0 = ti * Traits::BC; + bf16* dK = sK + buf * Traits::BC * Traits::LD; + bf16* dV = sV + buf * Traits::BC * Traits::LD; int logical_page = kv0 / p.page_size; int phys_page = p.page_table[batch * p.max_pages + logical_page]; bool page_valid = (phys_page >= 0); -#pragma unroll - for (int i = lane * VEC; i < TOTAL; i += 32 * VEC) { - int r = i / HEAD_DIM, d = i % HEAD_DIM; + #pragma unroll + for (int i = lane * Traits::VEC; i < Traits::TOTAL; + i += Traits::NUM_THREADS * Traits::VEC) { + int r = i / Traits::HEAD_DIM, d = i % Traits::HEAD_DIM; int kc = kv0 + r; bool valid = (kc < p.kv_len) && page_valid; int page_off = kc % p.page_size; int64_t gmem_base = (int64_t)phys_page * page_stride + (int64_t)page_off * pos_stride + head_off; - int off = r * LD + swiz_col(d, r, SWIZ_MASK); + int off = r * Traits::LD + swiz_col(d, r, Traits::SWIZ_MASK); cp_async_16_pred(&dK[off], &p.k_cache[gmem_base + d], valid); cp_async_16_pred(&dV[off], &p.v_cache[gmem_base + d], valid); } cp_async_commit(); }; - // ---- Prologue: issue first tile load ---- + constexpr int BUF_MASK = (Traits::STAGES > 1) ? (Traits::STAGES - 1) : 0; + if (ti_begin < ti_end) { load_tile(ti_begin, 0); } for (int ti = ti_begin; ti < ti_end; ti++) { - constexpr int BUF_MASK = (STAGES > 1) ? (STAGES - 1) : 0; int buf = (ti - ti_begin) & BUF_MASK; cp_async_wait_group<0>(); __syncwarp(); - if constexpr (STAGES > 1) { + if constexpr (Traits::STAGES > 1) { if (ti + 1 < ti_end) load_tile(ti + 1, (ti + 1 - ti_begin) & BUF_MASK); } - const bf16* bK = sK + buf * BC * LD; - const bf16* bV = sV + buf * BC * LD; - int kv0 = ti * BC; + const bf16* bK = sK + buf * Traits::BC * Traits::LD; + const bf16* bV = sV + buf * Traits::BC * Traits::LD; + int kv0 = ti * Traits::BC; - float Sacc[NC8][4]; - mma_compute_scores(Qa, bK, LD, SWIZ_MASK, lane, Sacc); + float Sacc[Traits::NC8][4]; + mma_compute_scores(Qa, bK, lane, Sacc); -#pragma unroll - for (int n8 = 0; n8 < NC8; n8++) + #pragma unroll + for (int n8 = 0; n8 < Traits::NC8; n8++) Sacc[n8][0] *= p.scale, Sacc[n8][1] *= p.scale, Sacc[n8][2] *= p.scale, Sacc[n8][3] *= p.scale; - // Decode: q_len=1, so qrow0=qrow1=0, mask_q_stride irrelevant - int maxc = (p.causal_offset >= 0) ? min(p.kv_len, p.causal_offset + 1) : p.kv_len; - mma_softmax_tile(kv0, maxc, maxc, - 0, 0, - p.mask_b_stride, 0, - batch, - p.mask, has_mask, - Sacc, Oacc, m0, m1, l0, l1, lane); + int maxc = IsCausal ? min(p.kv_len, p.causal_offset + 1) : p.kv_len; + mma_softmax_tile(kv0, maxc, maxc, + 0, 0, + p.mask_b_stride, 0, + batch, + p.mask, + Sacc, Oacc, m0, m1, l0, l1, lane); - mma_pv_accumulate(Sacc, bV, LD, SWIZ_MASK, lane, Oacc); + mma_pv_accumulate(Sacc, bV, lane, Oacc); __syncwarp(); - if constexpr (STAGES == 1) { + if constexpr (Traits::STAGES == 1) { if (ti + 1 < ti_end) load_tile(ti + 1, 0); } } - // ---- write UN-normalised partials for this split ---- auto split_slot = [&](int h) -> size_t { size_t bh = (size_t)batch * p.q_head + h; return bh * p.num_splits + split; }; -#pragma unroll - for (int dn8 = 0; dn8 < DN8; dn8++) { + #pragma unroll + for (int dn8 = 0; dn8 < Traits::DN8; dn8++) { int d = dn8 * 8 + 2 * tid4; int r0 = gid, r1 = gid + 8; if (r0 < G) { int h = q_head0 + r0; - float* op = p.o_part + split_slot(h) * HEAD_DIM; + float* op = p.o_part + split_slot(h) * Traits::HEAD_DIM; op[d] = Oacc[dn8][0]; op[d + 1] = Oacc[dn8][1]; } if (r1 < G) { int h = q_head0 + r1; - float* op = p.o_part + split_slot(h) * HEAD_DIM; + float* op = p.o_part + split_slot(h) * Traits::HEAD_DIM; op[d] = Oacc[dn8][2]; op[d + 1] = Oacc[dn8][3]; } diff --git a/csrc/kernels/attn_prefill.cu b/csrc/kernels/attn_prefill.cu index d3c979f..c06b73d 100644 --- a/csrc/kernels/attn_prefill.cu +++ b/csrc/kernels/attn_prefill.cu @@ -3,30 +3,48 @@ #ifndef ASTRAI_NO_MMA #include "attn_prefill_split_q_mma.cuh" + +template +static void launch_mma_prefill(AttentionParams& p) { + constexpr int WARPS = 4; + constexpr int BC = (HEAD_DIM <= 128) ? 32 : 16; + using Traits = KernelTraits; + dim3 grid((p.q_len + Traits::BR * WARPS - 1) / (Traits::BR * WARPS), + p.q_head, p.batch); + dim3 block(Traits::NUM_THREADS, 1, 1); + attn_prefill_split_q_mma_kernel<<>>(p); +} #endif -template -static void dispatch_prefill(AttentionParams& p) { -#ifndef ASTRAI_NO_MMA - constexpr int WARPS = 4, BR = 16; - // KV tile: bigger tiles amortize the per-tile cp.async wait + barrier + - // loop overhead over more tensor-core work (this kernel is latency-bound, - // not compute/bandwidth-bound), so BC=32 wins ~6-8% over BC=16 for - // D<=128. D=256 stays at 16: BC=32 double-buffered would need 64KB smem, - // over the 48KB static cap. - constexpr int BC = (HEAD_DIM <= 128) ? 32 : 16; - dim3 grid((p.q_len + BR * WARPS - 1) / (BR * WARPS), p.q_head, p.batch); - dim3 block(WARPS * 32, 1, 1); - // Static shared memory — double-buffered K/V only (no sQ: Q goes direct - // to registers). 2*BC*LD bf16 each for sK and sV → 4*BC*HEAD_DIM*2 bytes. - // Occupancy is smem-capped: D=64→3 blocks/SM (16KB), D=128→1 (32KB), - // D=256→1 (32KB, BC=16). - attn_prefill_split_q_mma_kernel<<>>(p); -#else +template +static void launch_scalar_prefill(AttentionParams& p) { constexpr int G = 8, ROWS = 32, P_BC = 32; dim3 grid((p.q_len + ROWS - 1) / ROWS, p.q_head, p.batch); dim3 block(G, ROWS, 1); - attn_prefill_split_q_kernel_t<<>>(p); + attn_prefill_split_q_kernel_t<<>>(p); +} + +template +static void dispatch_prefill(AttentionParams& p) { + bool is_causal = (p.causal_offset >= 0); + bool has_mask = (p.use_mask && p.mask); + +#ifndef ASTRAI_NO_MMA + if (is_causal) { + if (has_mask) launch_mma_prefill(p); + else launch_mma_prefill(p); + } else { + if (has_mask) launch_mma_prefill(p); + else launch_mma_prefill(p); + } +#else + if (is_causal) { + if (has_mask) launch_scalar_prefill(p); + else launch_scalar_prefill(p); + } else { + if (has_mask) launch_scalar_prefill(p); + else launch_scalar_prefill(p); + } #endif } diff --git a/csrc/kernels/attn_prefill_split_q.cuh b/csrc/kernels/attn_prefill_split_q.cuh index 6faa6a1..8f43304 100644 --- a/csrc/kernels/attn_prefill_split_q.cuh +++ b/csrc/kernels/attn_prefill_split_q.cuh @@ -6,12 +6,9 @@ using bf16 = __nv_bfloat16; // v9: group-split register blocking. G threads cooperate on one query row, -// each owning HEAD_DIM/G dims of qreg[]/acc[]. Small per-thread footprint keeps -// occupancy high; the S dot product is reduced across the G-lane group with a -// short shuffle chain (log2(G) shuffles) instead of a full 32-lane warp reduce. -// Online (per-kv) softmax — cheap because acc[] is only HEAD_DIM/G long. -// Templated on . Block = (G, ROWS). G power-of-two, -// G*ROWS a multiple of 32 with groups warp-aligned. +// each owning HEAD_DIM/G dims of qreg[]/acc[]. IsCausal and HasMask are +// compile-time bools — the compiler eliminates dead branches. +// Templated on . template __device__ __forceinline__ float group_reduce_sum(float v, unsigned mask) { @@ -21,8 +18,7 @@ __device__ __forceinline__ float group_reduce_sum(float v, unsigned mask) { return v; } -// load 8 contiguous bf16 from (16-byte aligned) smem as one float4, unpack to -// 8 floats — cuts shared-load instructions 8x vs scalar bf16 loads. +// load 8 contiguous bf16 from (16-byte aligned) smem as one float4 __device__ __forceinline__ void ld8(const bf16* p, float* o) { float4 raw = *reinterpret_cast(p); const __nv_bfloat162* h = reinterpret_cast(&raw); @@ -34,7 +30,7 @@ __device__ __forceinline__ void ld8(const bf16* p, float* o) { } } -template +template __global__ void attn_prefill_split_q_kernel_t(AttentionParams p) { constexpr int DPT = HEAD_DIM / G; @@ -73,8 +69,6 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams p) { int tt = G * ROWS; int lid = row * G + gpos; - // per-group shuffle mask: only the G lanes of this row's group participate, - // so causal masking (differing loop bounds across rows in a warp) is safe. int lane_in_warp = lid & 31; unsigned gmask = (G == 32) ? 0xFFFFFFFFu : (((1u << G) - 1u) << (lane_in_warp & ~(G - 1))); @@ -95,12 +89,14 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams p) { __syncthreads(); int lim = tlen; - if (p.causal_offset >= 0 && q_row < p.q_len) { - int ep = q_row + p.causal_offset + 1; - if (kv0 >= ep) - lim = 0; - else if (kv0 + tlen > ep) - lim = ep - kv0; + if constexpr (IsCausal) { + if (q_row < p.q_len) { + int ep = q_row + p.causal_offset + 1; + if (kv0 >= ep) + lim = 0; + else if (kv0 + tlen > ep) + lim = ep - kv0; + } } int mask_row_base = mask_batch_base + q_row * p.mask_q_stride; @@ -118,8 +114,10 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams p) { float dot = group_reduce_sum(part, gmask); int kv_idx = kv0 + s; - if (p.use_mask && p.mask && !p.mask[mask_row_base + kv_idx]) - dot = -FLT_MAX; + if constexpr (HasMask) { + if (!p.mask[mask_row_base + kv_idx]) + dot = -FLT_MAX; + } float nm = fmaxf(m, dot); float al = __expf(m - nm); @@ -141,7 +139,6 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams p) { } if (q_row < p.q_len) { - // O: stride-based write int o_off = batch * p.q_stride_b + q_head * p.q_stride_h + q_row * p.q_stride_l + gpos * DPT * p.q_stride_d; float rl = (l > 1e-10f) ? (1.0f / l) : 0.0f; diff --git a/csrc/kernels/attn_prefill_split_q_mma.cuh b/csrc/kernels/attn_prefill_split_q_mma.cuh index edf3890..5bf546b 100644 --- a/csrc/kernels/attn_prefill_split_q_mma.cuh +++ b/csrc/kernels/attn_prefill_split_q_mma.cuh @@ -4,121 +4,76 @@ #include "attn_common.h" #include "attn_mma_utils.cuh" -using bf16 = __nv_bfloat16; - // Tensor-core prefill flash attention (raw mma.sync PTX). // One warp owns BR=16 query rows. S = Q@K^T and O = P@V run on bf16 tensor -// cores via mma.sync.m16n8k16 (f32 accumulate). Q fragments are loaded once -// straight from global into the mma A-operand layout (no smem staging) and -// kept resident in registers across the tile loop. S, O, and the online-softmax -// stats (m, l) also live in registers. -// Shared memory is statically sized via template parameters — no dynamic -// allocation. The mma fragment layout is used directly: the S accumulator -// (f32) maps element-for-element onto the P matrix_a (bf16) operand, so -// softmax needs no shuffle repack; row reductions fold across the 4-lane -// thread group. Templated on with BC a multiple of 16. +// cores via mma.sync.m16n8k16 (f32 accumulate). // -// Software pipeline: K/V are double-buffered and loaded via cp.async one tile -// ahead, so the next tile streams from global memory while the current tile's -// tensor-core math runs — hiding load latency (long_scoreboard). A single -// __syncthreads per tile both publishes the freshly loaded tile cross-warp and -// (because it runs before the next prefetch) guards the buffer being refilled, -// so no second barrier is needed. Predicated cp.async (cp_async_16_pred) -// zero-fills rows past kv_len, unifying full and partial tiles on one path. -// BC=32 (D<=128) amortizes the per-tile wait+barrier+loop overhead over more -// tensor-core work — this kernel is latency-bound (low occupancy from high -// register pressure), so fewer, larger tiles beat many tiny ones. +// IsCausal and HasMask are compile-time bools — the compiler eliminates all +// dead branches in the inner compute loop (FA2-style). // -// Optimizations: load Q fragments directly from global in mma A-operand layout -// (no sQ staging, no prologue barriers); post-multiply scale in float after -// S=Q@K^T to avoid bf16 precision loss; packed bf16x2 output stores; -// causal tile skipping (block-level prefetch bound + warp-level compute skip); -// XOR swizzle (swiz_col) → eliminates ldmatrix bank conflicts without LD -// padding (LD=HEAD_DIM). - -template +// Traits = KernelTraits. +template __global__ void attn_prefill_split_q_mma_kernel(AttentionParams p) { - constexpr int BR = 16; - constexpr int KD = HEAD_DIM / 16; // Q/K k-tiles - constexpr int NC8 = BC / 8; // S n-tiles (N=8 each) - constexpr int KT2 = BC / 16; // P k-tiles (K=16 each) - constexpr int DN8 = HEAD_DIM / 8; // O n-tiles (N=8 each) - constexpr int LD = HEAD_DIM; // XOR swizzle (swiz_col) handles bank conflicts - constexpr int SWIZ_MASK = (HEAD_DIM >= 64) ? 7 : (HEAD_DIM / 8 - 1); // chunk bits, stay within LD - const int warp = threadIdx.x / 32; const int lane = threadIdx.x % 32; - const int gid = lane >> 2; // 0..7 → rows gid, gid+8 + const int gid = lane >> 2; // 0..7 const int tid4 = lane & 3; // 0..3 - const int nthreads = WARPS * 32; const int q_head = blockIdx.y; const int batch = blockIdx.z; const int kv_head = q_head / (p.q_head / p.kv_head); - const int qrow0 = (blockIdx.x * WARPS + warp) * BR; + const int qrow0 = (blockIdx.x * Traits::WARPS + warp) * Traits::BR; - // ---- Static shared memory: double-buffered K/V ---- - // K/V are double-buffered (STAGES=2): the next tile's cp.async load runs - // while the current tile's tensor-core math executes, hiding global-load - // latency (FA2-style software pipeline). No dynamic smem / carveout opt-in. - constexpr int STAGES = 2; - __shared__ __align__(16) bf16 sK[STAGES * BC * LD]; - __shared__ __align__(16) bf16 sV[STAGES * BC * LD]; + // Static shared memory: double-buffered K/V (no sQ — Q goes direct + // to registers in mma A-operand layout). + __shared__ __align__(16) bf16 sK[Traits::STAGES * Traits::BC * Traits::LD]; + __shared__ __align__(16) bf16 sV[Traits::STAGES * Traits::BC * Traits::LD]; // Load Q fragments straight from global into mma A-operand layout. - // stride_row = p.q_stride_l for prefill (multi-q rows across q_len). - // See attn_mma_utils.cuh for the shared template. const int q_base = batch * p.q_stride_b + q_head * p.q_stride_h; const int qra = qrow0 + gid; const int qrb = qrow0 + gid + 8; const bool va = qra < p.q_len, vb = qrb < p.q_len; - unsigned Qa[KD][4]; - load_q_mma_frags(p.q + q_base, p.q_stride_l, p.q_stride_d, - qra, qrb, va, vb, tid4, Qa); + unsigned Qa[Traits::KD][4]; + load_q_mma_frags(p.q + q_base, p.q_stride_l, p.q_stride_d, + qra, qrb, va, vb, tid4, Qa); - float Oacc[DN8][4]; -#pragma unroll - for (int j = 0; j < DN8; j++) + float Oacc[Traits::DN8][4]; + #pragma unroll + for (int j = 0; j < Traits::DN8; j++) Oacc[j][0] = Oacc[j][1] = Oacc[j][2] = Oacc[j][3] = 0.0f; float m0 = -FLT_MAX, m1 = -FLT_MAX, l0 = 0.0f, l1 = 0.0f; // KV: stride-based base const int kv_base = batch * p.kv_stride_b + kv_head * p.kv_stride_h; - const int tiles = (p.kv_len + BC - 1) / BC; - const int qr0 = qrow0 + gid; // row for c0/c1 - const int qr1 = qrow0 + gid + 8; // row for c2/c3 + const int tiles = (p.kv_len + Traits::BC - 1) / Traits::BC; + const int qr0 = qrow0 + gid; + const int qr1 = qrow0 + gid + 8; - // Causal tile-skip bounds (no-op when causal_offset < 0) - const int use_skip = (p.causal_offset >= 0) ? 1 : 0; - const int max_kv = qrow0 + BR - 1 + p.causal_offset; + // Causal tile-skip bounds (dead code when IsCausal == false) + const int max_kv = qrow0 + Traits::BR - 1 + p.causal_offset; const int block_max_kv = - blockIdx.x * WARPS * BR + WARPS * BR - 1 + p.causal_offset; - const int has_mask = p.use_mask && p.mask; + blockIdx.x * Traits::WARPS * Traits::BR + Traits::WARPS * Traits::BR - 1 + + p.causal_offset; - // Last active tile: block-level causal bound (all warps in the block share - // the K/V load, so the prefetch range is the block max, not per-warp). int t_end = tiles - 1; - if (use_skip) { - int bt = block_max_kv / BC; + if constexpr (IsCausal) { + int bt = block_max_kv / Traits::BC; if (bt < t_end) t_end = bt; } - constexpr int VEC = 8; // bf16 per cp.async unit (16 bytes) - constexpr int TOTAL = BC * HEAD_DIM; - // ---- Load tile lambda: predicated cp.async ---- - // Issue cp.async loads for tile `ti` into shared buffer `buf`. Predicated - // loads zero-fill rows past kv_len, so partial tiles need no scalar path. auto load_tile = [&](int ti, int buf) { - int kv0 = ti * BC; - bf16* dK = sK + buf * BC * LD; - bf16* dV = sV + buf * BC * LD; -#pragma unroll - for (int i = threadIdx.x * VEC; i < TOTAL; i += nthreads * VEC) { - int r = i / HEAD_DIM, d = i % HEAD_DIM; + int kv0 = ti * Traits::BC; + bf16* dK = sK + buf * Traits::BC * Traits::LD; + bf16* dV = sV + buf * Traits::BC * Traits::LD; + #pragma unroll + for (int i = threadIdx.x * Traits::VEC; i < Traits::TOTAL; + i += Traits::NUM_THREADS * Traits::VEC) { + int r = i / Traits::HEAD_DIM, d = i % Traits::HEAD_DIM; int kc = kv0 + r; bool valid = kc < p.kv_len; - int off = r * LD + swiz_col(d, r, SWIZ_MASK); + int off = r * Traits::LD + swiz_col(d, r, Traits::SWIZ_MASK); int g_off = kv_base + kc * p.kv_stride_l + d * p.kv_stride_d; cp_async_16_pred(&dK[off], &p.k[g_off], valid); cp_async_16_pred(&dV[off], &p.v[g_off], valid); @@ -132,65 +87,60 @@ __global__ void attn_prefill_split_q_mma_kernel(AttentionParams p) { for (int ti = 0; ti <= t_end; ti++) { int buf = ti & 1; - // Wait for the current tile's async copies, then a single barrier: it - // both publishes this tile's data cross-warp AND guarantees the prior - // compute on the buffer we are about to refill has finished. Issuing - // the next tile's load *after* this barrier lets one barrier cover both - // hazards (vs two), while the load still overlaps this tile's math. + // Wait for current tile, then publish cross-warp + guard buffer reuse. cp_async_wait_group<0>(); __syncthreads(); if (ti < t_end) load_tile(ti + 1, (ti + 1) & 1); - const bf16* bK = sK + buf * BC * LD; - const bf16* bV = sV + buf * BC * LD; - int kv0 = ti * BC; + const bf16* bK = sK + buf * Traits::BC * Traits::LD; + const bf16* bV = sV + buf * Traits::BC * Traits::LD; + int kv0 = ti * Traits::BC; - // Warp-level causal skip - if (!use_skip || kv0 <= max_kv) { + // Warp-level causal skip (dead branch eliminated when IsCausal == false) + if (!IsCausal || kv0 <= max_kv) { - // S = Q @ K^T + scale + online softmax + O += P @ V - float Sacc[NC8][4]; - mma_compute_scores(Qa, bK, LD, SWIZ_MASK, lane, Sacc); + float Sacc[Traits::NC8][4]; + mma_compute_scores(Qa, bK, lane, Sacc); - // post-multiply scale in float (no bf16 precision loss from pre-scaling Q) - #pragma unroll - for (int n8 = 0; n8 < NC8; n8++) - Sacc[n8][0] *= p.scale, Sacc[n8][1] *= p.scale, - Sacc[n8][2] *= p.scale, Sacc[n8][3] *= p.scale; + // Post-multiply scale in float (no bf16 precision loss) + #pragma unroll + for (int n8 = 0; n8 < Traits::NC8; n8++) + Sacc[n8][0] *= p.scale, Sacc[n8][1] *= p.scale, + Sacc[n8][2] *= p.scale, Sacc[n8][3] *= p.scale; - int maxc0 = (p.causal_offset >= 0) ? min(p.kv_len, qr0 + p.causal_offset + 1) - : p.kv_len; - int maxc1 = (p.causal_offset >= 0) ? min(p.kv_len, qr1 + p.causal_offset + 1) - : p.kv_len; - mma_softmax_tile(kv0, maxc0, maxc1, - qr0, qr1, - p.mask_b_stride, p.mask_q_stride, - batch, - p.mask, has_mask, - Sacc, Oacc, m0, m1, l0, l1, lane); + int maxc0 = IsCausal ? min(p.kv_len, qr0 + p.causal_offset + 1) + : p.kv_len; + int maxc1 = IsCausal ? min(p.kv_len, qr1 + p.causal_offset + 1) + : p.kv_len; + mma_softmax_tile(kv0, maxc0, maxc1, + qr0, qr1, + p.mask_b_stride, p.mask_q_stride, + batch, + p.mask, + Sacc, Oacc, m0, m1, l0, l1, lane); - mma_pv_accumulate(Sacc, bV, LD, SWIZ_MASK, lane, Oacc); - } // if active (warp-level causal skip) + mma_pv_accumulate(Sacc, bV, lane, Oacc); + } } - // ---- write output ---- (packed bf16x2 stores: one 32-bit STG per pair, - // halves store count and removes the uncoalesced scalar-store penalty) + // ---- write output: packed bf16x2 stores ---- float rl0 = (l0 > 1e-20f) ? (1.0f / l0) : 0.0f; float rl1 = (l1 > 1e-20f) ? (1.0f / l1) : 0.0f; - // O: stride-based write const int o_base = batch * p.q_stride_b + q_head * p.q_stride_h; -#pragma unroll - for (int dn8 = 0; dn8 < DN8; dn8++) { + #pragma unroll + for (int dn8 = 0; dn8 < Traits::DN8; dn8++) { int d = dn8 * 8 + 2 * tid4; if (qr0 < p.q_len) { __nv_bfloat162 v = __floats2bfloat162_rn(Oacc[dn8][0] * rl0, - Oacc[dn8][1] * rl0); - *reinterpret_cast<__nv_bfloat162*>(&p.o[o_base + qr0 * p.q_stride_l + d * p.q_stride_d]) = v; + Oacc[dn8][1] * rl0); + *reinterpret_cast<__nv_bfloat162*>( + &p.o[o_base + qr0 * p.q_stride_l + d * p.q_stride_d]) = v; } if (qr1 < p.q_len) { __nv_bfloat162 v = __floats2bfloat162_rn(Oacc[dn8][2] * rl1, - Oacc[dn8][3] * rl1); - *reinterpret_cast<__nv_bfloat162*>(&p.o[o_base + qr1 * p.q_stride_l + d * p.q_stride_d]) = v; + Oacc[dn8][3] * rl1); + *reinterpret_cast<__nv_bfloat162*>( + &p.o[o_base + qr1 * p.q_stride_l + d * p.q_stride_d]) = v; } } } diff --git a/csrc/tests/attn_decode_test.cu b/csrc/tests/attn_decode_test.cu index 221f505..cda0bc3 100644 --- a/csrc/tests/attn_decode_test.cu +++ b/csrc/tests/attn_decode_test.cu @@ -1,5 +1,5 @@ /* -Pure-C test: +Pure-C test — updated for KernelTraits + IsCausal/HasMask. nvcc -I csrc -arch=sm_89 -O3 \ --use_fast_math --ptxas-options=-O3 --extra-device-vectorization \ csrc/tests/attn_decode_test.cu -o test && ./test @@ -11,35 +11,29 @@ nvcc -I csrc -arch=sm_89 -O3 \ #include "../kernels/attn_decode_split_kv_mma.cuh" #endif -// Split-K scratch (torch-free): the production launcher allocates these from -// torch; here we pass pre-allocated device buffers so the bench loop doesn't -// pay a cudaMalloc per iteration. Size for the maximum split count (32). +// Split-K scratch (torch-free) struct DecodeScratch { float* o_part = nullptr; float* ml_part = nullptr; }; -// Launch the production decode path (tensor-core head-packing MMA on sm_80+, -// scalar fallback otherwise), mirroring dispatch_decode() in attn_decode.cu. #ifndef ASTRAI_NO_MMA -static bool decode_use_mma(const AttentionParams& p) { - int G = p.q_head / p.kv_head; - return !p.use_mask && G > 1 && G <= 16; -} - -template +template static void launch_mma_decode(AttentionParams& p, DecodeScratch& sc) { + constexpr int STAGES = (HEAD_DIM <= 128) ? 2 : 1; + using Traits = KernelTraits; int tiles_total = (p.kv_len + BC - 1) / BC; p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total); p.o_part = sc.o_part; p.ml_part = sc.ml_part; - attn_decode_split_kv_mma_kernel + attn_decode_split_kv_mma_kernel <<>>(p); attn_decode_combine_kernel<<>>(p); } #endif +template static void launch_scalar_decode(AttentionParams& p, DecodeScratch& sc) { int gs = p.q_head / p.kv_head; int chunks_total = (p.kv_len + DC_CHUNK - 1) / DC_CHUNK; @@ -48,16 +42,36 @@ static void launch_scalar_decode(AttentionParams& p, DecodeScratch& sc) { p.ml_part = sc.ml_part; size_t smem = DC_CHUNK * p.head_dim * sizeof(bf16); - attn_decode_split_kv_kernel<<>>(p); + attn_decode_split_kv_kernel + <<>>(p); attn_decode_combine_kernel<<>>(p); } template static void dispatch_decode_t(AttentionParams& p, DecodeScratch& sc) { + bool is_causal = (p.causal_offset >= 0); + bool has_mask = (p.use_mask && p.mask); + #ifndef ASTRAI_NO_MMA - if (decode_use_mma(p)) { launch_mma_decode(p, sc); return; } + int G = p.q_head / p.kv_head; + if (G >= 1 && G <= 16) { + if (is_causal) { + if (has_mask) launch_mma_decode(p, sc); + else launch_mma_decode(p, sc); + } else { + if (has_mask) launch_mma_decode(p, sc); + else launch_mma_decode(p, sc); + } + return; + } #endif - launch_scalar_decode(p, sc); + if (is_causal) { + if (has_mask) launch_scalar_decode(p, sc); + else launch_scalar_decode(p, sc); + } else { + if (has_mask) launch_scalar_decode(p, sc); + else launch_scalar_decode(p, sc); + } } static void dispatch_decode(AttentionParams& p, DecodeScratch& sc) { @@ -123,77 +137,92 @@ static void bench() { } } +static int run_test(int B, int Hq, int Hk, int sl, int D, int causal) { + int gs = Hq / Hk; + printf("=== B=%d Hq=%d Hk=%d seq=%d D=%d gs=%d causal=%d ===\n", + B,Hq,Hk,sl,D,gs,causal); + + size_t nQ = B*Hq*1*D, nKV = B*Hk*sl*D; + float *hQ=new float[nQ], *hK=new float[nKV], *hV=new float[nKV]; + for (size_t i=0;i p; + p.batch=B; p.q_head=Hq; p.kv_head=Hk; p.q_len=1; p.kv_len=sl; p.head_dim=D; + p.use_mask=0; p.causal_offset=causal?0:-1; + p.scale=1.0f/sqrtf((float)D); + set_default_strides(p); + p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO; + + DecodeScratch sc; + cudaMalloc(&sc.o_part, (size_t)B*Hq*32*D*sizeof(float)); + cudaMalloc(&sc.ml_part, (size_t)B*Hq*32*2*sizeof(float)); + + double t0=now_ms(); + dispatch_decode(p, sc); + cudaDeviceSynchronize(); + double kms=now_ms()-t0; + cudaError_t err=cudaGetLastError(); + if (err!=cudaSuccess){printf("CUDA err: %s\n",cudaGetErrorString(err));return 1;} + + bf16* hOut=new bf16[nQ]; + cudaMemcpy(hOut,dO,nQ*2,cudaMemcpyDeviceToHost); + + float* ref=new float[nQ]; + cpu_attention_ref(hQ, hK, hV, hMask, ref, B, Hq, Hk, 1, sl, D, causal ? 0 : -1); + + float max_err=0; + for (size_t i=0;imax_err) max_err=d; + } + printf("kernel: %.3f ms max_err: %.6e\n\n",kms,max_err); + + cudaFree(dQ);cudaFree(dK);cudaFree(dV);cudaFree(dO);cudaFree(dMask); + cudaFree(sc.o_part);cudaFree(sc.ml_part); + delete[]hQ;delete[]hK;delete[]hV;delete[]hMask;delete[]hOut;delete[]ref;delete[]tmp; + + return (max_err < 0.05f) ? 0 : 1; +} + int main() { - const int configs[][5] = { - {1, 2, 1, 64, 32}, // B,Hq,Hk,seq_len,D - {1, 32, 4, 512, 128}, - {1, 32, 4, 1024, 128}, + const int configs[][6] = { + {1, 2, 1, 64, 32, 0}, // B,Hq,Hk,seq_len,D,causal + {1, 32, 4, 512, 128, 0}, + {1, 32, 4, 1024, 128, 0}, + {1, 32, 4, 512, 128, 1}, // causal decode }; int n_cfgs = sizeof(configs) / sizeof(configs[0]); + int fail = 0; for (int ci = 0; ci < n_cfgs; ci++) { int B = configs[ci][0], Hq = configs[ci][1], Hk = configs[ci][2]; - int sl = configs[ci][3], D = configs[ci][4], gs = Hq / Hk; - printf("=== B=%d Hq=%d Hk=%d seq=%d D=%d gs=%d ===\n", B,Hq,Hk,sl,D,gs); + int sl = configs[ci][3], D = configs[ci][4], causal = configs[ci][5]; + fail += run_test(B, Hq, Hk, sl, D, causal); + if (fail) break; + } - size_t nQ = B*Hq*1*D, nKV = B*Hk*sl*D; - float *hQ=new float[nQ], *hK=new float[nKV], *hV=new float[nKV]; - for (size_t i=0;i p; - p.batch=B; p.q_head=Hq; p.kv_head=Hk; p.q_len=1; p.kv_len=sl; p.head_dim=D; - p.use_mask=0; p.causal_offset=-1; - p.scale=1.0f/sqrtf((float)D); - set_default_strides(p); - p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO; - - // Split-K scratch (max 32 splits), sized for the production MMA path. - DecodeScratch sc; - cudaMalloc(&sc.o_part, (size_t)B*Hq*32*D*sizeof(float)); - cudaMalloc(&sc.ml_part, (size_t)B*Hq*32*2*sizeof(float)); - - double t0=now_ms(); - dispatch_decode(p, sc); - cudaDeviceSynchronize(); - double kms=now_ms()-t0; - cudaError_t err=cudaGetLastError(); - if (err!=cudaSuccess){printf("CUDA err: %s\n",cudaGetErrorString(err));return 1;} - - bf16* hOut=new bf16[nQ]; - cudaMemcpy(hOut,dO,nQ*2,cudaMemcpyDeviceToHost); - - float* ref=new float[nQ]; - cpu_attention_ref(hQ, hK, hV, hMask, ref, B, Hq, Hk, 1, sl, D, -1); - - float max_err=0; - for (size_t i=0;imax_err) max_err=d; - } - printf("kernel: %.3f ms max_err: %.6e\n\n",kms,max_err); - - cudaFree(dQ);cudaFree(dK);cudaFree(dV);cudaFree(dO);cudaFree(dMask); - cudaFree(sc.o_part);cudaFree(sc.ml_part); - delete[]hQ;delete[]hK;delete[]hV;delete[]hMask;delete[]hOut;delete[]ref;delete[]tmp; + if (fail) { + printf("FAILED\n"); + return fail; } printf("All tests passed!\n"); bench(); diff --git a/csrc/tests/attn_paged_decode_test.cu b/csrc/tests/attn_paged_decode_test.cu index fd1ade9..115bf67 100644 --- a/csrc/tests/attn_paged_decode_test.cu +++ b/csrc/tests/attn_paged_decode_test.cu @@ -36,34 +36,63 @@ static void gather_kv_cpu( } } +#ifndef ASTRAI_NO_MMA +template +static void launch_paged_mma_decode(PagedAttentionParams& p) { + constexpr int STAGES = (HEAD_DIM <= 128) ? 2 : 1; + using Traits = KernelTraits; + int tiles_total = (p.kv_len + 32 - 1) / 32; + p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total); + paged_attn_decode_split_kv_mma_kernel + <<>>(p); +} +#endif + +template +static void launch_paged_scalar_decode(PagedAttentionParams& p) { + int group_sz = p.q_head / p.kv_head; + int chunks_total = (p.kv_len + PDC_CHUNK - 1) / PDC_CHUNK; + p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total); + size_t smem = PDC_CHUNK * p.head_dim * sizeof(bf16); + paged_attn_decode_split_kv_kernel<<< + dim3(p.batch * p.kv_head, 1, p.num_splits), + dim3(32, group_sz), smem>>>(p); +} + template static void launch_paged_decode(PagedAttentionParams& p) { + bool is_causal = (p.causal_offset >= 0); + bool has_mask = (p.use_mask && p.mask); + #ifndef ASTRAI_NO_MMA int G_check = p.q_head / p.kv_head; - bool use_mma = !p.use_mask && G_check >= 1 && G_check <= 16 && p.page_size >= 32; + bool use_mma = G_check >= 1 && G_check <= 16 && p.page_size >= 32; if (use_mma) { - constexpr int STAGES = (HEAD_DIM <= 128) ? 2 : 1; - int tiles_total = (p.kv_len + 32 - 1) / 32; - p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total); - paged_attn_decode_split_kv_mma_kernel - <<>>(p); + if (is_causal) { + if (has_mask) launch_paged_mma_decode(p); + else launch_paged_mma_decode(p); + } else { + if (has_mask) launch_paged_mma_decode(p); + else launch_paged_mma_decode(p); + } } else #endif { - int group_sz = p.q_head / p.kv_head; - int chunks_total = (p.kv_len + PDC_CHUNK - 1) / PDC_CHUNK; - p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total); - size_t smem = PDC_CHUNK * p.head_dim * sizeof(bf16); - paged_attn_decode_split_kv_kernel<<< - dim3(p.batch * p.kv_head, 1, p.num_splits), - dim3(32, group_sz), smem>>>(p); + if (is_causal) { + if (has_mask) launch_paged_scalar_decode(p); + else launch_paged_scalar_decode(p); + } else { + if (has_mask) launch_paged_scalar_decode(p); + else launch_paged_scalar_decode(p); + } } paged_attn_decode_combine_kernel<<>>(p); } template -static int run_test(int B, int Hq, int Hkv, int kv_len, int page_size, int seed) { - printf("B=%d Hq=%d Hkv=%d kv_len=%d page_sz=%d head_dim=%d ... ", B, Hq, Hkv, kv_len, page_size, HEAD_DIM); +static int run_test(int B, int Hq, int Hkv, int kv_len, int page_size, int causal, int seed) { + printf("B=%d Hq=%d Hkv=%d kv_len=%d page_sz=%d head_dim=%d causal=%d ... ", + B, Hq, Hkv, kv_len, page_size, HEAD_DIM, causal); fflush(stdout); int max_pages = (kv_len + page_size - 1) / page_size; @@ -138,13 +167,14 @@ static int run_test(int B, int Hq, int Hkv, int kv_len, int page_size, int seed) } float* h_o_ref = (float*)calloc(B * Hq * HEAD_DIM, sizeof(float)); - cpu_attention_ref(h_q_f, h_k_f, h_v_f, nullptr, h_o_ref, B, Hq, Hkv, 1, kv_len, HEAD_DIM, -1); + cpu_attention_ref(h_q_f, h_k_f, h_v_f, nullptr, h_o_ref, B, Hq, Hkv, + 1, kv_len, HEAD_DIM, causal ? 0 : -1); float scale_val = 1.0f / sqrtf((float)HEAD_DIM); PagedAttentionParams p; p.batch = B; p.q_head = Hq; p.kv_head = Hkv; p.q_len = 1; p.kv_len = kv_len; p.head_dim = HEAD_DIM; - p.use_mask = 0; p.causal_offset = -1; + p.use_mask = 0; p.causal_offset = causal ? 0 : -1; set_default_paged_strides(p); p.num_splits = 1; p.scale = scale_val; p.page_size = page_size; p.max_pages = max_pages; @@ -201,25 +231,24 @@ static int run_test(int B, int Hq, int Hkv, int kv_len, int page_size, int seed) struct TestCase { int head_dim; - int B, Hq, Hkv, kv_len, page_size, seed; + int B, Hq, Hkv, kv_len, page_size, causal, seed; }; static const TestCase TESTS[] = { - {128, 1, 1, 1, 8, 128, 1}, - {128, 1, 4, 4, 128, 128, 2}, - {128, 2, 4, 4, 256, 128, 3}, - {128, 1, 4, 1, 64, 64, 4}, - {128, 1, 8, 2, 64, 128, 5}, - {128, 2, 16, 4, 128, 128, 6}, - {64, 1, 4, 2, 32, 128, 7}, - {256, 1, 2, 1, 16, 128, 8}, - {32, 1, 4, 2, 32, 64, 9}, - {128, 3, 8, 2, 256, 128, 10}, - {128, 2, 32, 8, 512, 128, 11}, -#ifndef ASTRAI_NO_MMA - {128, 1, 16, 2, 256, 128, 12}, - {128, 2, 32, 4, 512, 128, 13}, -#endif + {128, 1, 1, 1, 8, 128, 0, 1}, + {128, 1, 4, 4, 128, 128, 0, 2}, + {128, 2, 4, 4, 256, 128, 0, 3}, + {128, 1, 4, 1, 64, 64, 0, 4}, + {128, 1, 8, 2, 64, 128, 0, 5}, + {128, 2, 16, 4, 128, 128, 0, 6}, + {64, 1, 4, 2, 32, 128, 0, 7}, + {256, 1, 2, 1, 16, 128, 0, 8}, + {32, 1, 4, 2, 32, 64, 0, 9}, + {128, 3, 8, 2, 256, 128, 0, 10}, + {128, 2, 32, 8, 512, 128, 0, 11}, + {128, 1, 16, 2, 256, 128, 0, 12}, + {128, 2, 32, 4, 512, 128, 0, 13}, + {128, 2, 8, 2, 128, 128, 1, 14}, // causal paged decode }; static int dispatch_test(const TestCase& tc) { @@ -227,13 +256,11 @@ static int dispatch_test(const TestCase& tc) { int r = 0; dispatch_by_head_dim(tc.head_dim, [&]() { matched = true; - r = run_test(tc.B, tc.Hq, tc.Hkv, tc.kv_len, tc.page_size, tc.seed); + r = run_test(tc.B, tc.Hq, tc.Hkv, tc.kv_len, tc.page_size, tc.causal, tc.seed); }); return matched ? r : 1; } -// Warmed-up, CUDA-event timed sweep over paged decode configs. -// Bytes = K + V read through page table (B*Hk*kv*D each), bf16. template static void bench_config(int B, int Hq, int Hkv, int kv_len, int page_size) { int max_pages = (kv_len + page_size - 1) / page_size; diff --git a/csrc/tests/attn_prefill_test.cu b/csrc/tests/attn_prefill_test.cu index 30bfe40..8fd3dab 100644 --- a/csrc/tests/attn_prefill_test.cu +++ b/csrc/tests/attn_prefill_test.cu @@ -1,5 +1,5 @@ /* -Pure-C test: +Pure-C test — updated for KernelTraits + IsCausal/HasMask. nvcc -I csrc -arch=sm_89 -O3 \ --use_fast_math --ptxas-options=-O3 --extra-device-vectorization \ csrc/tests/attn_prefill_test.cu -o test && ./test @@ -11,35 +11,60 @@ nvcc -I csrc -arch=sm_89 -O3 \ #include "../kernels/attn_prefill_split_q_mma.cuh" #endif -// Launch the production prefill path (tensor-core MMA on sm_80+, else the -// scalar fallback), mirroring dispatch_prefill() in attn_prefill.cu. -template -static void launch_prefill(AttentionParams& p) { #ifndef ASTRAI_NO_MMA - constexpr int WARPS = 4, BR = 16; +template +static void launch_mma_prefill(AttentionParams& p) { + constexpr int WARPS = 4; constexpr int BC = (HEAD_DIM <= 128) ? 32 : 16; - dim3 grid((p.q_len + BR * WARPS - 1) / (BR * WARPS), p.q_head, p.batch); - dim3 block(WARPS * 32, 1, 1); - attn_prefill_split_q_mma_kernel<<>>(p); -#else + using Traits = KernelTraits; + dim3 grid((p.q_len + Traits::BR * WARPS - 1) / (Traits::BR * WARPS), + p.q_head, p.batch); + dim3 block(Traits::NUM_THREADS, 1, 1); + attn_prefill_split_q_mma_kernel<<>>(p); +} +#endif + +template +static void launch_scalar_prefill(AttentionParams& p) { constexpr int G = 8, ROWS = 32, P_BC = 32; dim3 grid((p.q_len + ROWS - 1) / ROWS, p.q_head, p.batch); dim3 block(G, ROWS, 1); - attn_prefill_split_q_kernel_t<<>>(p); + attn_prefill_split_q_kernel_t<<>>(p); +} + +template +static void launch_prefill_dispatch(AttentionParams& p) { + bool is_causal = (p.causal_offset >= 0); + bool has_mask = (p.use_mask && p.mask); +#ifndef ASTRAI_NO_MMA + if (is_causal) { + if (has_mask) launch_mma_prefill(p); + else launch_mma_prefill(p); + } else { + if (has_mask) launch_mma_prefill(p); + else launch_mma_prefill(p); + } +#else + if (is_causal) { + if (has_mask) launch_scalar_prefill(p); + else launch_scalar_prefill(p); + } else { + if (has_mask) launch_scalar_prefill(p); + else launch_scalar_prefill(p); + } #endif } static void dispatch_prefill(AttentionParams& p) { switch (p.head_dim) { - case 64: launch_prefill<64>(p); break; - case 128: launch_prefill<128>(p); break; + case 64: launch_prefill_dispatch<64>(p); break; + case 128: launch_prefill_dispatch<128>(p); break; default: printf("bench: unsupported D=%d\n", p.head_dim); } } // Warmed-up, CUDA-event timed throughput sweep over the production MMA path. -// Reports per-call latency and effective tensor-core TFLOP/s (2 matmuls: -// QK^T and P@V, each 2*B*Hq*ql*kl*D flops; halved for causal). static void bench() { const int cfgs[][7] = { {1,32,4,512,512,128,0}, @@ -94,7 +119,6 @@ static void bench() { double flops = 4.0*B*Hq*(double)ql*kl*D; if (causal) flops *= 0.5; double tflops = flops/(ms*1e-3)/1e12; - // HBM traffic: Q + O (B*Hq*ql*D each) + K + V (B*Hk*kl*D each), bf16. double bytes = 2.0 * (2.0*nQ + 2.0*nKV); double gbps = bytes/(ms*1e-3)/1e9; @@ -110,6 +134,59 @@ static void bench() { } } +static int run_test(int B, int Hq, int Hk, int ql, int kl, int D, int causal) { + printf("=== B=%d Hq=%d Hk=%d q=%d kv=%d D=%d causal=%d ===\n", + B,Hq,Hk,ql,kl,D,causal); + + size_t nQ = B*Hq*ql*D, nKV = B*Hk*kl*D; + float *hQ=new float[nQ], *hK=new float[nKV], *hV=new float[nKV]; + for (size_t i=0;i p; + p.batch=B; p.q_head=Hq; p.kv_head=Hk; p.q_len=ql; p.kv_len=kl; p.head_dim=D; + p.use_mask=0; p.causal_offset=causal?0:-1; + set_default_strides(p); + p.scale=1.0f/sqrtf((float)D); + p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO; + + double t0=now_ms(); + dispatch_prefill(p); + cudaDeviceSynchronize(); + double kms=now_ms()-t0; + cudaError_t err=cudaGetLastError(); + if (err!=cudaSuccess){printf("CUDA err: %s\n",cudaGetErrorString(err));return 1;} + + bf16* hOut=new bf16[nQ]; + cudaMemcpy(hOut,dO,nQ*2,cudaMemcpyDeviceToHost); + + float* ref=new float[nQ]; + cpu_attention_ref(hQ, hK, hV, nullptr, ref, B, Hq, Hk, ql, kl, D, causal ? 0 : -1); + + float max_err=0; + for (size_t i=0;imax_err) max_err=d; + } + printf("kernel: %.3f ms max_err: %.6e\n\n",kms,max_err); + + cudaFree(dQ);cudaFree(dK);cudaFree(dV);cudaFree(dO); + delete[]hQ;delete[]hK;delete[]hV;delete[]hOut;delete[]ref;delete[]tmp; + + return (max_err < 0.05f) ? 0 : 1; +} + int main() { const int configs[][7] = { {1,2,1,64,128,64,0}, // tiny: B,Hq,Hk,q,kv,D,causal @@ -118,59 +195,19 @@ int main() { {1,4,2,256,256,128,1}, // causal }; int n_configs = sizeof(configs) / sizeof(configs[0]); + int fail = 0; for (int ci = 0; ci < n_configs; ci++) { int B=configs[ci][0], Hq=configs[ci][1], Hk=configs[ci][2]; int ql=configs[ci][3], kl=configs[ci][4], D=configs[ci][5]; int causal=configs[ci][6]; - printf("=== B=%d Hq=%d Hk=%d q=%d kv=%d D=%d causal=%d ===\n", - B,Hq,Hk,ql,kl,D,causal); + fail += run_test(B, Hq, Hk, ql, kl, D, causal); + if (fail) break; + } - size_t nQ = B*Hq*ql*D, nKV = B*Hk*kl*D; - float *hQ=new float[nQ], *hK=new float[nKV], *hV=new float[nKV]; - for (size_t i=0;i p; - p.batch=B; p.q_head=Hq; p.kv_head=Hk; p.q_len=ql; p.kv_len=kl; p.head_dim=D; - p.use_mask=0; p.causal_offset=causal?0:-1; - set_default_strides(p); - p.scale=1.0f/sqrtf((float)D); - p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO; - - double t0=now_ms(); - dispatch_prefill(p); - cudaDeviceSynchronize(); - double kms=now_ms()-t0; - cudaError_t err=cudaGetLastError(); - if (err!=cudaSuccess){printf("CUDA err: %s\n",cudaGetErrorString(err));return 1;} - - bf16* hOut=new bf16[nQ]; - cudaMemcpy(hOut,dO,nQ*2,cudaMemcpyDeviceToHost); - - float* ref=new float[nQ]; - cpu_attention_ref(hQ, hK, hV, nullptr, ref, B, Hq, Hk, ql, kl, D, causal ? 0 : -1); - - float max_err=0; - for (size_t i=0;imax_err) max_err=d; - } - printf("kernel: %.3f ms max_err: %.6e\n\n",kms,max_err); - - cudaFree(dQ);cudaFree(dK);cudaFree(dV);cudaFree(dO); - delete[]hQ;delete[]hK;delete[]hV;delete[]hOut;delete[]ref;delete[]tmp; + if (fail) { + printf("FAILED\n"); + return fail; } printf("All tests passed!\n"); bench();