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:
+25
-19
@@ -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]
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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);
|
||||
|
||||
|
||||
@@ -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};
|
||||
}
|
||||
};
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -68,7 +68,7 @@ static void cpu_paged_prefill_ref(
|
||||
const float* Q, const float* K_pool, const float* V_pool,
|
||||
const int64_t* req_to_token, const int64_t* req_pool_indices,
|
||||
const int* kv_indptr, const int* qo_indptr,
|
||||
const bool* mask, int mask_q_stride, int mask_kv_stride,
|
||||
const bool* mask, int mask_l_stride, int mask_kv_stride,
|
||||
int B, int Hq, int Hkv, int D, int max_ctx_len, int causal,
|
||||
float* O)
|
||||
{
|
||||
@@ -87,7 +87,7 @@ static void cpu_paged_prefill_ref(
|
||||
float accum[256] = {0.0f};
|
||||
int lim = causal ? min(seq_len, causal_off + qi + 1) : seq_len;
|
||||
for (int kj = 0; kj < lim; kj++) {
|
||||
if (mask && !mask[b * mask_q_stride * mask_kv_stride
|
||||
if (mask && !mask[b * mask_l_stride * mask_kv_stride
|
||||
+ qi * mask_kv_stride + kj]) continue;
|
||||
int64_t slot = req_to_token[req_idx * max_ctx_len + kj];
|
||||
float dot = 0.0f;
|
||||
@@ -219,13 +219,13 @@ static int run_decode_test(int B, int Hq, int Hkv, int max_seq,
|
||||
AttentionParams<bf16> 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;
|
||||
p.q_l_stride = Hq * HEAD_DIM; p.q_h_stride = HEAD_DIM; p.q_d_stride = 1;
|
||||
p.max_context_len = max_ctx; p.max_seq_len = max_sl;
|
||||
p.causal_offset = causal ? 0 : -1; p.use_mask = 0;
|
||||
p.mask = nullptr; p.mask_b_stride = 0;
|
||||
p.mask_h_stride = 0; p.mask_q_stride = 0;
|
||||
p.mask_h_stride = 0; p.mask_l_stride = 0;
|
||||
p.scale = 1.0f / sqrtf((float)HEAD_DIM);
|
||||
p.q = d_q; p.k_cache = d_k_pool; p.v_cache = d_v_pool;
|
||||
p.q_ptr = d_q; p.k_ptr = d_k_pool; p.v_ptr = d_v_pool;
|
||||
p.req_to_token = d_rtt; p.req_pool_indices = d_rpi;
|
||||
p.kv_indptr = d_kvi; p.qo_indptr = nullptr;
|
||||
p.o = d_o; p.o_part = d_op; p.ml_part = d_ml;
|
||||
@@ -354,13 +354,13 @@ static int run_decode_mask_test(int B, int Hq, int Hkv, int max_seq,
|
||||
AttentionParams<bf16> 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;
|
||||
p.q_l_stride = Hq * HEAD_DIM; p.q_h_stride = HEAD_DIM; p.q_d_stride = 1;
|
||||
p.max_context_len = max_ctx; p.max_seq_len = max_sl;
|
||||
p.causal_offset = -1; p.use_mask = 1;
|
||||
p.mask = d_mask; p.mask_b_stride = max_sl;
|
||||
p.mask_h_stride = 0; p.mask_q_stride = 0;
|
||||
p.mask_h_stride = 0; p.mask_l_stride = 0;
|
||||
p.scale = 1.0f / sqrtf((float)HEAD_DIM);
|
||||
p.q = d_q; p.k_cache = d_k_pool; p.v_cache = d_v_pool;
|
||||
p.q_ptr = d_q; p.k_ptr = d_k_pool; p.v_ptr = d_v_pool;
|
||||
p.req_to_token = d_rtt; p.req_pool_indices = d_rpi;
|
||||
p.kv_indptr = d_kvi; p.qo_indptr = nullptr;
|
||||
p.o = d_o; p.o_part = d_op; p.ml_part = d_ml;
|
||||
@@ -487,16 +487,16 @@ static int run_prefill_test(int B, int Hq, int Hkv,
|
||||
AttentionParams<bf16> 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;
|
||||
p.q_l_stride = Hq * HEAD_DIM; p.q_h_stride = HEAD_DIM; p.q_d_stride = 1;
|
||||
p.max_context_len = max_ctx; p.max_seq_len = max_sl;
|
||||
int max_ql = 0;
|
||||
for (int b = 0; b < B; b++) max_ql = max(max_ql, q_lens[b]);
|
||||
p.max_q_len = max_ql;
|
||||
p.causal_offset = causal ? 0 : -1; p.use_mask = 0;
|
||||
p.mask = nullptr; p.mask_b_stride = 0;
|
||||
p.mask_h_stride = 0; p.mask_q_stride = 0;
|
||||
p.mask_h_stride = 0; p.mask_l_stride = 0;
|
||||
p.scale = 1.0f / sqrtf((float)HEAD_DIM);
|
||||
p.q = d_q; p.k_cache = d_k_pool; p.v_cache = d_v_pool;
|
||||
p.q_ptr = d_q; p.k_ptr = d_k_pool; p.v_ptr = d_v_pool;
|
||||
p.req_to_token = d_rtt; p.req_pool_indices = d_rpi;
|
||||
p.kv_indptr = d_kvi; p.qo_indptr = d_qoi;
|
||||
p.o = d_o; p.o_part = nullptr; p.ml_part = nullptr;
|
||||
@@ -624,14 +624,14 @@ static int run_prefill_mask_test(int Hq, int Hkv, int q_len, int seed) {
|
||||
AttentionParams<bf16> 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;
|
||||
p.q_l_stride = Hq * HEAD_DIM; p.q_h_stride = HEAD_DIM; p.q_d_stride = 1;
|
||||
p.max_context_len = max_ctx; p.max_seq_len = q_len;
|
||||
p.max_q_len = q_len;
|
||||
p.causal_offset = -1; p.use_mask = 1;
|
||||
p.mask = d_mask; p.mask_b_stride = q_len * q_len;
|
||||
p.mask_h_stride = 0; p.mask_q_stride = q_len;
|
||||
p.mask_h_stride = 0; p.mask_l_stride = q_len;
|
||||
p.scale = 1.0f / sqrtf((float)HEAD_DIM);
|
||||
p.q = d_q; p.k_cache = d_k_pool; p.v_cache = d_v_pool;
|
||||
p.q_ptr = d_q; p.k_ptr = d_k_pool; p.v_ptr = d_v_pool;
|
||||
p.req_to_token = d_rtt; p.req_pool_indices = d_rpi;
|
||||
p.kv_indptr = d_kvi; p.qo_indptr = d_qoi;
|
||||
p.o = d_o; p.o_part = nullptr; p.ml_part = nullptr;
|
||||
@@ -714,12 +714,12 @@ static void bench_decode(int B, int Hq, int Hkv, int seq_len) {
|
||||
AttentionParams<bf16> 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;
|
||||
p.q_l_stride = Hq * HEAD_DIM; p.q_h_stride = HEAD_DIM; p.q_d_stride = 1;
|
||||
p.max_context_len = max_ctx; p.max_seq_len = seq_len;
|
||||
p.causal_offset = 0; p.use_mask = 0;
|
||||
p.mask = nullptr; p.mask_b_stride = 0;
|
||||
p.scale = 1.0f / sqrtf((float)HEAD_DIM);
|
||||
p.q = d_q; p.k_cache = d_k_pool; p.v_cache = d_v_pool;
|
||||
p.q_ptr = d_q; p.k_ptr = d_k_pool; p.v_ptr = d_v_pool;
|
||||
p.req_to_token = d_rtt; p.req_pool_indices = d_rpi;
|
||||
p.kv_indptr = d_kvi; p.qo_indptr = nullptr;
|
||||
p.o = d_o; p.o_part = d_op; p.ml_part = d_ml;
|
||||
@@ -791,13 +791,13 @@ static void bench_prefill(int B, int Hq, int Hkv, int q_len, int kv_len, int cau
|
||||
AttentionParams<bf16> 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;
|
||||
p.q_l_stride = Hq * HEAD_DIM; p.q_h_stride = HEAD_DIM; p.q_d_stride = 1;
|
||||
p.max_context_len = max_ctx; p.max_seq_len = kv_len;
|
||||
p.total_q = total_q; p.max_q_len = q_len;
|
||||
p.causal_offset = causal ? 0 : -1; p.use_mask = 0;
|
||||
p.mask = nullptr; p.mask_b_stride = 0;
|
||||
p.scale = 1.0f / sqrtf((float)HEAD_DIM);
|
||||
p.q = d_q; p.k_cache = d_k_pool; p.v_cache = d_v_pool;
|
||||
p.q_ptr = d_q; p.k_ptr = d_k_pool; p.v_ptr = d_v_pool;
|
||||
p.req_to_token = d_rtt; p.req_pool_indices = d_rpi;
|
||||
p.kv_indptr = d_kvi; p.qo_indptr = d_qoi;
|
||||
p.o = d_o; p.o_part = nullptr; p.ml_part = nullptr;
|
||||
|
||||
@@ -61,7 +61,7 @@ static int run_decode_test(int B, int Hq, int Hk, int sl, int D, int causal) {
|
||||
p.use_mask=0; p.causal_offset=causal?0:-1;
|
||||
p.scale=1.0f/sqrtf((float)D);
|
||||
set_default_strides(p);
|
||||
p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO;
|
||||
p.q_ptr=dQ; p.k_ptr=dK; p.v_ptr=dV; p.mask=nullptr; p.o=dO;
|
||||
|
||||
DecodeScratch sc;
|
||||
setup_scratch(p, sc);
|
||||
@@ -141,7 +141,7 @@ static void bench_decode() {
|
||||
p.head_dim = D; p.use_mask = 0; p.causal_offset = -1;
|
||||
p.scale = 1.0f / sqrtf((float)D);
|
||||
set_default_strides(p);
|
||||
p.q = dQ; p.k = dK; p.v = dV; p.mask = nullptr; p.o = dO;
|
||||
p.q_ptr = dQ; p.k_ptr = dK; p.v_ptr = dV; p.mask = nullptr; p.o = dO;
|
||||
|
||||
DecodeScratch sc;
|
||||
setup_scratch(p, sc);
|
||||
@@ -188,7 +188,7 @@ static int run_prefill_test(int B, int Hq, int Hk, int ql, int kl, int D, int ca
|
||||
p.use_mask=0; p.causal_offset=causal?0:-1;
|
||||
set_default_strides(p);
|
||||
p.scale=1.0f/sqrtf((float)D);
|
||||
p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO;
|
||||
p.q_ptr=dQ; p.k_ptr=dK; p.v_ptr=dV; p.mask=nullptr; p.o=dO;
|
||||
|
||||
double t0=now_ms();
|
||||
dispatch_by_head_dim(D, PrefillDispatch{p});
|
||||
@@ -262,7 +262,7 @@ static void bench_prefill() {
|
||||
p.use_mask=0; p.causal_offset=causal?0:-1;
|
||||
set_default_strides(p);
|
||||
p.scale=1.0f/sqrtf((float)D);
|
||||
p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO;
|
||||
p.q_ptr=dQ; p.k_ptr=dK; p.v_ptr=dV; p.mask=nullptr; p.o=dO;
|
||||
|
||||
auto launch = [&]() { dispatch_by_head_dim(D, PrefillDispatch{p}); };
|
||||
for (int i=0;i<WARMUP;i++) launch();
|
||||
|
||||
+14
-14
@@ -107,29 +107,29 @@ void dispatch_by_head_dim(int head_dim, Fn&& fn) {
|
||||
// Set default strides for contiguous b h l d layout on AttentionParams.
|
||||
template<typename P>
|
||||
inline void set_default_strides(P& p) {
|
||||
p.q_stride_b = p.q_head * p.q_len * p.head_dim;
|
||||
p.q_stride_h = p.q_len * p.head_dim;
|
||||
p.q_stride_l = p.head_dim;
|
||||
p.q_stride_d = 1;
|
||||
p.kv_stride_b = p.kv_head * p.kv_len * p.head_dim;
|
||||
p.kv_stride_h = p.kv_len * p.head_dim;
|
||||
p.kv_stride_l = p.head_dim;
|
||||
p.kv_stride_d = 1;
|
||||
p.q_b_stride = p.q_head * p.q_len * p.head_dim;
|
||||
p.q_h_stride = p.q_len * p.head_dim;
|
||||
p.q_l_stride = p.head_dim;
|
||||
p.q_d_stride = 1;
|
||||
p.kv_b_stride = p.kv_head * p.kv_len * p.head_dim;
|
||||
p.kv_h_stride = p.kv_len * p.head_dim;
|
||||
p.kv_l_stride = p.head_dim;
|
||||
p.kv_d_stride = 1;
|
||||
p.mask_b_stride = p.kv_len;
|
||||
p.mask_h_stride = 0;
|
||||
p.mask_q_stride = 0;
|
||||
p.mask_l_stride = 0;
|
||||
}
|
||||
|
||||
// Set default Q strides for a paged decode params struct.
|
||||
template<typename P>
|
||||
inline void set_default_paged_strides(P& p) {
|
||||
p.q_stride_b = p.q_head * p.q_len * p.head_dim;
|
||||
p.q_stride_h = p.q_len * p.head_dim;
|
||||
p.q_stride_l = p.head_dim;
|
||||
p.q_stride_d = 1;
|
||||
p.q_b_stride = p.q_head * p.q_len * p.head_dim;
|
||||
p.q_h_stride = p.q_len * p.head_dim;
|
||||
p.q_l_stride = p.head_dim;
|
||||
p.q_d_stride = 1;
|
||||
p.mask_b_stride = p.kv_len;
|
||||
p.mask_h_stride = 0;
|
||||
p.mask_q_stride = 0;
|
||||
p.mask_l_stride = 0;
|
||||
}
|
||||
|
||||
// Generic CPU reference for multi-query / grouped-query attention.
|
||||
|
||||
Reference in New Issue
Block a user