- Thread a cudaStream_t through attn dispatchers onto torch's current stream - Scope the device guard to the entry function so kernels run on tensor device - DISPATCH_HEAD_DIM now forwards varargs so stream reaches each dispatch - Parallelize CPU reference kernels with OpenMP (paged test 31s -> 7s) - Merge decode/prefill standalone tests into attn_test.cu with correctness tables - Drop bench error column (CPU ref too slow at large sizes) - Update cuda_kernels.md for the merged test layout
312 lines
12 KiB
Plaintext
312 lines
12 KiB
Plaintext
#pragma once
|
|
#include <float.h>
|
|
#include <torch/extension.h>
|
|
#include <c10/cuda/CUDAGuard.h>
|
|
#include "attn_common.h"
|
|
#include "attn_warp_utils.cuh"
|
|
|
|
using bf16 = __nv_bfloat16;
|
|
|
|
// Dispatch head_dim: shared macro — avoids C++20 lambda template syntax.
|
|
// Usage: DISPATCH_HEAD_DIM(hd, fn, args...)
|
|
// Expands to: fn<32>(args...); fn<64>(args...); etc.
|
|
#define DISPATCH_HEAD_DIM(hd, fn, ...) \
|
|
switch (hd) { \
|
|
case 32: fn<32>(__VA_ARGS__); break; \
|
|
case 64: fn<64>(__VA_ARGS__); break; \
|
|
case 128: fn<128>(__VA_ARGS__); break; \
|
|
case 256: fn<256>(__VA_ARGS__); break; \
|
|
default: \
|
|
TORCH_CHECK(false, "unsupported head_dim ", hd, \
|
|
" (supported: 32, 64, 128, 256)"); \
|
|
}
|
|
|
|
// The split kernel unconditionally writes every (batch, q_head, split) slot it
|
|
// owns — including empty split ranges, which store m = -FLT_MAX so the combine
|
|
// skips them. Allocators are therefore left uninitialized (torch::empty); the
|
|
// per-call memset (torch::zeros / torch::full) was pure overhead.
|
|
template<typename P>
|
|
inline void alloc_split_partials(P& p) {
|
|
auto fopt = torch::TensorOptions().dtype(torch::kFloat32).device(torch::kCUDA);
|
|
auto o_part = torch::empty(at::IntArrayRef{p.batch, p.q_head, MAX_SPLITS, p.head_dim}, fopt);
|
|
auto ml_part = torch::empty(at::IntArrayRef{p.batch, p.q_head, MAX_SPLITS, 2}, fopt);
|
|
p.o_part = (float*)o_part.data_ptr();
|
|
p.ml_part = (float*)ml_part.data_ptr();
|
|
}
|
|
|
|
// ---- Shared Q-dims + strides extraction ----
|
|
template <typename P>
|
|
inline void extract_q_dims_and_strides(torch::Tensor& q, int64_t layout, P& p) {
|
|
if (layout == BLHD) q = q.transpose(1, 2);
|
|
p.batch = (int)q.size(0);
|
|
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);
|
|
}
|
|
|
|
// ---- Shared mask packing ----
|
|
// Accepts 2D [batch, kv_len], 3D [batch, q_len, kv_len],
|
|
// or 4D [batch, n_heads, q_len, kv_len].
|
|
// Head/q dimensions with size 1 broadcast (stride set to 0).
|
|
template <typename P>
|
|
inline void pack_mask(const c10::optional<torch::Tensor>& mask, P& p) {
|
|
if (p.use_mask) {
|
|
auto m = mask.value();
|
|
TORCH_CHECK(m.is_cuda(), "mask must be on CUDA");
|
|
TORCH_CHECK(m.dtype() == torch::kBool, "mask must be bool");
|
|
TORCH_CHECK(m.size(0) == p.batch, "mask batch mismatch");
|
|
TORCH_CHECK(m.size(m.dim() - 1) == p.kv_len, "mask kv_len mismatch");
|
|
if (m.dim() == 2) {
|
|
p.mask_b_stride = (int)m.stride(0);
|
|
p.mask_h_stride = 0;
|
|
p.mask_q_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);
|
|
} 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);
|
|
} else {
|
|
TORCH_CHECK(false, "mask must be 2D, 3D, 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;
|
|
}
|
|
}
|
|
|
|
// ---- attn_pack_params (contiguous KV) ----
|
|
template<typename T>
|
|
inline void attn_pack_params(
|
|
torch::Tensor q,
|
|
torch::Tensor k,
|
|
torch::Tensor v,
|
|
c10::optional<torch::Tensor> mask,
|
|
int64_t causal_offset,
|
|
double scale,
|
|
int64_t layout,
|
|
AttentionParams<T>& p
|
|
) {
|
|
const at::cuda::OptionalCUDAGuard device_guard(device_of(q));
|
|
|
|
TORCH_CHECK(q.is_cuda() && k.is_cuda() && v.is_cuda());
|
|
TORCH_CHECK(q.dtype() == torch::kBFloat16);
|
|
TORCH_CHECK(k.dtype() == torch::kBFloat16);
|
|
TORCH_CHECK(v.dtype() == torch::kBFloat16);
|
|
TORCH_CHECK(k.sizes() == v.sizes(), "K and V must have identical shapes");
|
|
TORCH_CHECK(q.dim() == 4 && k.dim() == 4, "Q/K/V must be 4D");
|
|
|
|
extract_q_dims_and_strides(q, layout, p);
|
|
|
|
if (layout == BLHD) k = k.transpose(1, 2), v = v.transpose(1, 2);
|
|
|
|
p.kv_head = (int)k.size(1);
|
|
p.kv_len = (int)k.size(2);
|
|
TORCH_CHECK(k.size(3) == p.head_dim, "K/V head_dim must match Q");
|
|
|
|
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.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.o = nullptr;
|
|
p.o_part = nullptr;
|
|
p.ml_part = nullptr;
|
|
|
|
pack_mask(mask, p);
|
|
}
|
|
|
|
// ---- 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_decode_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,
|
|
int64_t max_seq_len,
|
|
c10::optional<torch::Tensor> mask,
|
|
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() && 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(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]");
|
|
|
|
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(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);
|
|
|
|
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;
|
|
}
|