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)
This commit is contained in:
2026-08-05 14:06:13 +08:00
parent 6dffb0305a
commit 2667b8116d
17 changed files with 415 additions and 890 deletions
+29 -21
View File
@@ -2,17 +2,21 @@
#include <cfloat>
#include <cuda_bf16.h>
#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<HEAD_DIM, BC, WARPS=4, STAGES=2>.
template <typename Traits, bool IsCausal, bool HasMask>
template <typename Traits, typename KV, bool IsCausal, bool HasMask>
__global__ void attn_prefill_split_q_mma_kernel(AttentionParams<bf16> 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<bf16> 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<Traits::HEAD_DIM>(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<Traits::KD>(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<bf16> 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<bf16> 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<bf16> 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<bf16> 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<Traits, HasMask>(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<bf16> 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*>(