From 2667b8116d3c956169324032138fe6f529fa0164 Mon Sep 17 00:00:00 2001 From: ViperEkura <3081035982@qq.com> Date: Wed, 5 Aug 2026 14:06:04 +0800 Subject: [PATCH] refactor: unify paged and contiguous attention kernels via KVSource policy - merge AttentionParams and PagedAttentionParams into one struct - add attn_kv_source.cuh with ContigKV/PagedKV addressing policies - template prefill/decode kernels (MMA + scalar) on the KV policy, deleting the four duplicated attn_paged_*.cuh variants - template dispatcher launchers on KV; single combine kernel - verify: all correctness tests pass and SASS matches baseline (no perf regression) --- csrc/kernels/attn_common.h | 65 ++--- csrc/kernels/attn_decode_split_kv.cuh | 50 ++-- csrc/kernels/attn_decode_split_kv_mma.cuh | 41 ++-- csrc/kernels/attn_dispatchers.cuh | 232 ++++++++---------- csrc/kernels/attn_entry_utils.cuh | 4 +- csrc/kernels/attn_kv_source.cuh | 159 ++++++++++++ csrc/kernels/attn_paged_decode.cu | 2 +- csrc/kernels/attn_paged_decode_split_kv.cuh | 151 ------------ .../attn_paged_decode_split_kv_mma.cuh | 178 -------------- csrc/kernels/attn_paged_prefill.cu | 2 +- csrc/kernels/attn_paged_prefill_split_q.cuh | 126 ---------- .../attn_paged_prefill_split_q_mma.cuh | 164 ------------- csrc/kernels/attn_prefill_split_q.cuh | 45 ++-- csrc/kernels/attn_prefill_split_q_mma.cuh | 50 ++-- csrc/tests/attn_paged_test.cu | 12 +- csrc/tests/test_utils.cuh | 2 +- docs/developer/cuda_kernels.md | 22 +- 17 files changed, 415 insertions(+), 890 deletions(-) create mode 100644 csrc/kernels/attn_kv_source.cuh delete mode 100644 csrc/kernels/attn_paged_decode_split_kv.cuh delete mode 100644 csrc/kernels/attn_paged_decode_split_kv_mma.cuh delete mode 100644 csrc/kernels/attn_paged_prefill_split_q.cuh delete mode 100644 csrc/kernels/attn_paged_prefill_split_q_mma.cuh diff --git a/csrc/kernels/attn_common.h b/csrc/kernels/attn_common.h index d02ab80..19b53c0 100644 --- a/csrc/kernels/attn_common.h +++ b/csrc/kernels/attn_common.h @@ -9,13 +9,19 @@ enum TensorLayout : int { }; +// Unified attention params covering BOTH addressing modes: +// - Contiguous K/V: dense [batch, kv_head, kv_len, head_dim] tensors (k/v). +// - Paged (SGLang-style): flat pool [size, kv_head, head_dim] + req_to_token. +// Each kernel selects the addressing via a KVSource policy (see +// attn_kv_source.cuh); a given call only touches the fields of one mode, so +// this is a POD shared by both paths rather than two parallel structs that +// drift out of sync. template struct AttentionParams { + // ---- shared across all paths ---- int batch; int q_head; int kv_head; - int q_len; - int kv_len; int head_dim; int use_mask; int causal_offset; // -1 = non-causal; >=0 = absolute position of first Q token @@ -24,54 +30,27 @@ struct AttentionParams { // Q strides (element offsets for each dim — layout-agnostic) int q_stride_b, q_stride_h, q_stride_l, q_stride_d; - // KV strides (K and V share the same layout — only base pointers differ) - int kv_stride_b, kv_stride_h, kv_stride_l, kv_stride_d; // Mask: 2D [batch, kv_len], 3D [batch, q_len, kv_len], // or 4D [batch, n_heads, q_len, kv_len] (head dim broadcasts when stride=0) int mask_b_stride; // batch stride int mask_h_stride; // head stride (0 = broadcast across heads) int mask_q_stride; // q stride (0 = all q rows share) - - const T* __restrict__ q; - const T* __restrict__ k; - const T* __restrict__ v; const bool* __restrict__ mask; + const T* __restrict__ q; T* __restrict__ o; AT* __restrict__ o_part; AT* __restrict__ ml_part; -}; -// ---- PagedAttentionParams ---- -// SGLang-style indirect params over a shared KV pool. -// k_cache/v_cache: [size, kv_head, head_dim] (bare buffers, no gather). -// req_to_token: [num_reqs, max_context_len] token -> slot. -// req_pool_indices:[batch] rows of the current batch into req_to_token. -// kv_indptr: [batch+1] prefix sum of per-request seq_lens (device). -// qo_indptr: [batch+1] prefix sum of per-request q_len (prefill) or -// nullptr for decode (q_len == 1 everywhere). -template -struct PagedAttentionParams { - int batch; - int q_head; - int kv_head; - int head_dim; - int num_splits; - int use_mask; - int causal_offset; // -1 = non-causal; >=0 = causal (per-request offset - // computed inside kernel from kv_indptr/qo_indptr) - float scale; + // ---- contiguous K/V mode ---- + int q_len; + int kv_len; + int kv_stride_b, kv_stride_h, kv_stride_l, kv_stride_d; + const T* __restrict__ k; + const T* __restrict__ v; - // Q: [total_q, q_head, head_dim] (3D flattened — no batch dim). - // For decode total_q == batch (q_len=1 per request). - // For prefill total_q == qo_indptr[batch]. - int q_stride_l, q_stride_h, q_stride_d; - - // Q: [total_q, q_head, head_dim] - const T* __restrict__ q; - - // Flat KV pool: [size, kv_head, head_dim] + // ---- paged (SGLang flat pool) mode ---- const T* __restrict__ k_cache; const T* __restrict__ v_cache; @@ -84,16 +63,4 @@ struct PagedAttentionParams { int max_seq_len; // max per-request seq_len (host-side, for split computation) int total_q; // total Q tokens across all requests (host-side, for grid) int max_q_len; // max per-request q_len (host-side, for prefill grid) - - // Mask: [batch, max_seq_len] (decode) or [batch, 1, q_len, kv_len] - // (prefill, optional). mask_h_stride/mask_q_stride are 0 when those - // dims are size 1 (broadcast). - int mask_b_stride; - int mask_h_stride; - int mask_q_stride; - const bool* __restrict__ mask; - - T* __restrict__ o; - AT* __restrict__ o_part; - AT* __restrict__ ml_part; }; diff --git a/csrc/kernels/attn_decode_split_kv.cuh b/csrc/kernels/attn_decode_split_kv.cuh index a5ad21d..7d6aabd 100644 --- a/csrc/kernels/attn_decode_split_kv.cuh +++ b/csrc/kernels/attn_decode_split_kv.cuh @@ -2,10 +2,16 @@ #include #include #include "attn_common.h" +#include "attn_kv_source.cuh" #include "attn_warp_utils.cuh" constexpr int DC_CHUNK = 64; -template +// Scalar split-KV decode (fallback for sm < 80, no tensor cores), unified +// across contiguous and paged (SGLang flat-pool) K/V via the KV template +// parameter. For decode the query is the last token, so its valid range +// [0, seq_len) IS the causal range; KV::decode_attend_len expresses that +// bound per addressing mode (contig clips to causal_offset, paged = seq_len). +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; @@ -15,15 +21,16 @@ __global__ void attn_decode_split_kv_kernel(AttentionParams p) { int lane = threadIdx.x; int hd_per_thread = p.head_dim / 32; + const int seq_len = KV::kv_len(p, batch); + const KVContext kctx = KV::template make_ctx(p, batch, kv_head); + // Q: [batch, q_head, q_len=1, head_dim] — stride-based float q_reg[8]; - int q_off = batch * p.q_stride_b + q_head * p.q_stride_h + int q_off = KV::q_decode_base(p, batch, q_head) + lane * hd_per_thread * p.q_stride_d; for (int i = 0; i < hd_per_thread; i++) q_reg[i] = __bfloat162float(p.q[q_off + i * p.q_stride_d]); - // KV: [batch, kv_head, kv_len, head_dim] — stride-based base - int kv_base = batch * p.kv_stride_b + kv_head * p.kv_stride_h; int mask_base = batch * p.mask_b_stride + q_head * p.mask_h_stride; float m = -FLT_MAX, d = 0.0f, acc_reg[8] = {0.0f}; @@ -31,24 +38,25 @@ __global__ void attn_decode_split_kv_kernel(AttentionParams p) { extern __shared__ __align__(16) bf16 k_smem[]; // Split-KV: each split processes a contiguous subset of chunks - int chunks_total = (p.kv_len + DC_CHUNK - 1) / DC_CHUNK; + int chunks_total = (seq_len + DC_CHUNK - 1) / DC_CHUNK; int chunks_per_split = (chunks_total + p.num_splits - 1) / p.num_splits; int ch_begin = split * chunks_per_split; int ch_end = min(chunks_total, ch_begin + chunks_per_split); for (int ci = ch_begin; ci < ch_end; ci++) { int chunk_start = ci * DC_CHUNK; - int this_chunk = min(DC_CHUNK, p.kv_len - chunk_start); + int this_chunk = min(DC_CHUNK, seq_len - chunk_start); - // Load K into shared memory (gather from strided global) + // Load K into shared memory (addressing via KV policy; paged guards + // empty slots with zero-fill). int total = this_chunk * p.head_dim; 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; - int g_off = kv_base + kv_idx * p.kv_stride_l + d_dim * p.kv_stride_d; - k_smem[i] = p.k[g_off]; + int kc = chunk_start + s; + KVAddr a = KV::kv_addr(p, kctx, kc, d_dim, true); + k_smem[i] = a.valid ? *reinterpret_cast(a.k) : (bf16)0.f; } __syncthreads(); @@ -65,7 +73,7 @@ __global__ void attn_decode_split_kv_kernel(AttentionParams p) { partial = -FLT_MAX; } if constexpr (IsCausal) { - if (kv_idx > p.causal_offset) + if (kv_idx >= KV::decode_attend_len(p, batch)) partial = -FLT_MAX; } @@ -74,11 +82,15 @@ __global__ void attn_decode_split_kv_kernel(AttentionParams p) { float beta = __expf(partial - new_m); d = d * alpha + beta; - 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] = fmaf(acc_reg[i], alpha, - __bfloat162float(p.v[v_off + i * p.kv_stride_d]) * beta); + // V read via KV policy; when masked (beta == 0) or the slot is + // empty the term vanishes, so no extra branches are needed. + for (int i = 0; i < hd_per_thread; i++) { + KVAddr a = KV::kv_addr(p, kctx, kv_idx, lane * hd_per_thread + i, true); + float vv = a.valid + ? __bfloat162float(*reinterpret_cast(a.v)) + : 0.0f; + acc_reg[i] = fmaf(acc_reg[i], alpha, vv * beta); + } m = new_m; } __syncthreads(); @@ -98,6 +110,10 @@ __global__ void attn_decode_split_kv_kernel(AttentionParams p) { } } +// Split-combine: merges the per-split partials (o_part/ml_part) into the +// final normalised O. KV selects the O addressing (contig batch stride vs +// paged row stride). +template __global__ void attn_decode_combine_kernel(AttentionParams p) { int bh = blockIdx.x; int d = threadIdx.x; @@ -124,6 +140,6 @@ __global__ void attn_decode_combine_kernel(AttentionParams p) { } float inv = (l > 1e-20f) ? (1.0f / l) : 0.0f; - int o_off = batch * p.q_stride_b + q_head * p.q_stride_h + d * p.q_stride_d; + int o_off = KV::q_decode_base(p, batch, q_head) + 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 6d45a51..b9ca271 100644 --- a/csrc/kernels/attn_decode_split_kv_mma.cuh +++ b/csrc/kernels/attn_decode_split_kv_mma.cuh @@ -2,19 +2,22 @@ #include #include #include "attn_common.h" +#include "attn_kv_source.cuh" #include "attn_mma_utils.cuh" #include "attn_warp_utils.cuh" -// 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. +// Split-K (FlashDecoding) tensor-core decode via GQA head-packing, unified +// across contiguous and paged (SGLang flat-pool) K/V via the KV template +// parameter. 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. // +// KV = ContigKV (dense tensors) or PagedKV (flat pool + req_to_token). // IsCausal and HasMask are compile-time bools — no runtime branch in the // inner compute loop. // -// Traits = KernelTraits>. -template +// Traits = KernelTraits. +template __global__ void attn_decode_split_kv_mma_kernel(AttentionParams p) { const int lane = threadIdx.x; const int gid = lane >> 2; @@ -31,13 +34,16 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams p) { const int G = min(MAX_G, G_total - g_begin); const int q_head0 = kv_head * G_total + g_begin; + // Per-request seq_len (paged reads kv_indptr; contig uses p.kv_len). + const int seq_len = KV::kv_len(p, batch); + const KVContext kctx = KV::template make_ctx(p, batch, kv_head); + // 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. - // 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 q_base = KV::q_decode_base(p, batch, q_head0); const int qra = gid; const int qrb = gid + 8; const bool va = qra < G, vb = qrb < G; @@ -51,13 +57,12 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams p) { 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 kv_base = batch * p.kv_stride_b + kv_head * p.kv_stride_h; - const int tiles_total = (p.kv_len + Traits::BC - 1) / Traits::BC; + const int tiles_total = (seq_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); - // ---- Load tile lambda: predicated cp.async ---- + // ---- Load tile lambda: predicated cp.async (addressing via KV policy) ---- auto load_tile = [&](int ti, int buf) { int kv0 = ti * Traits::BC; bf16* dK = sK + buf * Traits::BC * Traits::LD; @@ -67,11 +72,11 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams p) { 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; + bool valid = kc < seq_len; + KVAddr a = KV::kv_addr(p, kctx, kc, d, valid); 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); + cp_async_16_pred(&dK[off], a.k, a.valid); + cp_async_16_pred(&dV[off], a.v, a.valid); } cp_async_commit(); }; @@ -96,8 +101,10 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams p) { 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 - int maxc = IsCausal ? min(p.kv_len, p.causal_offset + 1) : p.kv_len; + // Decode: q_len=1, so qrow0=qrow1=0. Paged treats [0, seq_len) as + // the causal range (query is the last token); contig clips to the + // causal_offset bound. Dead code eliminated when IsCausal == false. + int maxc = IsCausal ? KV::decode_attend_len(p, batch) : seq_len; mma_softmax_tile(kv0, maxc, maxc, 0, 0, p.mask_b_stride, 0, 0, diff --git a/csrc/kernels/attn_dispatchers.cuh b/csrc/kernels/attn_dispatchers.cuh index acc5744..33515a7 100644 --- a/csrc/kernels/attn_dispatchers.cuh +++ b/csrc/kernels/attn_dispatchers.cuh @@ -1,19 +1,22 @@ #pragma once // Shared attention dispatchers — used by both production .cu and test .cu. // No torch dependency; pure CUDA. +// +// The paged and contiguous kernels are unified by the KVSource policy +// (ContigKV / PagedKV from attn_kv_source.cuh), so each launcher struct +// below is templated on KV and the paged dispatch is just the same launcher +// instantiated with PagedKV. Only the grid/split math differs, and that is +// covered by KV::host_q_len / KV::host_kv_len. #include #include #include "attn_warp_utils.cuh" +#include "attn_kv_source.cuh" #include "attn_prefill_split_q.cuh" #include "attn_decode_split_kv.cuh" -#include "attn_paged_decode_split_kv.cuh" -#include "attn_paged_prefill_split_q.cuh" #ifndef ASTRAI_NO_MMA #include "attn_prefill_split_q_mma.cuh" #include "attn_decode_split_kv_mma.cuh" -#include "attn_paged_decode_split_kv_mma.cuh" -#include "attn_paged_prefill_split_q_mma.cuh" #endif // Split-KV: compute number of splits to fill all SMs for small-batch decode. @@ -39,7 +42,7 @@ inline int compute_num_splits(int base_blocks, int tiles_total, // template ; HEAD_DIM is forwarded // as the first template argument so callers only spell it once. // -// Usage: DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_decode_mma, HEAD_DIM, p, group_size); +// Usage: DISPATCH_CAUSAL_MASK(is_causal, has_mask, launcher::template launch, HEAD_DIM, p, stream); #define DISPATCH_CAUSAL_MASK(is_causal, has_mask, FN, HEAD_DIM, ...) \ do { \ if (is_causal) { \ @@ -52,28 +55,39 @@ inline int compute_num_splits(int base_blocks, int tiles_total, } while (0) // ====================================================================== -// Prefill +// Prefill launchers (KV selects ContigKV or PagedKV addressing) // ====================================================================== #ifndef ASTRAI_NO_MMA -template -static inline void launch_prefill_mma(AttentionParams& p, cudaStream_t stream) { - 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); - attn_prefill_split_q_mma_kernel<<>>(p); -} +template +struct PrefillLauncherMMA { + template + static void launch(AttentionParams& p, cudaStream_t stream) { + constexpr int WARPS = 4; + constexpr int BC = (HEAD_DIM <= 128) ? 32 : 16; + using Traits = KernelTraits; + int q_len = KV::host_q_len(p); + dim3 grid((q_len + Traits::BR * WARPS - 1) / (Traits::BR * WARPS), + p.q_head, p.batch); + dim3 block(Traits::NUM_THREADS); + attn_prefill_split_q_mma_kernel + <<>>(p); + } +}; #endif -template -static inline void launch_prefill_scalar(AttentionParams& p, cudaStream_t stream) { - 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); - attn_prefill_split_q_kernel_t<<>>(p); -} +template +struct PrefillLauncherScalar { + template + static void launch(AttentionParams& p, cudaStream_t stream) { + constexpr int G = 8, ROWS = 32, P_BC = 32; + int q_len = KV::host_q_len(p); + dim3 grid((q_len + ROWS - 1) / ROWS, p.q_head, p.batch); + dim3 block(G, ROWS); + attn_prefill_split_q_kernel_t + <<>>(p); + } +}; template static inline void dispatch_prefill(AttentionParams& p, cudaStream_t stream) { @@ -81,14 +95,34 @@ static inline void dispatch_prefill(AttentionParams& p, cudaStream_t strea bool has_mask = (p.use_mask && p.mask); #ifndef ASTRAI_NO_MMA - DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_prefill_mma, HEAD_DIM, p, stream); + DISPATCH_CAUSAL_MASK(is_causal, has_mask, + PrefillLauncherMMA::template launch, + HEAD_DIM, p, stream); #else - DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_prefill_scalar, HEAD_DIM, p, stream); + DISPATCH_CAUSAL_MASK(is_causal, has_mask, + PrefillLauncherScalar::template launch, + HEAD_DIM, p, stream); +#endif +} + +template +static inline void dispatch_paged_prefill(AttentionParams& p, cudaStream_t stream) { + bool is_causal = (p.causal_offset >= 0); + bool has_mask = (p.use_mask && p.mask); + +#ifndef ASTRAI_NO_MMA + DISPATCH_CAUSAL_MASK(is_causal, has_mask, + PrefillLauncherMMA::template launch, + HEAD_DIM, p, stream); +#else + DISPATCH_CAUSAL_MASK(is_causal, has_mask, + PrefillLauncherScalar::template launch, + HEAD_DIM, p, stream); #endif } // ====================================================================== -// Decode +// Decode launchers (KV selects ContigKV or PagedKV addressing) // ====================================================================== #ifndef ASTRAI_NO_MMA @@ -96,31 +130,41 @@ static inline void dispatch_prefill(AttentionParams& p, cudaStream_t strea // For D=256, BC=16 also reduces register pressure (fewer Sacc/PV frags), // enabling STAGES=2 (double-buffer) within the 32KB smem budget — eliminates // the 176-byte spill that STAGES=1+BC=32 suffered. -template -static inline void launch_decode_mma(AttentionParams& p, int group_size, cudaStream_t stream) { - int G = p.q_head / p.kv_head; - constexpr int MAX_G = 16; - int num_passes = (G + MAX_G - 1) / MAX_G; - constexpr int BC = 16; - int tiles_total = (p.kv_len + BC - 1) / BC; - p.num_splits = compute_num_splits(p.batch * p.kv_head * num_passes, tiles_total, 2); - constexpr int STAGES = 2; - using Traits = KernelTraits; - dim3 grid(p.kv_head * num_passes, p.batch, p.num_splits); - attn_decode_split_kv_mma_kernel<<>>(p); -} +template +struct DecodeLauncherMMA { + template + static void launch(AttentionParams& p, int group_size, cudaStream_t stream) { + int G = p.q_head / p.kv_head; + constexpr int MAX_G = 16; + int num_passes = (G + MAX_G - 1) / MAX_G; + constexpr int BC = 16; + int kv_len = KV::host_kv_len(p); + int tiles_total = (kv_len + BC - 1) / BC; + p.num_splits = compute_num_splits(p.batch * p.kv_head * num_passes, tiles_total, 2); + constexpr int STAGES = 2; + using Traits = KernelTraits; + dim3 grid(p.kv_head * num_passes, p.batch, p.num_splits); + attn_decode_split_kv_mma_kernel + <<>>(p); + } +}; #endif -template -static inline void launch_decode_scalar(AttentionParams& p, int group_size, cudaStream_t stream) { - int chunks_total = (p.kv_len + DC_CHUNK - 1) / DC_CHUNK; - p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total); - size_t smem = DC_CHUNK * p.head_dim * sizeof(bf16); - int g = min(group_size, 32); // cap at 32 to respect 1024-thread limit - dim3 grid(p.batch * p.kv_head, 1, p.num_splits); - dim3 block(32, g); - attn_decode_split_kv_kernel<<>>(p); -} +template +struct DecodeLauncherScalar { + template + static void launch(AttentionParams& p, int group_size, cudaStream_t stream) { + int kv_len = KV::host_kv_len(p); + int chunks_total = (kv_len + DC_CHUNK - 1) / DC_CHUNK; + p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total); + size_t smem = DC_CHUNK * p.head_dim * sizeof(bf16); + int g = min(group_size, 32); // cap at 32 to respect 1024-thread limit + dim3 grid(p.batch * p.kv_head, 1, p.num_splits); + dim3 block(32, g); + attn_decode_split_kv_kernel + <<>>(p); + } +}; template static inline void dispatch_decode(AttentionParams& p, cudaStream_t stream) { @@ -129,95 +173,33 @@ static inline void dispatch_decode(AttentionParams& p, cudaStream_t stream int group_size = p.q_head / p.kv_head; #ifndef ASTRAI_NO_MMA - DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_decode_mma, HEAD_DIM, p, group_size, stream); + DISPATCH_CAUSAL_MASK(is_causal, has_mask, + DecodeLauncherMMA::template launch, + HEAD_DIM, p, group_size, stream); #else - DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_decode_scalar, HEAD_DIM, p, group_size, stream); + DISPATCH_CAUSAL_MASK(is_causal, has_mask, + DecodeLauncherScalar::template launch, + HEAD_DIM, p, group_size, stream); #endif - attn_decode_combine_kernel<<>>(p); -} - -// ====================================================================== -// Paged Decode (SGLang-style: flat pool + req_to_token + kv_indptr) -// ====================================================================== - -#ifndef ASTRAI_NO_MMA -template -static inline void launch_paged_decode_mma(PagedAttentionParams& p, cudaStream_t stream) { - int G = p.q_head / p.kv_head; - constexpr int MAX_G = 16; - constexpr int BC = 16; - int num_passes = (G + MAX_G - 1) / MAX_G; - int tiles_total = (p.max_seq_len + BC - 1) / BC; - p.num_splits = compute_num_splits(p.batch * p.kv_head * num_passes, tiles_total, 2); - constexpr int STAGES = 2; - using Traits = KernelTraits; - dim3 grid(p.kv_head * num_passes, p.batch, p.num_splits); - paged_attn_decode_split_kv_mma_kernel <<>>(p); -} -#endif - -template -static inline void launch_paged_decode_scalar(PagedAttentionParams& p, int group_size, cudaStream_t stream) { - int chunks_total = (p.max_seq_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); - int g = min(group_size, 32); - dim3 grid(p.batch * p.kv_head, 1, p.num_splits); - dim3 block(32, g); - paged_attn_decode_split_kv_kernel<<>>(p); + attn_decode_combine_kernel<<>>(p); } template -static inline void dispatch_paged_decode(PagedAttentionParams& p, cudaStream_t stream) { +static inline void dispatch_paged_decode(AttentionParams& p, cudaStream_t stream) { bool is_causal = (p.causal_offset >= 0); bool has_mask = (p.use_mask && p.mask); int group_size = p.q_head / p.kv_head; #ifndef ASTRAI_NO_MMA - DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_paged_decode_mma, HEAD_DIM, p, stream); + DISPATCH_CAUSAL_MASK(is_causal, has_mask, + DecodeLauncherMMA::template launch, + HEAD_DIM, p, group_size, stream); #else - DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_paged_decode_scalar, HEAD_DIM, p, group_size, stream); + DISPATCH_CAUSAL_MASK(is_causal, has_mask, + DecodeLauncherScalar::template launch, + HEAD_DIM, p, group_size, stream); #endif - paged_attn_decode_combine_kernel<<>>(p); -} - -// ====================================================================== -// Paged Prefill (SGLang-style: flat pool + ragged batch) -// ====================================================================== - -#ifndef ASTRAI_NO_MMA -template -static inline void launch_paged_prefill_mma(PagedAttentionParams& p, cudaStream_t stream) { - constexpr int WARPS = 4; - constexpr int BC = (HEAD_DIM <= 128) ? 32 : 16; - using Traits = KernelTraits; - int max_q_tiles = (p.max_q_len + Traits::BR * WARPS - 1) / (Traits::BR * WARPS); - dim3 grid(max_q_tiles, p.q_head, p.batch); - dim3 block(Traits::NUM_THREADS); - paged_attn_prefill_split_q_mma_kernel<<>>(p); -} -#endif - -template -static inline void launch_paged_prefill_scalar(PagedAttentionParams& p, cudaStream_t stream) { - constexpr int G = 8, ROWS = 32, P_BC = 32; - int max_q_tiles = (p.max_q_len + ROWS - 1) / ROWS; - dim3 grid(max_q_tiles, p.q_head, p.batch); - dim3 block(G, ROWS); - paged_attn_prefill_split_q_kernel - <<>>(p); -} - -template -static inline void dispatch_paged_prefill(PagedAttentionParams& p, cudaStream_t stream) { - bool is_causal = (p.causal_offset >= 0); - bool has_mask = (p.use_mask && p.mask); - -#ifndef ASTRAI_NO_MMA - DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_paged_prefill_mma, HEAD_DIM, p, stream); -#else - DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_paged_prefill_scalar, HEAD_DIM, p, stream); -#endif + attn_decode_combine_kernel<<>>(p); } diff --git a/csrc/kernels/attn_entry_utils.cuh b/csrc/kernels/attn_entry_utils.cuh index 5575bb1..90c3618 100644 --- a/csrc/kernels/attn_entry_utils.cuh +++ b/csrc/kernels/attn_entry_utils.cuh @@ -149,7 +149,7 @@ inline void attn_pack_paged_decode_params( c10::optional mask, int64_t causal_offset, double scale, - PagedAttentionParams& p + AttentionParams& p ) { const at::cuda::OptionalCUDAGuard device_guard(device_of(q)); @@ -229,7 +229,7 @@ inline void attn_pack_paged_prefill_params( int64_t max_q_len, int64_t causal_offset, double scale, - PagedAttentionParams& p + AttentionParams& p ) { const at::cuda::OptionalCUDAGuard device_guard(device_of(q)); diff --git a/csrc/kernels/attn_kv_source.cuh b/csrc/kernels/attn_kv_source.cuh new file mode 100644 index 0000000..416cd9c --- /dev/null +++ b/csrc/kernels/attn_kv_source.cuh @@ -0,0 +1,159 @@ +#pragma once +#include +#include "attn_common.h" + +// ============================================================================ +// KVSource policies — the single dimension along which the paged and +// non-paged attention kernels differ. Each kernel is templated on one of +// these (ContigKV / PagedKV) and stays fully generic: the policy owns every +// place where "where does K/V live" and "what is this request's seq_len" +// are answered. All methods are __host__ __device__ so the same policy +// serves both the device kernels (addressing, seq_len) and the host-side +// launchers (grid / split computation). +// +// ContigKV: K/V are dense [batch, kv_head, kv_len, head_dim] tensors. +// Params fields used: k, v, kv_stride_*, kv_len, q_len, +// q_stride_b, causal_offset. +// PagedKV: K/V live in a flat pool [size, kv_head, head_dim] indexed via +// req_to_token. Params fields used: k_cache, v_cache, +// req_to_token, req_pool_indices, kv_indptr, qo_indptr, +// max_context_len, q_stride_l. +// +// Addressing state that is constant across a whole kernel invocation for one +// (batch, kv_head) pair is captured once by make_ctx() and passed +// to kv_addr, so the load loops never redo the hoistable base computation +// (e.g. the req_pool_indices global read) element-by-element. +// ============================================================================ + +// Every policy method is static + callable from both host and device code. +#define HOST_DEV_FORCEINLINE static __host__ __device__ __forceinline__ + +using bf16 = __nv_bfloat16; + +// Hoisted per-(batch, kv_head) addressing context. +struct KVContext { + int kv_base; // contig: batch*kv_stride_b + kv_head*kv_stride_h + int64_t req_idx; // paged: req_pool_indices[batch] + int64_t rtt_stride; // paged: max_context_len + int64_t pool_stride; // paged: kv_head * HEAD_DIM + int64_t head_off; // paged: kv_head * HEAD_DIM +}; + +// Per-element K/V global addresses for one (kc, d) position of a K/V tile. +// The pointers are ALWAYS the computed addresses (never nullptr) — callers +// gate on `valid` (cp.async src_size=0, or a guarded scalar deref). `valid` +// starts as "within the request's seq_len"; the paged policy further degrades +// it when req_to_token maps the position to a negative slot (empty padding). +// This matches the original hand-rolled load loops, where the address was +// always formed and the predicate decided whether anything was read. +struct KVAddr { + const void* k; + const void* v; + bool valid; +}; + +// ---- Contiguous K/V ---- +struct ContigKV { + static constexpr bool kPaged = false; + + // host-side length hooks (grid + split computation in the launchers) + HOST_DEV_FORCEINLINE int host_q_len(const AttentionParams& p) { + return p.q_len; + } + HOST_DEV_FORCEINLINE int host_kv_len(const AttentionParams& p) { + return p.kv_len; + } + + // prefill: element offset of the request's Q rows (kernel adds qrow*q_stride_l) + HOST_DEV_FORCEINLINE int q_base( + const AttentionParams& p, int batch, int q_head) { + return batch * p.q_stride_b + q_head * p.q_stride_h; + } + // decode: same offset (q_len == 1, so there is no row stride component) + HOST_DEV_FORCEINLINE int q_decode_base( + const AttentionParams& p, int batch, int q_head) { + return batch * p.q_stride_b + q_head * p.q_stride_h; + } + + HOST_DEV_FORCEINLINE int kv_len(const AttentionParams& p, int batch) { + return p.kv_len; + } + HOST_DEV_FORCEINLINE int q_len(const AttentionParams& p, int batch) { + return p.q_len; + } + HOST_DEV_FORCEINLINE int causal_offset(const AttentionParams& p, int batch) { + return p.causal_offset; + } + // decode: exclusive bound of the single query's attend range + HOST_DEV_FORCEINLINE int decode_attend_len(const AttentionParams& p, int batch) { + return (p.kv_len < p.causal_offset + 1) ? p.kv_len : (p.causal_offset + 1); + } + + template + HOST_DEV_FORCEINLINE KVContext make_ctx( + const AttentionParams& p, int batch, int kv_head) { + KVContext c = {}; + c.kv_base = batch * p.kv_stride_b + kv_head * p.kv_stride_h; + return c; + } + HOST_DEV_FORCEINLINE KVAddr kv_addr( + const AttentionParams& p, const KVContext& c, int kc, int d, bool valid) { + const int g_off = c.kv_base + kc * p.kv_stride_l + d * p.kv_stride_d; + return {&p.k[g_off], &p.v[g_off], valid}; + } +}; + +// ---- Paged (SGLang-style flat pool) K/V ---- +struct PagedKV { + static constexpr bool kPaged = true; + + HOST_DEV_FORCEINLINE int host_q_len(const AttentionParams& p) { + return p.max_q_len; + } + HOST_DEV_FORCEINLINE int host_kv_len(const AttentionParams& p) { + return p.max_seq_len; + } + + // prefill: Q rows start at qo_indptr[batch] (ragged batch base) + HOST_DEV_FORCEINLINE int q_base( + const AttentionParams& p, int batch, int q_head) { + return p.qo_indptr[batch] * p.q_stride_l + q_head * p.q_stride_h; + } + // decode: Q is [batch, q_head, head_dim], so batch is the outer row + HOST_DEV_FORCEINLINE int q_decode_base( + const AttentionParams& p, int batch, int q_head) { + return batch * p.q_stride_l + q_head * p.q_stride_h; + } + + HOST_DEV_FORCEINLINE int kv_len(const AttentionParams& p, int batch) { + return p.kv_indptr[batch + 1] - p.kv_indptr[batch]; + } + HOST_DEV_FORCEINLINE int q_len(const AttentionParams& p, int batch) { + return p.qo_indptr[batch + 1] - p.qo_indptr[batch]; + } + HOST_DEV_FORCEINLINE int causal_offset(const AttentionParams& p, int batch) { + return kv_len(p, batch) - q_len(p, batch); + } + // decode: the query is the last token, so [0, seq_len) IS its causal range + HOST_DEV_FORCEINLINE int decode_attend_len(const AttentionParams& p, int batch) { + return kv_len(p, batch); + } + + template + HOST_DEV_FORCEINLINE KVContext make_ctx( + const AttentionParams& p, int batch, int kv_head) { + KVContext c = {}; + c.req_idx = p.req_pool_indices[batch]; + c.rtt_stride = (int64_t)p.max_context_len; + c.pool_stride = (int64_t)p.kv_head * HEAD_DIM; + c.head_off = (int64_t)kv_head * HEAD_DIM; + return c; + } + HOST_DEV_FORCEINLINE KVAddr kv_addr( + const AttentionParams& p, const KVContext& c, int kc, int d, bool valid) { + const int64_t slot = valid ? p.req_to_token[c.req_idx * c.rtt_stride + kc] : 0; + const bool ok = valid && (slot >= 0); + const int64_t gmem_off = slot * c.pool_stride + c.head_off + d; + return {&p.k_cache[gmem_off], &p.v_cache[gmem_off], ok}; + } +}; diff --git a/csrc/kernels/attn_paged_decode.cu b/csrc/kernels/attn_paged_decode.cu index 97dc44c..22caccc 100644 --- a/csrc/kernels/attn_paged_decode.cu +++ b/csrc/kernels/attn_paged_decode.cu @@ -16,7 +16,7 @@ torch::Tensor attn_paged_decode( const at::cuda::OptionalCUDAGuard device_guard(device_of(q)); auto stream = at::cuda::getCurrentCUDAStream(); - PagedAttentionParams p; + AttentionParams p; attn_pack_paged_decode_params(q, k_cache, v_cache, req_to_token, req_pool_indices, kv_indptr, max_seq_len, mask, causal_offset, scale, p); diff --git a/csrc/kernels/attn_paged_decode_split_kv.cuh b/csrc/kernels/attn_paged_decode_split_kv.cuh deleted file mode 100644 index 6d7a156..0000000 --- a/csrc/kernels/attn_paged_decode_split_kv.cuh +++ /dev/null @@ -1,151 +0,0 @@ -#pragma once -#include -#include -#include "attn_common.h" -#include "attn_warp_utils.cuh" -constexpr int PDC_CHUNK = 64; - -// Scalar paged decode (fallback for sm < 80, no tensor cores). -// Reads K/V from flat pool via req_to_token indexing. -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; - int split = blockIdx.z; - int group_size = blockDim.y; - int q_head = kv_head * group_size + threadIdx.y; - int lane = threadIdx.x; - int hd_per_thread = p.head_dim / 32; - - const int seq_len = p.kv_indptr[batch + 1] - p.kv_indptr[batch]; - const int64_t req_idx = p.req_pool_indices[batch]; - - float q_reg[8]; - int q_off = batch * p.q_stride_l + q_head * p.q_stride_h - + lane * hd_per_thread * p.q_stride_d; - #pragma unroll - for (int i = 0; i < hd_per_thread; i++) - q_reg[i] = __bfloat162float(p.q[q_off + i * p.q_stride_d]); - - float m = -FLT_MAX, d = 0.0f, acc_reg[8] = {0.0f}; - - extern __shared__ __align__(16) bf16 k_smem[]; - - int chunks_total = (seq_len + PDC_CHUNK - 1) / PDC_CHUNK; - int chunks_per_split = (chunks_total + p.num_splits - 1) / p.num_splits; - int ch_begin = split * chunks_per_split; - int ch_end = min(chunks_total, ch_begin + chunks_per_split); - - const int mask_base = batch * p.mask_b_stride; - const int64_t pool_stride = (int64_t)p.kv_head * p.head_dim; - const int64_t head_off = (int64_t)kv_head * p.head_dim; - const int64_t rtt_stride = (int64_t)p.max_context_len; - - for (int ci = ch_begin; ci < ch_end; ci++) { - int chunk_start = ci * PDC_CHUNK; - int this_chunk = min(PDC_CHUNK, seq_len - chunk_start); - - int total = this_chunk * p.head_dim; - 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; - int64_t slot = p.req_to_token[req_idx * rtt_stride + pos]; - if (slot >= 0) { - int64_t off = slot * pool_stride + head_off + d_dim; - k_smem[i] = p.k_cache[off]; - } else { - k_smem[i] = __float2bfloat16(0.0f); - } - } - __syncthreads(); - - for (int s = 0; s < this_chunk; s++) { - 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 = warp_reduce_sum(partial) * p.scale; - - int kv_idx = chunk_start + s; - bool masked = false; - if constexpr (HasMask) { - if (!p.mask[mask_base + kv_idx]) - masked = true; - } - // Decode: the query is the last token, so its valid range [0, - // seq_len) IS the causal range. IsCausal is accepted for dispatch - // uniformity but must not apply causal_offset masking here. - if (masked) - 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; - - int pos = chunk_start + s; - int64_t slot = p.req_to_token[req_idx * rtt_stride + pos]; - if (masked) { - #pragma unroll - for (int i = 0; i < hd_per_thread; i++) - acc_reg[i] = fmaf(acc_reg[i], alpha, 0.0f); - } else if (slot >= 0) { - int64_t v_base = slot * pool_stride + head_off; - #pragma unroll - for (int i = 0; i < hd_per_thread; i++) - acc_reg[i] = fmaf(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++) - acc_reg[i] = fmaf(acc_reg[i], alpha, 0.0f); - } - m = new_m; - } - __syncthreads(); - } - - size_t bh = (size_t)batch * p.q_head + q_head; - size_t slot = bh * MAX_SPLITS + split; - int d0 = lane * hd_per_thread; - #pragma unroll - for (int i = 0; i < hd_per_thread; i++) - p.o_part[slot * p.head_dim + (d0 + i)] = acc_reg[i]; - if (lane == 0) { - p.ml_part[slot * 2] = m; - p.ml_part[slot * 2 + 1] = d; - } -} - -__global__ void paged_attn_decode_combine_kernel(PagedAttentionParams p) { - int bh = blockIdx.x; - int d = threadIdx.x; - if (d >= p.head_dim) return; - - int batch = bh / p.q_head; - int q_head = bh % p.q_head; - - size_t split_base = (size_t)bh * MAX_SPLITS; - const float* mlp = p.ml_part + split_base * 2; - const float* op = p.o_part + split_base * p.head_dim; - - float m = -FLT_MAX, l = 0.0f, acc = 0.0f; - for (int s = 0; s < p.num_splits; s++) { - float mi = mlp[s * 2]; - if (mi <= -FLT_MAX) continue; - float li = mlp[s * 2 + 1]; - float nm = fmaxf(m, mi); - float corr = __expf(m - nm); - float e = __expf(mi - nm); - acc = fmaf(acc, corr, op[s * p.head_dim + d] * e); - l = fmaf(l, corr, li * e); - m = nm; - } - - float inv = (l > 1e-20f) ? (1.0f / l) : 0.0f; - int o_off = batch * p.q_stride_l + q_head * p.q_stride_h + d * p.q_stride_d; - p.o[o_off] = __float2bfloat16(acc * inv); -} diff --git a/csrc/kernels/attn_paged_decode_split_kv_mma.cuh b/csrc/kernels/attn_paged_decode_split_kv_mma.cuh deleted file mode 100644 index a67bf72..0000000 --- a/csrc/kernels/attn_paged_decode_split_kv_mma.cuh +++ /dev/null @@ -1,178 +0,0 @@ -#pragma once -#include -#include -#include "attn_common.h" -#include "attn_mma_utils.cuh" -#include "attn_warp_utils.cuh" - -// SGLang-style split-KV tensor-core decode. -// -// Reads K/V directly from a flat pool [size, kv_head, head_dim] via -// req_to_token indexing — no gather, no page-table dimension. -// Each batch element has its own seq_len (from kv_indptr), eliminating -// padding waste: short sequences only process the tiles they own. -// -// For decode (q_len=1), causal masking is implicit — each request attends -// to [0, seq_len) which is exactly its valid range. The IsCausal flag -// is accepted for dispatch uniformity but does not change maxc. -template -__global__ void paged_attn_decode_split_kv_mma_kernel(PagedAttentionParams p) { - const int lane = threadIdx.x; - const int gid = lane >> 2; - const int tid4 = lane & 3; - - const int pass = blockIdx.x / p.kv_head; - const int kv_head = blockIdx.x % p.kv_head; - const int batch = blockIdx.y; - const int split = blockIdx.z; - - // Per-request seq_len from device-side kv_indptr — no padding. - const int seq_len = p.kv_indptr[batch + 1] - p.kv_indptr[batch]; - const int64_t req_idx = p.req_pool_indices[batch]; - - constexpr int MAX_G = 16; - const int G_total = p.q_head / p.kv_head; - const int g_begin = pass * MAX_G; - const int G = min(MAX_G, G_total - g_begin); - const int q_head0 = kv_head * G_total + g_begin; - - __shared__ __align__(16) bf16 sK[Traits::STAGES * Traits::BC * Traits::LD]; - __shared__ __align__(16) bf16 sV[Traits::STAGES * Traits::BC * Traits::LD]; - - const int q_base = batch * p.q_stride_l + 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[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[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 = (seq_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); - - // Flat pool stride: [size, kv_head, head_dim] — contiguous. - const int64_t pool_stride = (int64_t)p.kv_head * Traits::HEAD_DIM; - const int64_t head_off = (int64_t)kv_head * Traits::HEAD_DIM; - const int64_t rtt_stride = (int64_t)p.max_context_len; - - // ---- Load tile lambda: SGLang addressing ---- - // slot = req_to_token[req_idx * max_context_len + kc] - // gmem = k_cache[slot * pool_stride + head_off + d] - auto load_tile = [&](int ti, int buf) { - 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 < seq_len); - if constexpr (HasMask) { - valid = valid && p.mask[batch * p.mask_b_stride + kc]; - } - int64_t slot = valid ? p.req_to_token[req_idx * rtt_stride + kc] : 0; - valid = valid && (slot >= 0); - int64_t gmem_base = slot * pool_stride + head_off; - 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(); - }; - - constexpr int STAGES = Traits::STAGES; - const int ntiles = ti_end - ti_begin; - - auto process_tile = [&](int it, int buf) { - const bf16* bK = sK + buf * Traits::BC * Traits::LD; - const bf16* bV = sV + buf * Traits::BC * Traits::LD; - int kv0 = (ti_begin + it) * Traits::BC; - - float Sacc[Traits::NC8][4]; - mma_compute_scores(Qa, bK, lane, Sacc); - - #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; - - // For decode, maxc = seq_len regardless of IsCausal — the valid - // range [0, seq_len) IS the causal range (query is the last token). - mma_softmax_tile(kv0, seq_len, seq_len, - 0, 0, - p.mask_b_stride, 0, 0, - batch, 0, - p.mask, - Sacc, Oacc, m0, m1, l0, l1, lane); - - mma_pv_accumulate(Sacc, bV, lane, Oacc); - }; - - if (ntiles >= STAGES) { - #pragma unroll - for (int i = 0; i < STAGES; i++) - load_tile(ti_begin + i, i); - - for (int it = 0; it < ntiles; it++) { - cp_async_wait_group(); - __syncwarp(); - process_tile(it, it & (STAGES - 1)); - __syncwarp(); - if (it + STAGES < ntiles) - load_tile(ti_begin + it + STAGES, (it + STAGES) & (STAGES - 1)); - } - } else { - for (int i = 0; i < ntiles; i++) - load_tile(ti_begin + i, i); - cp_async_wait_group<0>(); - __syncwarp(); - for (int it = 0; it < ntiles; it++) - process_tile(it, it); - } - - // ---- write partials ---- - auto split_slot = [&](int h) -> size_t { - size_t bh = (size_t)batch * p.q_head + h; - return bh * MAX_SPLITS + split; - }; - #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) * 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) * Traits::HEAD_DIM; - op[d] = Oacc[dn8][2]; - op[d + 1] = Oacc[dn8][3]; - } - } - if (tid4 == 0) { - int r0 = gid, r1 = gid + 8; - if (r0 < G) { - int h = q_head0 + r0; - float* mp = p.ml_part + split_slot(h) * 2; - mp[0] = m0; mp[1] = l0; - } - if (r1 < G) { - int h = q_head0 + r1; - float* mp = p.ml_part + split_slot(h) * 2; - mp[0] = m1; mp[1] = l1; - } - } -} diff --git a/csrc/kernels/attn_paged_prefill.cu b/csrc/kernels/attn_paged_prefill.cu index 4848ced..87ecd1b 100644 --- a/csrc/kernels/attn_paged_prefill.cu +++ b/csrc/kernels/attn_paged_prefill.cu @@ -17,7 +17,7 @@ torch::Tensor attn_paged_prefill( const at::cuda::OptionalCUDAGuard device_guard(device_of(q)); auto stream = at::cuda::getCurrentCUDAStream(); - PagedAttentionParams p; + AttentionParams p; attn_pack_paged_prefill_params(q, k_cache, v_cache, req_to_token, req_pool_indices, kv_indptr, qo_indptr, mask, diff --git a/csrc/kernels/attn_paged_prefill_split_q.cuh b/csrc/kernels/attn_paged_prefill_split_q.cuh deleted file mode 100644 index a0e17b8..0000000 --- a/csrc/kernels/attn_paged_prefill_split_q.cuh +++ /dev/null @@ -1,126 +0,0 @@ -#pragma once -#include -#include -#include "attn_common.h" - -using bf16 = __nv_bfloat16; - -// Scalar paged prefill (fallback for sm < 80, no tensor cores). -// Reads K/V from a flat pool via req_to_token, supports ragged batches -// via qo_indptr + kv_indptr. Mirrors the split-Q MMA kernel's indexing: -// grid (max_q_tiles, q_head, batch), block (G, ROWS). -// -// HasMask: 4D mask [batch, 1, q_len, kv_len] (True=keep), columns are -// request-local kv positions. q_head is the q-index (mask_h broadcast). -// -// group_reduce_sum is provided by attn_prefill_split_q.cuh (already -// included via the dispatcher). -template -__global__ void paged_attn_prefill_split_q_kernel(PagedAttentionParams p) { - constexpr int DPT = HEAD_DIM / G; - - const int q_tile = blockIdx.x; - const int q_head = blockIdx.y; - const int req_b = blockIdx.z; - const int gpos = threadIdx.x; // 0..G-1 (d-chunk) - const int row = threadIdx.y; // 0..ROWS-1 (q row within tile) - const int q_row = q_tile * ROWS + row; - - const int seq_len = p.kv_indptr[req_b + 1] - p.kv_indptr[req_b]; - const int q_len = p.qo_indptr[req_b + 1] - p.qo_indptr[req_b]; - const int causal_off = seq_len - q_len; - const int64_t req_idx = p.req_pool_indices[req_b]; - const int kv_head = q_head / (p.q_head / p.kv_head); - - __shared__ __align__(16) bf16 sK[P_BC * HEAD_DIM]; - __shared__ __align__(16) bf16 sV[P_BC * HEAD_DIM]; - - // Q base: absolute token = qo_indptr[req_b] + q_row. - float qreg[DPT]; - if (q_row < q_len) { - int q_off = (p.qo_indptr[req_b] + q_row) * p.q_stride_l - + q_head * p.q_stride_h + gpos * DPT * p.q_stride_d; - #pragma unroll - for (int i = 0; i < DPT; i++) - qreg[i] = __bfloat162float(p.q[q_off + i * p.q_stride_d]); - } - - float m = -FLT_MAX, l = 0.0f, acc[DPT]; - #pragma unroll - for (int i = 0; i < DPT; i++) acc[i] = 0.0f; - - const int64_t pool_stride = (int64_t)p.kv_head * p.head_dim; - const int64_t head_off = (int64_t)kv_head * p.head_dim; - const int64_t rtt_stride = (int64_t)p.max_context_len; - const int mask_base = req_b * p.mask_b_stride + q_head * p.mask_h_stride - + q_row * p.mask_q_stride; - - int tiles = (seq_len + P_BC - 1) / P_BC; - int tt = G * ROWS; - int lid = row * G + gpos; - - // Each warp holds (32/G) q-rows; reduce only within this row's G lanes. - int lane_in_warp = lid & 31; - unsigned gmask = (G == 32) ? 0xFFFFFFFFu - : (((1u << G) - 1u) << (lane_in_warp & ~(G - 1))); - - for (int ti = 0; ti < tiles; ti++) { - int kv0 = ti * P_BC; - int tlen = min(P_BC, seq_len - kv0); - - // Load K/V tile into shared memory via req_to_token (request-local pos). - for (int i = lid; i < tlen * HEAD_DIM; i += tt) { - int s = i / HEAD_DIM, d_dim = i % HEAD_DIM; - int pos = kv0 + s; - int64_t slot = p.req_to_token[req_idx * rtt_stride + pos]; - int64_t off = slot * pool_stride + head_off + d_dim; - sK[i] = (slot >= 0) ? p.k_cache[off] : __float2bfloat16(0.0f); - sV[i] = (slot >= 0) ? p.v_cache[off] : __float2bfloat16(0.0f); - } - __syncthreads(); - - int lim = tlen; - if constexpr (IsCausal) { - if (q_row < q_len) { - int ep = causal_off + q_row + 1; - if (kv0 >= ep) - lim = 0; - else if (kv0 + tlen > ep) - lim = ep - kv0; - } - } - - for (int s = 0; s < lim; s++) { - bool keep = true; - if constexpr (HasMask) { - if (q_row < q_len && !p.mask[mask_base + kv0 + s]) - keep = false; - } - float w = 0.0f; - #pragma unroll - for (int i = 0; i < DPT; i++) - w += qreg[i] * __bfloat162float(sK[s * HEAD_DIM + gpos * DPT + i]); - w = group_reduce_sum(w, gmask) * p.scale; - if (!keep) w = -FLT_MAX; - - float nm = fmaxf(m, w); - float alpha = __expf(m - nm); - float beta = __expf(w - nm); - l = l * alpha + beta; - #pragma unroll - for (int i = 0; i < DPT; i++) - acc[i] = acc[i] * alpha - + __bfloat162float(sV[s * HEAD_DIM + gpos * DPT + i]) * beta; - m = nm; - } - __syncthreads(); - } - - if (q_row >= q_len) return; - float inv = (l > 1e-20f) ? (1.0f / l) : 0.0f; - int o_off = (p.qo_indptr[req_b] + q_row) * p.q_stride_l - + q_head * p.q_stride_h + gpos * DPT * p.q_stride_d; - #pragma unroll - for (int i = 0; i < DPT; i++) - p.o[o_off + i * p.q_stride_d] = __float2bfloat16(acc[i] * inv); -} diff --git a/csrc/kernels/attn_paged_prefill_split_q_mma.cuh b/csrc/kernels/attn_paged_prefill_split_q_mma.cuh deleted file mode 100644 index a20aee6..0000000 --- a/csrc/kernels/attn_paged_prefill_split_q_mma.cuh +++ /dev/null @@ -1,164 +0,0 @@ -#pragma once -#include -#include -#include "attn_common.h" -#include "attn_mma_utils.cuh" - -// SGLang-style split-Q tensor-core prefill. -// -// Reads K/V directly from a flat pool [size, kv_head, head_dim] via -// req_to_token — no gather, no temporary tensor. Supports ragged batches: -// each request has its own q_len and kv_len, addressed via qo_indptr and -// kv_indptr. -// -// Grid: (max_q_tiles, q_head, batch) — one batch element per blockIdx.z. -// Blocks beyond a request's q_len exit early after writing sentinel-free -// no-ops. This avoids the binary-search approach and guarantees every Q -// token is covered, even when q_len < BR*WARPS (e.g. decode-like prefill). -// -// Q layout: [total_q, q_head, head_dim] (3D, flattened across requests). -// O layout: same as Q. -// -// IsCausal is a compile-time bool. When true, each Q row qi (within its -// request) attends to [0, causal_offset_b + qi + 1) where -// causal_offset_b = kv_len_b - q_len_b (position of first Q token). -template -__global__ void paged_attn_prefill_split_q_mma_kernel(PagedAttentionParams p) { - const int warp = threadIdx.x / 32; - const int lane = threadIdx.x % 32; - const int gid = lane >> 2; - const int tid4 = lane & 3; - - const int q_head = blockIdx.y; - const int req_b = blockIdx.z; - const int qrow0 = (blockIdx.x * Traits::WARPS + warp) * Traits::BR; - - const int seq_len = p.kv_indptr[req_b + 1] - p.kv_indptr[req_b]; - const int q_len = p.qo_indptr[req_b + 1] - p.qo_indptr[req_b]; - const int causal_off = seq_len - q_len; - const int64_t req_idx = p.req_pool_indices[req_b]; - - // No per-warp early exit — all warps must participate in __syncthreads. - // Warps beyond q_len get zero-filled Q frags (va=vb=false) and skip output. - const int kv_head = q_head / (p.q_head / p.kv_head); - - __shared__ __align__(16) bf16 sK[Traits::STAGES * Traits::BC * Traits::LD]; - __shared__ __align__(16) bf16 sV[Traits::STAGES * Traits::BC * Traits::LD]; - - // Q base: offset by qo_indptr[req_b] to get absolute token address. - const int q_base = p.qo_indptr[req_b] * p.q_stride_l + q_head * p.q_stride_h; - const int qra = qrow0 + gid; - const int qrb = qrow0 + gid + 8; - const bool va = qra < q_len, vb = qrb < q_len; - 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[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 int64_t pool_stride = (int64_t)p.kv_head * Traits::HEAD_DIM; - const int64_t head_off = (int64_t)kv_head * Traits::HEAD_DIM; - const int64_t rtt_stride = (int64_t)p.max_context_len; - - const int tiles = (seq_len + Traits::BC - 1) / Traits::BC; - const int qr0 = qrow0 + gid; - const int qr1 = qrow0 + gid + 8; - - // Causal tile-skip (dead code when IsCausal == false) - const int max_kv = qrow0 + Traits::BR - 1 + causal_off; - const int block_max_kv = - blockIdx.x * Traits::WARPS * Traits::BR + Traits::WARPS * Traits::BR - 1 - + causal_off; - - int t_end = tiles - 1; - if constexpr (IsCausal) { - int bt = block_max_kv / Traits::BC; - if (bt < t_end) t_end = bt; - } - - // ---- Load tile lambda: SGLang addressing ---- - auto load_tile = [&](int ti, int buf) { - 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 < seq_len; - int64_t slot = valid ? p.req_to_token[req_idx * rtt_stride + kc] : 0; - valid = valid && (slot >= 0); - int64_t gmem_base = slot * pool_stride + head_off; - 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 + main loop (FA2-style double-buffer) ---- - load_tile(0, 0); - - for (int ti = 0; ti <= t_end; ti++) { - int buf = ti & 1; - - cp_async_wait_group<0>(); - __syncthreads(); - if (ti < t_end) load_tile(ti + 1, (ti + 1) & 1); - - const bf16* bK = sK + buf * Traits::BC * Traits::LD; - const bf16* bV = sV + buf * Traits::BC * Traits::LD; - int kv0 = ti * Traits::BC; - - if (!IsCausal || kv0 <= max_kv) { - float Sacc[Traits::NC8][4]; - mma_compute_scores(Qa, bK, lane, Sacc); - - #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 = IsCausal ? min(seq_len, causal_off + qr0 + 1) - : seq_len; - int maxc1 = IsCausal ? min(seq_len, causal_off + qr1 + 1) - : seq_len; - // HasMask: mask[batch, q_head, qi, kc] — kc is request-local. - mma_softmax_tile(kv0, maxc0, maxc1, - qr0, qr1, - p.mask_b_stride, p.mask_h_stride, - p.mask_q_stride, - req_b, q_head, - p.mask, - Sacc, Oacc, m0, m1, l0, l1, lane); - - mma_pv_accumulate(Sacc, bV, lane, Oacc); - } - } - - // ---- 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; - const int o_base = p.qo_indptr[req_b] * p.q_stride_l + q_head * p.q_stride_h; - #pragma unroll - for (int dn8 = 0; dn8 < Traits::DN8; dn8++) { - int d = dn8 * 8 + 2 * tid4; - if (qr0 < 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; - } - if (qr1 < 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; - } - } -} diff --git a/csrc/kernels/attn_prefill_split_q.cuh b/csrc/kernels/attn_prefill_split_q.cuh index 2969f98..2fe794b 100644 --- a/csrc/kernels/attn_prefill_split_q.cuh +++ b/csrc/kernels/attn_prefill_split_q.cuh @@ -2,13 +2,15 @@ #include #include #include "attn_common.h" +#include "attn_kv_source.cuh" 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[]. IsCausal and HasMask are // compile-time bools — the compiler eliminates dead branches. -// Templated on . +// Unified across contiguous and paged (SGLang flat-pool) K/V via KV. +// Templated on . template __device__ __forceinline__ float group_reduce_sum(float v, unsigned mask) { @@ -30,7 +32,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; @@ -41,16 +43,21 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams p) { int row = threadIdx.y; // 0..ROWS-1 int q_row = q_tile * ROWS + row; - int kv_head = q_head / (p.q_head / p.kv_head); + // Per-request dims (from KV policy — paged reads kv_indptr/qo_indptr). + const int seq_len = KV::kv_len(p, batch); + const int q_len = KV::q_len(p, batch); + const int causal_off = KV::causal_offset(p, batch); + const int kv_head = q_head / (p.q_head / p.kv_head); + const KVContext kctx = KV::template make_ctx(p, batch, kv_head); __shared__ __align__(16) bf16 sK[P_BC * HEAD_DIM]; __shared__ __align__(16) bf16 sV[P_BC * HEAD_DIM]; // Q: stride-based load [batch, q_head, q_len, head_dim] + const int q_base = KV::q_base(p, batch, q_head); float qreg[DPT]; - if (q_row < p.q_len) { - int q_off = batch * p.q_stride_b + q_head * p.q_stride_h - + q_row * p.q_stride_l + gpos * DPT * p.q_stride_d; + if (q_row < q_len) { + int q_off = q_base + q_row * p.q_stride_l + gpos * DPT * p.q_stride_d; #pragma unroll for (int i = 0; i < DPT; i++) qreg[i] = __bfloat162float(p.q[q_off + i * p.q_stride_d]); @@ -62,10 +69,8 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams p) { for (int i = 0; i < DPT; i++) acc[i] = 0.0f; - // KV: stride-based base - int kv_base = batch * p.kv_stride_b + kv_head * p.kv_stride_h; int mask_batch_base = batch * p.mask_b_stride + q_head * p.mask_h_stride; - int tiles = (p.kv_len + P_BC - 1) / P_BC; + int tiles = (seq_len + P_BC - 1) / P_BC; int tt = G * ROWS; int lid = row * G + gpos; @@ -75,23 +80,24 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams p) { for (int ti = 0; ti < tiles; ti++) { int kv0 = ti * P_BC; - int tlen = min(P_BC, p.kv_len - kv0); + int tlen = min(P_BC, seq_len - kv0); - // Load K/V into shared memory from strided global + // Load K/V into shared memory (addressing via KV policy; paged + // guards empty slots with zero-fill). for (int i = lid; i < tlen * HEAD_DIM; i += tt) { int s = i / HEAD_DIM; int d_dim = i % HEAD_DIM; - int kv_idx = kv0 + s; - int g_off = kv_base + kv_idx * p.kv_stride_l + d_dim * p.kv_stride_d; - sK[i] = p.k[g_off]; - sV[i] = p.v[g_off]; + int kc = kv0 + s; + KVAddr a = KV::kv_addr(p, kctx, kc, d_dim, true); + sK[i] = a.valid ? *reinterpret_cast(a.k) : (bf16)0.f; + sV[i] = a.valid ? *reinterpret_cast(a.v) : (bf16)0.f; } __syncthreads(); int lim = tlen; if constexpr (IsCausal) { - if (q_row < p.q_len) { - int ep = q_row + p.causal_offset + 1; + if (q_row < q_len) { + int ep = causal_off + q_row + 1; if (kv0 >= ep) lim = 0; else if (kv0 + tlen > ep) @@ -138,9 +144,8 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams p) { __syncthreads(); } - if (q_row < p.q_len) { - 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; + if (q_row < q_len) { + int o_off = q_base + q_row * p.q_stride_l + gpos * DPT * p.q_stride_d; float rl = (l > 1e-20f) ? (1.0f / l) : 0.0f; #pragma unroll for (int i = 0; i < DPT; i++) diff --git a/csrc/kernels/attn_prefill_split_q_mma.cuh b/csrc/kernels/attn_prefill_split_q_mma.cuh index e1a4e24..b3b308b 100644 --- a/csrc/kernels/attn_prefill_split_q_mma.cuh +++ b/csrc/kernels/attn_prefill_split_q_mma.cuh @@ -2,17 +2,21 @@ #include #include #include "attn_common.h" +#include "attn_kv_source.cuh" #include "attn_mma_utils.cuh" -// Tensor-core prefill flash attention (raw mma.sync PTX). +// Tensor-core prefill flash attention (raw mma.sync PTX), unified across +// contiguous and paged (SGLang flat-pool) K/V via the KV template parameter. // 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). // +// KV = ContigKV (dense [batch, kv_head, kv_len, head_dim]) or PagedKV +// (flat pool + req_to_token, ragged batches via qo_indptr/kv_indptr). // IsCausal and HasMask are compile-time bools — the compiler eliminates all // dead branches in the inner compute loop (FA2-style). // // Traits = KernelTraits. -template +template __global__ void attn_prefill_split_q_mma_kernel(AttentionParams p) { const int warp = threadIdx.x / 32; const int lane = threadIdx.x % 32; @@ -24,16 +28,22 @@ __global__ void attn_prefill_split_q_mma_kernel(AttentionParams p) { const int kv_head = q_head / (p.q_head / p.kv_head); const int qrow0 = (blockIdx.x * Traits::WARPS + warp) * Traits::BR; + // Per-request dims (from KV policy — paged reads kv_indptr/qo_indptr). + const int seq_len = KV::kv_len(p, batch); + const int q_len = KV::q_len(p, batch); + const int causal_off = KV::causal_offset(p, batch); + const KVContext kctx = KV::template make_ctx(p, batch, kv_head); + // 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. - const int q_base = batch * p.q_stride_b + q_head * p.q_stride_h; + const int q_base = KV::q_base(p, batch, q_head); const int qra = qrow0 + gid; const int qrb = qrow0 + gid + 8; - const bool va = qra < p.q_len, vb = qrb < p.q_len; + const bool va = qra < q_len, vb = qrb < q_len; 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); @@ -44,17 +54,15 @@ __global__ void attn_prefill_split_q_mma_kernel(AttentionParams p) { 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 + Traits::BC - 1) / Traits::BC; + const int tiles = (seq_len + Traits::BC - 1) / Traits::BC; const int qr0 = qrow0 + gid; const int qr1 = qrow0 + gid + 8; // Causal tile-skip bounds (dead code when IsCausal == false) - const int max_kv = qrow0 + Traits::BR - 1 + p.causal_offset; + const int max_kv = qrow0 + Traits::BR - 1 + causal_off; const int block_max_kv = blockIdx.x * Traits::WARPS * Traits::BR + Traits::WARPS * Traits::BR - 1 - + p.causal_offset; + + causal_off; int t_end = tiles - 1; if constexpr (IsCausal) { @@ -62,7 +70,7 @@ __global__ void attn_prefill_split_q_mma_kernel(AttentionParams p) { if (bt < t_end) t_end = bt; } - // ---- Load tile lambda: predicated cp.async ---- + // ---- Load tile lambda: predicated cp.async (addressing via KV policy) ---- auto load_tile = [&](int ti, int buf) { int kv0 = ti * Traits::BC; bf16* dK = sK + buf * Traits::BC * Traits::LD; @@ -72,11 +80,11 @@ __global__ void attn_prefill_split_q_mma_kernel(AttentionParams p) { 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; + bool valid = kc < seq_len; + KVAddr a = KV::kv_addr(p, kctx, kc, d, valid); 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); + cp_async_16_pred(&dK[off], a.k, a.valid); + cp_async_16_pred(&dV[off], a.v, a.valid); } cp_async_commit(); }; @@ -108,10 +116,10 @@ __global__ void attn_prefill_split_q_mma_kernel(AttentionParams p) { Sacc[n8][0] *= p.scale, Sacc[n8][1] *= p.scale, Sacc[n8][2] *= p.scale, Sacc[n8][3] *= p.scale; - 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; + int maxc0 = IsCausal ? min(seq_len, causal_off + qr0 + 1) + : seq_len; + int maxc1 = IsCausal ? min(seq_len, causal_off + qr1 + 1) + : seq_len; mma_softmax_tile(kv0, maxc0, maxc1, qr0, qr1, p.mask_b_stride, p.mask_h_stride, p.mask_q_stride, @@ -126,17 +134,17 @@ __global__ void attn_prefill_split_q_mma_kernel(AttentionParams p) { // ---- 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; - const int o_base = batch * p.q_stride_b + q_head * p.q_stride_h; + const int o_base = KV::q_base(p, batch, q_head); #pragma unroll for (int dn8 = 0; dn8 < Traits::DN8; dn8++) { int d = dn8 * 8 + 2 * tid4; - if (qr0 < p.q_len) { + if (qr0 < 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; } - if (qr1 < p.q_len) { + if (qr1 < q_len) { __nv_bfloat162 v = __floats2bfloat162_rn(Oacc[dn8][2] * rl1, Oacc[dn8][3] * rl1); *reinterpret_cast<__nv_bfloat162*>( diff --git a/csrc/tests/attn_paged_test.cu b/csrc/tests/attn_paged_test.cu index 4e1f502..8c79441 100644 --- a/csrc/tests/attn_paged_test.cu +++ b/csrc/tests/attn_paged_test.cu @@ -212,7 +212,7 @@ static int run_decode_test(int B, int Hq, int Hkv, int max_seq, B, Hq, Hkv, HEAD_DIM, max_ctx, h_o_ref); // Kernel launch - PagedAttentionParams p; + AttentionParams p; p.batch = B; p.q_head = Hq; p.kv_head = Hkv; p.head_dim = HEAD_DIM; p.total_q = B; p.q_stride_l = Hq * HEAD_DIM; p.q_stride_h = HEAD_DIM; p.q_stride_d = 1; @@ -347,7 +347,7 @@ static int run_decode_mask_test(int B, int Hq, int Hkv, int max_seq, h_mask, max_sl, B, Hq, Hkv, HEAD_DIM, max_ctx, h_o_ref); - PagedAttentionParams p; + AttentionParams p; p.batch = B; p.q_head = Hq; p.kv_head = Hkv; p.head_dim = HEAD_DIM; p.total_q = B; p.q_stride_l = Hq * HEAD_DIM; p.q_stride_h = HEAD_DIM; p.q_stride_d = 1; @@ -480,7 +480,7 @@ static int run_prefill_test(int B, int Hq, int Hkv, B, Hq, Hkv, HEAD_DIM, max_ctx, causal, h_o_ref); // Kernel launch - PagedAttentionParams p; + AttentionParams p; p.batch = B; p.q_head = Hq; p.kv_head = Hkv; p.head_dim = HEAD_DIM; p.total_q = total_q; p.q_stride_l = Hq * HEAD_DIM; p.q_stride_h = HEAD_DIM; p.q_stride_d = 1; @@ -617,7 +617,7 @@ static int run_prefill_mask_test(int Hq, int Hkv, int q_len, int seed) { h_mask, q_len, q_len, B, Hq, Hkv, HEAD_DIM, max_ctx, 0, h_o_ref); - PagedAttentionParams p; + AttentionParams p; p.batch = B; p.q_head = Hq; p.kv_head = Hkv; p.head_dim = HEAD_DIM; p.total_q = total_q; p.q_stride_l = Hq * HEAD_DIM; p.q_stride_h = HEAD_DIM; p.q_stride_d = 1; @@ -707,7 +707,7 @@ static void bench_decode(int B, int Hq, int Hkv, int seq_len) { for (int b = 0; b < B; b++) h_kvi[b + 1] = h_kvi[b] + seq_len; cudaMemcpy(d_kvi, h_kvi, sz_kvi, cudaMemcpyHostToDevice); - PagedAttentionParams p; + AttentionParams p; p.batch = B; p.q_head = Hq; p.kv_head = Hkv; p.head_dim = HEAD_DIM; p.total_q = B; p.q_stride_l = Hq * HEAD_DIM; p.q_stride_h = HEAD_DIM; p.q_stride_d = 1; @@ -784,7 +784,7 @@ static void bench_prefill(int B, int Hq, int Hkv, int q_len, int kv_len, int cau for (int b = 0; b < B; b++) h_qoi[b + 1] = h_qoi[b] + q_len; cudaMemcpy(d_qoi, h_qoi, sz_qoi, cudaMemcpyHostToDevice); - PagedAttentionParams p; + AttentionParams p; p.batch = B; p.q_head = Hq; p.kv_head = Hkv; p.head_dim = HEAD_DIM; p.total_q = total_q; p.q_stride_l = Hq * HEAD_DIM; p.q_stride_h = HEAD_DIM; p.q_stride_d = 1; diff --git a/csrc/tests/test_utils.cuh b/csrc/tests/test_utils.cuh index b998757..33181cc 100644 --- a/csrc/tests/test_utils.cuh +++ b/csrc/tests/test_utils.cuh @@ -120,7 +120,7 @@ inline void set_default_strides(P& p) { p.mask_q_stride = 0; } -// Set default Q strides for contiguous b h l d layout on PagedAttentionParams. +// Set default Q strides for a paged decode params struct. template inline void set_default_paged_strides(P& p) { p.q_stride_b = p.q_head * p.q_len * p.head_dim; diff --git a/docs/developer/cuda_kernels.md b/docs/developer/cuda_kernels.md index 5090051..9fcd08f 100644 --- a/docs/developer/cuda_kernels.md +++ b/docs/developer/cuda_kernels.md @@ -18,7 +18,10 @@ Additionally, optimized `.cuh` variants with tensor-core MMA (Matrix Multiply-Ac |---------|------|--------------| | Split-KV MMA decode | `attn_decode_split_kv_mma.cuh` | Split KV across warps + MMA (sm_80+) | | Split-Q MMA prefill | `attn_prefill_split_q_mma.cuh` | Split Q across warps + MMA (sm_80+) | -| Paged split-KV MMA decode | `attn_paged_decode_split_kv_mma.cuh` | Paged cache + split-KV + MMA | + +> The paged and non-paged paths are ONE kernel templated on a `KVSource` +> policy (`ContigKV` / `PagedKV` in `attn_kv_source.cuh`); there are no +> separate `attn_paged_*.cuh` files anymore. ### Rotary Embedding Kernel @@ -157,21 +160,18 @@ nvcc -I csrc -arch=sm_89 -O3 --use_fast_math \ csrc/ ├── CMakeLists.txt # CMake build: 5 kernel targets, torch/pybind11 linking ├── kernels/ -│ ├── attn_common.h # Shared attention params (AttentionParams, PagedAttentionParams) +│ ├── attn_common.h # Unified attention params (contig + paged modes) │ ├── attn_decode.cu # Basic decode kernel (registered) │ ├── attn_prefill.cu # Basic prefill kernel (registered) │ ├── attn_paged_decode.cu # Paged decode kernel (registered) │ ├── attn_paged_prefill.cu # Paged prefill kernel (registered) │ ├── rotary_emb.cu # Fused rotary embedding kernel (registered) -│ ├── attn_decode_split_kv.cuh # Split-KV variant -│ ├── attn_decode_split_kv_mma.cuh # Split-KV + MMA variant -│ ├── attn_prefill_split_q.cuh # Split-Q variant -│ ├── attn_prefill_split_q_mma.cuh # Split-Q + MMA variant -│ ├── attn_paged_decode_split_kv.cuh # Paged + split-KV variant -│ ├── attn_paged_decode_split_kv_mma.cuh # Paged + split-KV + MMA variant -│ ├── attn_paged_prefill_split_q.cuh # Paged + split-Q variant -│ ├── attn_paged_prefill_split_q_mma.cuh # Paged + split-Q + MMA variant -│ ├── attn_dispatchers.cuh # Kernel dispatch macros +│ ├── attn_decode_split_kv.cuh # Split-KV variant (contig + paged via KVSource) +│ ├── attn_decode_split_kv_mma.cuh # Split-KV + MMA variant (contig + paged) +│ ├── attn_prefill_split_q.cuh # Split-Q variant (contig + paged via KVSource) +│ ├── attn_prefill_split_q_mma.cuh # Split-Q + MMA variant (contig + paged) +│ ├── attn_kv_source.cuh # KVSource policies (ContigKV / PagedKV) +│ ├── attn_dispatchers.cuh # Kernel dispatch macros + KV-templated launchers │ ├── attn_entry_utils.cuh # Entry point helpers │ ├── attn_mma_utils.cuh # MMA utilities │ └── attn_warp_utils.cuh # Warp-level utilities