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:
2026-08-01 15:41:25 +08:00
parent 9960f79920
commit 41dcf0feb9
17 changed files with 1683 additions and 535 deletions
+34 -14
View File
@@ -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;
+35 -7
View File
@@ -12,6 +12,7 @@
#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
// Cached SM count — cudaDeviceGetAttribute is a host-side call that was
@@ -145,18 +146,18 @@ static inline void dispatch_decode(AttentionParams<bf16>& p) {
}
// ======================================================================
// Paged Decode
// Paged Decode (SGLang-style: flat pool + req_to_token + kv_indptr)
// ======================================================================
#ifndef ASTRAI_NO_MMA
template <int HEAD_DIM, bool IsCausal, bool HasMask>
static inline void launch_paged_decode_mma(PagedAttentionParams<bf16>& p, int group_size) {
static inline void launch_paged_decode_mma(PagedAttentionParams<bf16>& p, int) {
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.kv_len + BC - 1) / BC;
p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total, 2);
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<HEAD_DIM, BC, 1, STAGES>;
dim3 grid(p.kv_head * num_passes, p.batch, p.num_splits);
@@ -166,10 +167,10 @@ static inline void launch_paged_decode_mma(PagedAttentionParams<bf16>& p, int gr
template <int HEAD_DIM, bool IsCausal, bool HasMask>
static inline void launch_paged_decode_scalar(PagedAttentionParams<bf16>& p, int group_size) {
int chunks_total = (p.kv_len + PDC_CHUNK - 1) / PDC_CHUNK;
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); // cap at 32 to respect 1024-thread limit
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<HEAD_DIM, IsCausal, HasMask><<<grid, block, smem>>>(p);
@@ -182,10 +183,37 @@ static inline void dispatch_paged_decode(PagedAttentionParams<bf16>& p) {
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, group_size);
DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_paged_decode_mma, HEAD_DIM, p, 0);
#else
DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_paged_decode_scalar, HEAD_DIM, p, group_size);
#endif
paged_attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
}
// ======================================================================
// Paged Prefill (SGLang-style: flat pool + ragged batch)
// ======================================================================
#ifndef ASTRAI_NO_MMA
template <int HEAD_DIM, bool IsCausal, bool HasMask>
static inline void launch_paged_prefill_mma(PagedAttentionParams<bf16>& p) {
constexpr int WARPS = 4;
constexpr int BC = (HEAD_DIM <= 128) ? 32 : 16;
using Traits = KernelTraits<HEAD_DIM, BC, WARPS, 2>;
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<Traits, IsCausal, HasMask><<<grid, block>>>(p);
}
#endif
template <int HEAD_DIM>
static inline void dispatch_paged_prefill(PagedAttentionParams<bf16>& p) {
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);
#endif
}
+147 -23
View File
@@ -134,54 +134,178 @@ inline void attn_pack_params(
pack_mask(mask, p);
}
// ---- attn_pack_paged_params ----
// ---- attn_pack_paged_decode_params ----
// SGLang-style: flat KV pool + req_to_token indexing + variable
// seq_lens via kv_indptr. Q is [batch, q_head, head_dim] (q_len=1 per req).
template<typename T>
inline void attn_pack_paged_params(
inline void attn_pack_paged_decode_params(
torch::Tensor q,
torch::Tensor page_table,
torch::Tensor k_cache,
torch::Tensor v_cache,
int64_t page_size,
int64_t kv_len,
torch::Tensor req_to_token,
torch::Tensor req_pool_indices,
torch::Tensor kv_indptr,
int64_t max_seq_len,
c10::optional<torch::Tensor> mask,
int64_t causal_offset,
double scale,
int64_t layout,
PagedAttentionParams<T>& p
) {
const at::cuda::OptionalCUDAGuard device_guard(device_of(q));
TORCH_CHECK(q.is_cuda() && page_table.is_cuda() && k_cache.is_cuda() && v_cache.is_cuda());
TORCH_CHECK(q.is_cuda() && k_cache.is_cuda() && v_cache.is_cuda());
TORCH_CHECK(req_to_token.is_cuda() && req_pool_indices.is_cuda() && kv_indptr.is_cuda());
TORCH_CHECK(q.dtype() == torch::kBFloat16, "q must be bf16");
TORCH_CHECK(k_cache.dtype() == torch::kBFloat16, "k_cache must be bf16");
TORCH_CHECK(v_cache.dtype() == torch::kBFloat16, "v_cache must be bf16");
TORCH_CHECK(page_table.dtype() == torch::kLong, "page_table must be int64");
TORCH_CHECK(k_cache.sizes() == v_cache.sizes(), "k_cache and v_cache must have identical shapes");
TORCH_CHECK(req_to_token.dtype() == torch::kLong, "req_to_token must be int64");
TORCH_CHECK(req_pool_indices.dtype() == torch::kLong, "req_pool_indices must be int64");
TORCH_CHECK(kv_indptr.dtype() == torch::kInt32, "kv_indptr must be int32");
TORCH_CHECK(k_cache.sizes() == v_cache.sizes(), "k_cache and v_cache must match");
TORCH_CHECK(k_cache.dim() == 3, "k_cache must be 3D [size, kv_head, head_dim]");
TORCH_CHECK(q.dim() == 3, "q must be 3D [batch, q_head, head_dim]");
extract_q_dims_and_strides(q, layout, p);
p.kv_head = (int)k_cache.size(2);
p.kv_len = (int)kv_len;
p.page_size = (int)page_size;
p.max_pages = (int)page_table.size(1);
TORCH_CHECK(q.size(2) == 1, "Q seq_len must be 1 (decode)");
p.batch = (int)q.size(0);
p.q_head = (int)q.size(1);
p.head_dim = (int)q.size(2);
p.kv_head = (int)k_cache.size(1);
TORCH_CHECK(k_cache.size(2) == p.head_dim, "k_cache head_dim mismatch");
TORCH_CHECK(p.head_dim % 32 == 0, "head_dim must be multiple of 32");
TORCH_CHECK(k_cache.size(1) == page_size,
"k_cache dim 1 must equal page_size, got ",
k_cache.size(1), " vs ", page_size);
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.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.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>();
p.qo_indptr = nullptr;
p.max_context_len = (int)req_to_token.size(1);
p.max_seq_len = (int)max_seq_len;
p.total_q = p.batch; // decode: 1 Q token per request
p.max_q_len = 1;
p.causal_offset = (int)causal_offset;
p.use_mask = (mask.has_value() && mask.value().defined()) ? 1 : 0;
p.scale = (scale > 0.0) ? (float)scale : 1.0f / sqrtf((float)p.head_dim);
p.page_table = page_table.data_ptr<int64_t>();
if (p.use_mask) {
auto m = mask.value();
TORCH_CHECK(m.is_cuda() && m.dtype() == torch::kBool, "mask must be bool CUDA");
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 = m.data_ptr<bool>();
} else {
p.mask = nullptr;
p.mask_b_stride = 0;
p.mask_h_stride = 0;
p.mask_q_stride = 0;
}
p.o = nullptr;
p.o_part = nullptr;
p.ml_part = nullptr;
}
// ---- attn_pack_paged_prefill_params ----
// SGLang-style: flat KV pool + req_to_token + ragged batch via qo_indptr.
// Q is [total_q, q_head, head_dim] (flattened across all requests).
template<typename T>
inline void attn_pack_paged_prefill_params(
torch::Tensor q,
torch::Tensor k_cache,
torch::Tensor v_cache,
torch::Tensor req_to_token,
torch::Tensor req_pool_indices,
torch::Tensor kv_indptr,
torch::Tensor qo_indptr,
c10::optional<torch::Tensor> mask,
int64_t max_q_len,
int64_t causal_offset,
double scale,
PagedAttentionParams<T>& p
) {
const at::cuda::OptionalCUDAGuard device_guard(device_of(q));
TORCH_CHECK(q.is_cuda() && k_cache.is_cuda() && v_cache.is_cuda());
TORCH_CHECK(req_to_token.is_cuda() && req_pool_indices.is_cuda());
TORCH_CHECK(kv_indptr.is_cuda() && qo_indptr.is_cuda());
TORCH_CHECK(q.dtype() == torch::kBFloat16, "q must be bf16");
TORCH_CHECK(k_cache.dtype() == torch::kBFloat16, "k_cache must be bf16");
TORCH_CHECK(v_cache.dtype() == torch::kBFloat16, "v_cache must be bf16");
TORCH_CHECK(req_to_token.dtype() == torch::kLong, "req_to_token must be int64");
TORCH_CHECK(req_pool_indices.dtype() == torch::kLong, "req_pool_indices must be int64");
TORCH_CHECK(kv_indptr.dtype() == torch::kInt32, "kv_indptr must be int32");
TORCH_CHECK(qo_indptr.dtype() == torch::kInt32, "qo_indptr must be int32");
TORCH_CHECK(k_cache.sizes() == v_cache.sizes(), "k_cache and v_cache must match");
TORCH_CHECK(k_cache.dim() == 3, "k_cache must be 3D [size, kv_head, head_dim]");
TORCH_CHECK(q.dim() == 3, "q must be 3D [total_q, q_head, head_dim]");
p.q_head = (int)q.size(1);
p.head_dim = (int)q.size(2);
p.kv_head = (int)k_cache.size(1);
p.batch = (int)req_pool_indices.size(0);
TORCH_CHECK(k_cache.size(2) == p.head_dim, "k_cache head_dim mismatch");
TORCH_CHECK(p.head_dim % 16 == 0, "head_dim must be multiple of 16");
TORCH_CHECK(p.q_head % p.kv_head == 0, "q_head must be divisible by kv_head");
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.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.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>();
p.qo_indptr = qo_indptr.data_ptr<int>();
p.max_context_len = (int)req_to_token.size(1);
p.total_q = (int)q.size(0); // prefill: flattened Q across all requests
p.max_q_len = (int)max_q_len;
// max_seq_len is unused by the prefill path (decode uses it for split
// computation); fill with max_q_len only to keep the POD struct defined.
p.max_seq_len = p.max_q_len;
p.causal_offset = (int)causal_offset;
p.use_mask = (mask.has_value() && mask.value().defined()) ? 1 : 0;
if (p.use_mask) {
auto m = mask.value();
TORCH_CHECK(m.is_cuda() && m.dtype() == torch::kBool, "mask must be bool CUDA");
TORCH_CHECK(m.size(0) == p.batch, "mask batch mismatch");
if (m.dim() == 2) {
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;
} 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);
} else {
TORCH_CHECK(false, "mask must be 2D or 4D");
}
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.scale = (scale > 0.0) ? (float)scale : 1.0f / sqrtf((float)p.head_dim);
p.o = nullptr;
p.o_part = nullptr;
p.ml_part = nullptr;
pack_mask(mask, p);
}
+15 -15
View File
@@ -3,23 +3,23 @@
torch::Tensor attn_paged_decode(
torch::Tensor q,
torch::Tensor page_table,
torch::Tensor k_cache,
torch::Tensor v_cache,
int64_t page_size,
int64_t kv_len,
torch::Tensor req_to_token,
torch::Tensor req_pool_indices,
torch::Tensor kv_indptr,
int64_t max_seq_len,
c10::optional<torch::Tensor> mask,
int64_t causal_offset,
double scale,
int64_t layout
double scale
) {
PagedAttentionParams<bf16> p;
attn_pack_paged_params(q, page_table, k_cache, v_cache,
page_size, kv_len, mask, causal_offset, scale, layout, 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);
auto O = torch::empty_strided(q.sizes(), q.strides(), q.options());
auto O_view = (layout == BLHD) ? O.transpose(1, 2) : O;
p.o = (bf16*)O_view.data_ptr();
auto O = torch::empty({q.size(0), q.size(1), q.size(2)}, q.options());
p.o = (bf16*)O.data_ptr();
alloc_split_partials(p);
DISPATCH_HEAD_DIM(p.head_dim, dispatch_paged_decode, p);
@@ -30,14 +30,14 @@ torch::Tensor attn_paged_decode(
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("attn_paged_decode", &attn_paged_decode,
py::arg("q"),
py::arg("page_table"),
py::arg("k_cache"),
py::arg("v_cache"),
py::arg("page_size"),
py::arg("kv_len"),
py::arg("req_to_token"),
py::arg("req_pool_indices"),
py::arg("kv_indptr"),
py::arg("max_seq_len"),
py::arg("mask") = py::none(),
py::arg("causal_offset") = -1,
py::arg("scale") = 0.0,
py::arg("layout") = (int64_t)BHLD,
"Paged GQA decode — split-KV with direct page-table access.");
"SGLang-style paged decode: flat KV pool + req_to_token + kv_indptr.");
}
+19 -20
View File
@@ -5,6 +5,8 @@
#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 <int HEAD_DIM, bool IsCausal, bool HasMask>
__global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams<bf16> p) {
int batch = blockIdx.x / p.kv_head;
@@ -15,8 +17,11 @@ __global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams<bf16> p)
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_b + q_head * p.q_stride_h
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++)
@@ -26,16 +31,19 @@ __global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams<bf16> p)
extern __shared__ __align__(16) bf16 k_smem[];
int chunks_total = (p.kv_len + PDC_CHUNK - 1) / PDC_CHUNK;
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 + q_head * p.mask_h_stride;
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, p.kv_len - chunk_start);
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;
@@ -43,14 +51,9 @@ __global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams<bf16> p)
int s = i / p.head_dim;
int d_dim = i % p.head_dim;
int pos = chunk_start + s;
int logical_page = pos / p.page_size;
int page_offset = pos % p.page_size;
int phys_page = p.page_table[batch * p.max_pages + logical_page];
if (phys_page >= 0) {
int64_t off = (int64_t)phys_page * p.page_size * p.kv_head * p.head_dim
+ (int64_t)page_offset * p.kv_head * p.head_dim
+ (int64_t)kv_head * p.head_dim
+ d_dim;
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);
@@ -85,17 +88,13 @@ __global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams<bf16> p)
d = d * alpha + beta;
int pos = chunk_start + s;
int logical_page = pos / p.page_size;
int page_offset = pos % p.page_size;
int phys_page = p.page_table[batch * p.max_pages + logical_page];
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 (phys_page >= 0) {
int64_t v_base = (int64_t)phys_page * p.page_size * p.kv_head * p.head_dim
+ (int64_t)page_offset * p.kv_head * p.head_dim
+ (int64_t)kv_head * p.head_dim;
} 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,
@@ -148,6 +147,6 @@ __global__ void paged_attn_decode_combine_kernel(PagedAttentionParams<bf16> 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 = batch * p.q_stride_l + q_head * p.q_stride_h + d * p.q_stride_d;
p.o[o_off] = __float2bfloat16(acc * inv);
}
+32 -29
View File
@@ -5,12 +5,16 @@
#include "attn_mma_utils.cuh"
#include "attn_warp_utils.cuh"
// Paged split-KV tensor-core decode via GQA head-packing.
// Reads K/V directly from the page pool through a page table — one tile
// (BC=32) fits within a single page (page_size >= 32), so the page-table
// lookup happens once per tile for cp.async.
// SGLang-style split-KV tensor-core decode.
//
// IsCausal and HasMask are compile-time bools.
// 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 <typename Traits, bool IsCausal, bool HasMask>
__global__ void paged_attn_decode_split_kv_mma_kernel(PagedAttentionParams<bf16> p) {
const int lane = threadIdx.x;
@@ -22,6 +26,10 @@ __global__ void paged_attn_decode_split_kv_mma_kernel(PagedAttentionParams<bf16>
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;
@@ -38,13 +46,14 @@ __global__ void paged_attn_decode_split_kv_mma_kernel(PagedAttentionParams<bf16>
}
__syncwarp();
const int q_base = batch * p.q_stride_b + q_head0 * p.q_stride_h;
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<Traits::KD>(p.q + q_base, p.q_stride_h, p.q_stride_d,
qra, qrb, va, vb, tid4, Qa);
load_q_mma_frags<Traits::KD>(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
@@ -52,19 +61,19 @@ __global__ void paged_attn_decode_split_kv_mma_kernel(PagedAttentionParams<bf16>
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 = (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);
const int64_t page_stride = (int64_t)p.page_size * p.kv_head * Traits::HEAD_DIM;
const int64_t pos_stride = (int64_t)p.kv_head * Traits::HEAD_DIM;
// 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: paged addressing ----
// Unified per-element page-table lookup. When page_size >= BC, all
// elements in a tile share the same page, so the lookup is redundant
// but harmless (L1-cached). This avoids a branch on page_size.
// ---- 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;
@@ -74,16 +83,13 @@ __global__ void paged_attn_decode_split_kv_mma_kernel(PagedAttentionParams<bf16>
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);
if constexpr (HasMask) {
valid = valid && p.mask[batch * p.mask_b_stride + kc];
}
int phys_page = valid ? p.page_table[batch * p.max_pages + kc] : 0;
valid = valid && (phys_page >= 0);
int page_off = kc % p.page_size;
int64_t gmem_base = (int64_t)phys_page * page_stride
+ (int64_t)page_off * pos_stride
+ head_off;
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);
@@ -91,10 +97,6 @@ __global__ void paged_attn_decode_split_kv_mma_kernel(PagedAttentionParams<bf16>
cp_async_commit();
};
// ---- Multi-stage cp.async pipeline ----
// Prologue loads STAGES tiles; each loop iteration waits only for the
// oldest outstanding group (wait_group<STAGES-1>) so the STAGES-1 newer
// tile loads stay in flight and overlap with the current tile's compute.
constexpr int STAGES = Traits::STAGES;
const int ntiles = ti_end - ti_begin;
@@ -111,8 +113,9 @@ __global__ void paged_attn_decode_split_kv_mma_kernel(PagedAttentionParams<bf16>
Sacc[n8][0] *= p.scale, Sacc[n8][1] *= p.scale,
Sacc[n8][2] *= p.scale, Sacc[n8][3] *= p.scale;
int maxc = IsCausal ? min(p.kv_len, p.causal_offset + 1) : p.kv_len;
mma_softmax_tile<Traits, HasMask>(kv0, maxc, maxc,
// 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<Traits, HasMask>(kv0, seq_len, seq_len,
0, 0,
p.mask_b_stride, 0, 0,
batch, 0,
@@ -136,7 +139,6 @@ __global__ void paged_attn_decode_split_kv_mma_kernel(PagedAttentionParams<bf16>
load_tile(ti_begin + it + STAGES, (it + STAGES) & (STAGES - 1));
}
} else {
// Fewer tiles than stages: load all, wait for all, process.
for (int i = 0; i < ntiles; i++)
load_tile(ti_begin + i, i);
cp_async_wait_group<0>();
@@ -145,6 +147,7 @@ __global__ void paged_attn_decode_split_kv_mma_kernel(PagedAttentionParams<bf16>
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;
+45
View File
@@ -0,0 +1,45 @@
#include "attn_dispatchers.cuh"
#include "attn_entry_utils.cuh"
torch::Tensor attn_paged_prefill(
torch::Tensor q,
torch::Tensor k_cache,
torch::Tensor v_cache,
torch::Tensor req_to_token,
torch::Tensor req_pool_indices,
torch::Tensor kv_indptr,
torch::Tensor qo_indptr,
c10::optional<torch::Tensor> mask,
int64_t max_q_len,
int64_t causal_offset,
double scale
) {
PagedAttentionParams<bf16> p;
attn_pack_paged_prefill_params(q, k_cache, v_cache,
req_to_token, req_pool_indices,
kv_indptr, qo_indptr, mask,
max_q_len, causal_offset, scale, p);
auto O = torch::empty({q.size(0), q.size(1), q.size(2)}, q.options());
p.o = (bf16*)O.data_ptr();
DISPATCH_HEAD_DIM(p.head_dim, dispatch_paged_prefill, p);
C10_CUDA_CHECK(cudaGetLastError());
return O;
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("attn_paged_prefill", &attn_paged_prefill,
py::arg("q"),
py::arg("k_cache"),
py::arg("v_cache"),
py::arg("req_to_token"),
py::arg("req_pool_indices"),
py::arg("kv_indptr"),
py::arg("qo_indptr"),
py::arg("mask") = py::none(),
py::arg("max_q_len"),
py::arg("causal_offset") = -1,
py::arg("scale") = 0.0,
"SGLang-style paged prefill: flat KV pool + ragged batch.");
}
@@ -0,0 +1,164 @@
#pragma once
#include <cfloat>
#include <cuda_bf16.h>
#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 <typename Traits, bool IsCausal, bool HasMask>
__global__ void paged_attn_prefill_split_q_mma_kernel(PagedAttentionParams<bf16> 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<Traits::KD>(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<Traits>(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<Traits, HasMask>(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<Traits>(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;
}
}
}