refactor: unify attention param field names

- rename q_stride_* to q_*_stride to match mask stride convention
- rename mask_q_stride to mask_l_stride for consistent l-dim naming
- merge k/v and k_cache/v_cache into k_ptr/v_ptr; rename q to q_ptr
- KVSource policy selects contiguous vs paged mode at compile time
This commit is contained in:
2026-08-09 20:23:58 +08:00
parent d565d44c43
commit a5a3cc1fc2
11 changed files with 124 additions and 118 deletions
+25 -19
View File
@@ -23,22 +23,35 @@ struct AttentionParams {
int q_head;
int kv_head;
int head_dim;
int use_mask;
int causal_offset; // -1 = non-causal; >=0 = absolute position of first Q token
int num_splits;
float scale;
// Q strides (element offsets for each dim — layout-agnostic)
int q_stride_b, q_stride_h, q_stride_l, q_stride_d;
// -1 = non-causal; >=0 = absolute position of first Q token
int causal_offset;
int use_mask;
int num_splits;
// 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)
// Q strides
int q_b_stride;
int q_h_stride;
int q_l_stride;
int q_d_stride;
// K/V strides
int kv_b_stride;
int kv_h_stride;
int kv_l_stride;
int kv_d_stride;
// Mask strides
int mask_b_stride;
int mask_h_stride;
int mask_l_stride;
const T* __restrict__ q_ptr;
const T* __restrict__ k_ptr;
const T* __restrict__ v_ptr;
const bool* __restrict__ mask;
const T* __restrict__ q;
T* __restrict__ o;
AT* __restrict__ o_part;
AT* __restrict__ ml_part;
@@ -46,13 +59,6 @@ struct AttentionParams {
// ---- 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;
// ---- paged (SGLang flat pool) mode ----
const T* __restrict__ k_cache;
const T* __restrict__ v_cache;
// Indexing
const int64_t* __restrict__ req_to_token; // [num_reqs, max_context_len]
+3 -3
View File
@@ -27,9 +27,9 @@ __global__ void attn_decode_split_kv_kernel(AttentionParams<bf16> p) {
// Q: [batch, q_head, q_len=1, head_dim] — stride-based
float q_reg[8];
int q_off = KV::q_decode_base(p, batch, q_head)
+ lane * hd_per_thread * p.q_stride_d;
+ lane * hd_per_thread * p.q_d_stride;
for (int i = 0; i < hd_per_thread; i++)
q_reg[i] = __bfloat162float(p.q[q_off + i * p.q_stride_d]);
q_reg[i] = __bfloat162float(p.q_ptr[q_off + i * p.q_d_stride]);
int mask_base = batch * p.mask_b_stride + q_head * p.mask_h_stride;
@@ -138,6 +138,6 @@ __global__ void attn_decode_combine_kernel(AttentionParams<bf16> p) {
}
float inv = (l > 1e-20f) ? (1.0f / l) : 0.0f;
int o_off = KV::q_decode_base(p, batch, q_head) + d * p.q_stride_d;
int o_off = KV::q_decode_base(p, batch, q_head) + d * p.q_d_stride;
p.o[o_off] = __float2bfloat16(acc * inv);
}
+2 -2
View File
@@ -48,7 +48,7 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
const int qrb = gid + 8;
const bool va = qra < G, vb = qrb < G;
unsigned Qa[Traits::KD][4];
load_q_mma_frags<Traits::KD>(p.q + q_base, p.q_stride_h, p.q_stride_d,
load_q_mma_frags<Traits::KD>(p.q_ptr + q_base, p.q_h_stride, p.q_d_stride,
qra, qrb, va, vb, tid4, Qa);
float Oacc[Traits::DN8][4];
@@ -107,7 +107,7 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
int maxc = IsCausal ? KV::decode_attend_len(p, batch) : seq_len;
mma_softmax_tile<Traits, HasMask>(kv0, maxc, maxc,
0, 0,
p.mask_b_stride, p.mask_h_stride, p.mask_q_stride,
p.mask_b_stride, p.mask_h_stride, p.mask_l_stride,
batch, q_head0 + gid, q_head0 + gid + 8,
p.mask,
va, vb,
+32 -32
View File
@@ -42,10 +42,10 @@ inline void extract_q_dims_and_strides(torch::Tensor& q, int64_t layout, P& p) {
p.q_head = (int)q.size(1);
p.q_len = (int)q.size(2);
p.head_dim = (int)q.size(3);
p.q_stride_b = (int)q.stride(0);
p.q_stride_h = (int)q.stride(1);
p.q_stride_l = (int)q.stride(2);
p.q_stride_d = (int)q.stride(3);
p.q_b_stride = (int)q.stride(0);
p.q_h_stride = (int)q.stride(1);
p.q_l_stride = (int)q.stride(2);
p.q_d_stride = (int)q.stride(3);
}
// ---- Shared mask packing ----
@@ -63,17 +63,17 @@ inline void pack_mask(const c10::optional<torch::Tensor>& mask, P& p) {
if (m.dim() == 2) {
p.mask_b_stride = (int)m.stride(0);
p.mask_h_stride = 0;
p.mask_q_stride = 0;
p.mask_l_stride = 0;
} else if (m.dim() == 3) {
TORCH_CHECK(m.size(1) == 1 || m.size(1) == p.q_len, "mask q_len mismatch");
p.mask_b_stride = (int)m.stride(0);
p.mask_h_stride = 0;
p.mask_q_stride = (m.size(1) == 1) ? 0 : (int)m.stride(1);
p.mask_l_stride = (m.size(1) == 1) ? 0 : (int)m.stride(1);
} else if (m.dim() == 4) {
TORCH_CHECK(m.size(2) == 1 || m.size(2) == p.q_len, "mask q_len mismatch");
p.mask_b_stride = (int)m.stride(0);
p.mask_h_stride = (m.size(1) == 1) ? 0 : (int)m.stride(1);
p.mask_q_stride = (m.size(2) == 1) ? 0 : (int)m.stride(2);
p.mask_l_stride = (m.size(2) == 1) ? 0 : (int)m.stride(2);
} else {
TORCH_CHECK(false, "mask must be 2D, 3D, or 4D");
}
@@ -82,7 +82,7 @@ inline void pack_mask(const c10::optional<torch::Tensor>& mask, P& p) {
p.mask = nullptr;
p.mask_b_stride = 0;
p.mask_h_stride = 0;
p.mask_q_stride = 0;
p.mask_l_stride = 0;
}
}
@@ -118,18 +118,18 @@ inline void attn_pack_params(
TORCH_CHECK(q.stride(3) == 1 && k.stride(3) == 1 && v.stride(3) == 1,
"Q/K/V head_dim must be contiguous");
p.kv_stride_b = (int)k.stride(0);
p.kv_stride_h = (int)k.stride(1);
p.kv_stride_l = (int)k.stride(2);
p.kv_stride_d = (int)k.stride(3);
p.kv_b_stride = (int)k.stride(0);
p.kv_h_stride = (int)k.stride(1);
p.kv_l_stride = (int)k.stride(2);
p.kv_d_stride = (int)k.stride(3);
p.causal_offset = (int)causal_offset;
p.use_mask = mask.has_value() ? 1 : 0;
p.scale = (scale > 0.0) ? (float)scale : 1.0f / sqrtf((float)p.head_dim);
p.q = (const T*)q.data_ptr();
p.k = (const T*)k.data_ptr();
p.v = (const T*)v.data_ptr();
p.q_ptr = (const T*)q.data_ptr();
p.k_ptr = (const T*)k.data_ptr();
p.v_ptr = (const T*)v.data_ptr();
p.o = nullptr;
p.o_part = nullptr;
p.ml_part = nullptr;
@@ -178,13 +178,13 @@ inline void attn_pack_paged_decode_params(
TORCH_CHECK(p.head_dim % 32 == 0, "head_dim must be multiple of 32");
TORCH_CHECK(p.q_head % p.kv_head == 0, "q_head must be divisible by kv_head");
p.q_stride_l = (int)q.stride(0);
p.q_stride_h = (int)q.stride(1);
p.q_stride_d = (int)q.stride(2);
p.q_l_stride = (int)q.stride(0);
p.q_h_stride = (int)q.stride(1);
p.q_d_stride = (int)q.stride(2);
p.k_cache = (const T*)k_cache.data_ptr();
p.v_cache = (const T*)v_cache.data_ptr();
p.q = (const T*)q.data_ptr();
p.k_ptr = (const T*)k_cache.data_ptr();
p.v_ptr = (const T*)v_cache.data_ptr();
p.q_ptr = (const T*)q.data_ptr();
p.req_to_token = req_to_token.data_ptr<int64_t>();
p.req_pool_indices = req_pool_indices.data_ptr<int64_t>();
p.kv_indptr = kv_indptr.data_ptr<int>();
@@ -204,13 +204,13 @@ inline void attn_pack_paged_decode_params(
TORCH_CHECK(m.size(0) == p.batch, "mask batch mismatch");
p.mask_b_stride = (int)m.stride(0);
p.mask_h_stride = 0;
p.mask_q_stride = 0;
p.mask_l_stride = 0;
p.mask = m.data_ptr<bool>();
} else {
p.mask = nullptr;
p.mask_b_stride = 0;
p.mask_h_stride = 0;
p.mask_q_stride = 0;
p.mask_l_stride = 0;
}
p.o = nullptr;
@@ -264,13 +264,13 @@ inline void attn_pack_paged_prefill_params(
TORCH_CHECK(kv_indptr.size(0) == p.batch + 1, "kv_indptr must be [batch+1]");
TORCH_CHECK(qo_indptr.size(0) == p.batch + 1, "qo_indptr must be [batch+1]");
p.q_stride_l = (int)q.stride(0);
p.q_stride_h = (int)q.stride(1);
p.q_stride_d = (int)q.stride(2);
p.q_l_stride = (int)q.stride(0);
p.q_h_stride = (int)q.stride(1);
p.q_d_stride = (int)q.stride(2);
p.k_cache = (const T*)k_cache.data_ptr();
p.v_cache = (const T*)v_cache.data_ptr();
p.q = (const T*)q.data_ptr();
p.k_ptr = (const T*)k_cache.data_ptr();
p.v_ptr = (const T*)v_cache.data_ptr();
p.q_ptr = (const T*)q.data_ptr();
p.req_to_token = req_to_token.data_ptr<int64_t>();
p.req_pool_indices = req_pool_indices.data_ptr<int64_t>();
p.kv_indptr = kv_indptr.data_ptr<int>();
@@ -292,14 +292,14 @@ inline void attn_pack_paged_prefill_params(
TORCH_CHECK(m.size(1) <= p.max_context_len, "mask kv_len mismatch");
p.mask_b_stride = (int)m.stride(0);
p.mask_h_stride = 0;
p.mask_q_stride = 0;
p.mask_l_stride = 0;
} else if (m.dim() == 4) {
TORCH_CHECK(m.size(1) == 1 || m.size(1) == p.q_head, "mask head mismatch");
TORCH_CHECK(m.size(2) == 1 || m.size(2) == p.max_q_len, "mask q_len mismatch");
TORCH_CHECK(m.size(3) <= p.max_context_len, "mask kv_len mismatch");
p.mask_b_stride = (int)m.stride(0);
p.mask_h_stride = (m.size(1) == 1) ? 0 : (int)m.stride(1);
p.mask_q_stride = (m.size(2) == 1) ? 0 : (int)m.stride(2);
p.mask_l_stride = (m.size(2) == 1) ? 0 : (int)m.stride(2);
} else {
TORCH_CHECK(false, "mask must be 2D or 4D");
}
@@ -308,7 +308,7 @@ inline void attn_pack_paged_prefill_params(
p.mask = nullptr;
p.mask_b_stride = 0;
p.mask_h_stride = 0;
p.mask_q_stride = 0;
p.mask_l_stride = 0;
}
p.scale = (scale > 0.0) ? (float)scale : 1.0f / sqrtf((float)p.head_dim);
+12 -12
View File
@@ -13,11 +13,11 @@
//
// 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.
// q_b_stride, 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.
// max_context_len, q_l_stride.
//
// Addressing state that is constant across a whole kernel invocation for one
// (batch, kv_head) pair is captured once by make_ctx<HEAD_DIM>() and passed
@@ -32,7 +32,7 @@ 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
int kv_base; // contig: batch*kv_b_stride + kv_head*kv_h_stride
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
@@ -64,15 +64,15 @@ struct ContigKV {
return p.kv_len;
}
// prefill: element offset of the request's Q rows (kernel adds qrow*q_stride_l)
// prefill: element offset of the request's Q rows (kernel adds qrow*q_l_stride)
HOST_DEV_FORCEINLINE int q_base(
const AttentionParams<bf16>& p, int batch, int q_head) {
return batch * p.q_stride_b + q_head * p.q_stride_h;
return batch * p.q_b_stride + q_head * p.q_h_stride;
}
// decode: same offset (q_len == 1, so there is no row stride component)
HOST_DEV_FORCEINLINE int q_decode_base(
const AttentionParams<bf16>& p, int batch, int q_head) {
return batch * p.q_stride_b + q_head * p.q_stride_h;
return batch * p.q_b_stride + q_head * p.q_h_stride;
}
HOST_DEV_FORCEINLINE int kv_len(const AttentionParams<bf16>& p, int batch) {
@@ -93,13 +93,13 @@ struct ContigKV {
HOST_DEV_FORCEINLINE KVContext make_ctx(
const AttentionParams<bf16>& p, int batch, int kv_head) {
KVContext c = {};
c.kv_base = batch * p.kv_stride_b + kv_head * p.kv_stride_h;
c.kv_base = batch * p.kv_b_stride + kv_head * p.kv_h_stride;
return c;
}
HOST_DEV_FORCEINLINE KVAddr kv_addr(
const AttentionParams<bf16>& 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};
const int g_off = c.kv_base + kc * p.kv_l_stride + d * p.kv_d_stride;
return {&p.k_ptr[g_off], &p.v_ptr[g_off], valid};
}
};
@@ -117,12 +117,12 @@ struct PagedKV {
// prefill: Q rows start at qo_indptr[batch] (ragged batch base)
HOST_DEV_FORCEINLINE int q_base(
const AttentionParams<bf16>& p, int batch, int q_head) {
return p.qo_indptr[batch] * p.q_stride_l + q_head * p.q_stride_h;
return p.qo_indptr[batch] * p.q_l_stride + q_head * p.q_h_stride;
}
// decode: Q is [batch, q_head, head_dim], so batch is the outer row
HOST_DEV_FORCEINLINE int q_decode_base(
const AttentionParams<bf16>& p, int batch, int q_head) {
return batch * p.q_stride_l + q_head * p.q_stride_h;
return batch * p.q_l_stride + q_head * p.q_h_stride;
}
HOST_DEV_FORCEINLINE int kv_len(const AttentionParams<bf16>& p, int batch) {
@@ -154,6 +154,6 @@ struct PagedKV {
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};
return {&p.k_ptr[gmem_off], &p.v_ptr[gmem_off], ok};
}
};
+5 -5
View File
@@ -133,8 +133,8 @@ __device__ __forceinline__ void cp_async_wait_group() {
// ---------------------------------------------------------------------------
// Q-load: load query rows directly from global memory into mma A-operand
// register layout. One call replaces ~15 duplicated lines in each MMA kernel.
// stride_row is p.q_stride_h for decode (q_len=1, G heads) or
// p.q_stride_l for prefill (multi-q rows).
// stride_row is p.q_h_stride for decode (q_len=1, G heads) or
// p.q_l_stride for prefill (multi-q rows).
// ---------------------------------------------------------------------------
template <int KD>
__device__ inline void load_q_mma_frags(
@@ -198,7 +198,7 @@ __device__ inline void mma_softmax_tile(
int kv0,
int maxc0, int maxc1,
int qrow0, int qrow1,
int mask_b_stride, int mask_h_stride, int mask_q_stride,
int mask_b_stride, int mask_h_stride, int mask_l_stride,
int mask_batch, int mask_head0, int mask_head1,
const bool* __restrict__ mask,
bool valid0, bool valid1,
@@ -211,8 +211,8 @@ __device__ inline void mma_softmax_tile(
int tid4 = lane & 3;
float rmax0 = -FLT_MAX, rmax1 = -FLT_MAX;
int mask_base0 = mask_batch * mask_b_stride + mask_head0 * mask_h_stride + qrow0 * mask_q_stride;
int mask_base1 = mask_batch * mask_b_stride + mask_head1 * mask_h_stride + qrow1 * mask_q_stride;
int mask_base0 = mask_batch * mask_b_stride + mask_head0 * mask_h_stride + qrow0 * mask_l_stride;
int mask_base1 = mask_batch * mask_b_stride + mask_head1 * mask_h_stride + qrow1 * mask_l_stride;
#pragma unroll
for (int n8 = 0; n8 < Traits::NC8; n8++) {
int cc = kv0 + n8 * 8 + 2 * tid4;
+5 -5
View File
@@ -57,10 +57,10 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
const int q_base = KV::q_base(p, batch, q_head);
float qreg[DPT];
if (q_row < q_len) {
int q_off = q_base + q_row * p.q_stride_l + gpos * DPT * p.q_stride_d;
int q_off = q_base + q_row * p.q_l_stride + gpos * DPT * p.q_d_stride;
#pragma unroll
for (int i = 0; i < DPT; i++)
qreg[i] = __bfloat162float(p.q[q_off + i * p.q_stride_d]);
qreg[i] = __bfloat162float(p.q_ptr[q_off + i * p.q_d_stride]);
}
float m = -FLT_MAX, l = 0.0f;
@@ -105,7 +105,7 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
}
}
int mask_row_base = mask_batch_base + q_row * p.mask_q_stride;
int mask_row_base = mask_batch_base + q_row * p.mask_l_stride;
for (int s = 0; s < lim; s++) {
const bf16* kr = sK + s * HEAD_DIM + gpos * DPT;
float part = 0.0f;
@@ -145,10 +145,10 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
}
if (q_row < q_len) {
int o_off = q_base + q_row * p.q_stride_l + gpos * DPT * p.q_stride_d;
int o_off = q_base + q_row * p.q_l_stride + gpos * DPT * p.q_d_stride;
float rl = (l > 1e-20f) ? (1.0f / l) : 0.0f;
#pragma unroll
for (int i = 0; i < DPT; i++)
p.o[o_off + i * p.q_stride_d] = __float2bfloat16(acc[i] * rl);
p.o[o_off + i * p.q_d_stride] = __float2bfloat16(acc[i] * rl);
}
}
+4 -4
View File
@@ -45,7 +45,7 @@ __global__ void attn_prefill_split_q_mma_kernel(AttentionParams<bf16> p) {
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<Traits::KD>(p.q + q_base, p.q_stride_l, p.q_stride_d,
load_q_mma_frags<Traits::KD>(p.q_ptr + q_base, p.q_l_stride, p.q_d_stride,
qra, qrb, va, vb, tid4, Qa);
float Oacc[Traits::DN8][4];
@@ -122,7 +122,7 @@ __global__ void attn_prefill_split_q_mma_kernel(AttentionParams<bf16> p) {
: seq_len;
mma_softmax_tile<Traits, HasMask>(kv0, maxc0, maxc1,
qr0, qr1,
p.mask_b_stride, p.mask_h_stride, p.mask_q_stride,
p.mask_b_stride, p.mask_h_stride, p.mask_l_stride,
batch, q_head, q_head,
p.mask,
va, vb,
@@ -143,13 +143,13 @@ __global__ void attn_prefill_split_q_mma_kernel(AttentionParams<bf16> p) {
__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;
&p.o[o_base + qr0 * p.q_l_stride + d * p.q_d_stride]) = 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;
&p.o[o_base + qr1 * p.q_l_stride + d * p.q_d_stride]) = v;
}
}
}