feat: SGLang-style paged attention kernels replace page-table path
- PagedAttentionParams uses flat KV pool + req_to_token + kv_indptr/qo_indptr instead of page_table - MMA split-KV decode and split-Q prefill kernels with indirect ragged-batch addressing - Prefill kernel accepts 4D mask (causal-aware); decode kernel supports 2D mask - CudaBackend is inference-only: kv_cache=None raises, no torch fallback - benchmark.py: required --ckpt, --backend/--compare options - Parallel build isolates build-temp/build-lib per subprocess - Standalone test covers decode/prefill with mask, 27 cases pass
This commit is contained in:
+34
-14
@@ -43,35 +43,55 @@ struct AttentionParams {
|
||||
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<typename T, typename AT = float>
|
||||
struct PagedAttentionParams {
|
||||
int batch;
|
||||
int q_head;
|
||||
int kv_head;
|
||||
int q_len;
|
||||
int kv_len;
|
||||
int head_dim;
|
||||
int num_splits;
|
||||
int use_mask;
|
||||
int causal_offset;
|
||||
int causal_offset; // -1 = non-causal; >=0 = causal (per-request offset
|
||||
// computed inside kernel from kv_indptr/qo_indptr)
|
||||
float scale;
|
||||
|
||||
int num_splits;
|
||||
int page_size;
|
||||
int max_pages;
|
||||
// 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 strides (layout-agnostic)
|
||||
int q_stride_b, q_stride_h, q_stride_l, q_stride_d;
|
||||
// Q: [total_q, q_head, head_dim]
|
||||
const T* __restrict__ q;
|
||||
|
||||
// Mask strides (2D, 3D, or 4D)
|
||||
// Flat KV pool: [size, kv_head, head_dim]
|
||||
const T* __restrict__ k_cache;
|
||||
const T* __restrict__ v_cache;
|
||||
|
||||
// Indexing
|
||||
const int64_t* __restrict__ req_to_token; // [num_reqs, max_context_len]
|
||||
const int64_t* __restrict__ req_pool_indices; // [batch]
|
||||
const int* __restrict__ kv_indptr; // [batch+1]
|
||||
const int* __restrict__ qo_indptr; // [batch+1] or nullptr (decode)
|
||||
int max_context_len; // req_to_token stride (dim 1)
|
||||
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 T* __restrict__ q;
|
||||
const T* __restrict__ k_cache;
|
||||
const T* __restrict__ v_cache;
|
||||
const bool* __restrict__ mask;
|
||||
const int64_t* __restrict__ page_table;
|
||||
|
||||
T* __restrict__ o;
|
||||
AT* __restrict__ o_part;
|
||||
|
||||
Reference in New Issue
Block a user