refactor: reorganize CUDA kernels into per-family directories
- move attention kernels to csrc/kernels/attention/ and rotary to rotary/ - add shared common/mma.cuh (mma_sync, ldmatrix) and device.cuh (sm checks) - split fp8_mm into three-layer fp8/common.h, gemm.cuh, mm.cu - fix fused FP8 GEMM ldmatrix lane indexing to fix OOB shared reads - update extension ops, loader, and kernel tests
This commit is contained in:
@@ -0,0 +1,77 @@
|
||||
#pragma once
|
||||
|
||||
// Tensor layout for Q/K/V tensors passed to attention kernels.
|
||||
// Internally, kernels always operate on BHLD [batch, n_heads, seq_len, head_dim].
|
||||
// When the caller passes BLHD, dims 1 and 2 are transposed at entry.
|
||||
enum TensorLayout : int {
|
||||
BHLD = 0, // [batch, n_heads, seq_len, head_dim]
|
||||
BLHD = 1, // [batch, seq_len, n_heads, head_dim]
|
||||
};
|
||||
|
||||
|
||||
// Unified attention params covering BOTH addressing modes:
|
||||
// - Contiguous K/V: dense [batch, kv_head, kv_len, head_dim] tensors (k/v).
|
||||
// - Paged (SGLang-style): flat pool [size, kv_head, head_dim] + req_to_token.
|
||||
// Each kernel selects the addressing via a KVSource policy (see
|
||||
// layout_policies.cuh); a given call only touches the fields of one mode, so
|
||||
// this is a POD shared by both paths rather than two parallel structs that
|
||||
// drift out of sync.
|
||||
template<typename T, typename AT = float>
|
||||
struct AttentionParams {
|
||||
// Shape
|
||||
int batch;
|
||||
int q_head;
|
||||
int kv_head;
|
||||
int head_dim;
|
||||
int q_len; // Per-request in contiguous mode; total_q in paged mode.
|
||||
int kv_len; // Contiguous mode; paged mode uses kv_indptr.
|
||||
|
||||
// Attention behavior
|
||||
float scale;
|
||||
// -1 = non-causal; >=0 = absolute position of first Q token
|
||||
int causal_offset;
|
||||
int use_mask;
|
||||
|
||||
// pointers
|
||||
const T* __restrict__ q_ptr;
|
||||
const T* __restrict__ k_ptr;
|
||||
const T* __restrict__ v_ptr;
|
||||
const T* __restrict__ new_k_ptr;
|
||||
const T* __restrict__ new_v_ptr;
|
||||
T* __restrict__ o_ptr;
|
||||
const bool* __restrict__ mask;
|
||||
|
||||
// strides
|
||||
int q_b_stride;
|
||||
int q_h_stride;
|
||||
int q_l_stride;
|
||||
int q_d_stride;
|
||||
|
||||
int kv_b_stride;
|
||||
int kv_h_stride;
|
||||
int kv_l_stride;
|
||||
int kv_d_stride;
|
||||
|
||||
int new_kv_b_stride;
|
||||
int new_kv_h_stride;
|
||||
|
||||
int mask_b_stride;
|
||||
int mask_h_stride;
|
||||
int mask_l_stride;
|
||||
|
||||
// Paged K/V addressing
|
||||
const int* __restrict__ req_to_token; // [num_reqs, max_context_len]
|
||||
const int* __restrict__ req_pool_indices; // [batch]
|
||||
const int* __restrict__ kv_indptr; // [batch + 1]
|
||||
const int* __restrict__ qo_indptr; // [batch + 1] or nullptr for decode
|
||||
const int* __restrict__ q_tile_to_batch; // [num_q_tiles], prefill only
|
||||
const int* __restrict__ q_tile_to_index; // [num_q_tiles], prefill only
|
||||
int num_q_tiles;
|
||||
int max_context_len; // req_to_token stride (dim 1)
|
||||
|
||||
// Decode split-KV workspace
|
||||
int num_splits;
|
||||
AT* __restrict__ o_part;
|
||||
AT* __restrict__ ml_part;
|
||||
|
||||
};
|
||||
@@ -0,0 +1,63 @@
|
||||
#include "dispatchers.cuh"
|
||||
#include "entry_utils.cuh"
|
||||
|
||||
torch::Tensor attn_decode(
|
||||
torch::Tensor q,
|
||||
torch::Tensor k,
|
||||
torch::Tensor v,
|
||||
c10::optional<torch::Tensor> mask,
|
||||
int64_t causal_offset,
|
||||
double scale,
|
||||
int64_t layout,
|
||||
c10::optional<torch::Tensor> o_part_buf,
|
||||
c10::optional<torch::Tensor> ml_part_buf
|
||||
) {
|
||||
const at::cuda::OptionalCUDAGuard device_guard(device_of(q));
|
||||
auto stream = at::cuda::getCurrentCUDAStream();
|
||||
|
||||
AttentionParams<bf16> p;
|
||||
attn_pack_params(q, k, v, mask, causal_offset, scale, layout, p);
|
||||
TORCH_CHECK(p.q_len == 1, "Q seq_len must be 1");
|
||||
TORCH_CHECK(p.head_dim % 32 == 0, "head_dim must be multiple of 32");
|
||||
|
||||
auto O = torch::empty_strided(q.sizes(), q.strides(), q.options());
|
||||
auto O_view = (layout == BLHD) ? O.transpose(1, 2) : O;
|
||||
p.o_ptr = (bf16*)O_view.data_ptr();
|
||||
|
||||
if (o_part_buf.has_value() && ml_part_buf.has_value()
|
||||
&& o_part_buf->defined() && ml_part_buf->defined()) {
|
||||
TORCH_CHECK(o_part_buf->scalar_type() == torch::kFloat32, "o_part_buf must be f32");
|
||||
TORCH_CHECK(ml_part_buf->scalar_type() == torch::kFloat32, "ml_part_buf must be f32");
|
||||
int64_t o_needed = (int64_t)p.batch * p.q_head * MAX_SPLITS * p.head_dim;
|
||||
int64_t ml_needed = (int64_t)p.batch * p.q_head * MAX_SPLITS * 2;
|
||||
TORCH_CHECK(o_part_buf->numel() >= o_needed,
|
||||
"o_part_buf too small: need ", o_needed, " got ", o_part_buf->numel());
|
||||
TORCH_CHECK(ml_part_buf->numel() >= ml_needed,
|
||||
"ml_part_buf too small: need ", ml_needed, " got ", ml_part_buf->numel());
|
||||
TORCH_CHECK(o_part_buf->is_cuda() && ml_part_buf->is_cuda(),
|
||||
"split buffers must be CUDA tensors");
|
||||
TORCH_CHECK(o_part_buf->is_contiguous() && ml_part_buf->is_contiguous(),
|
||||
"split buffers must be contiguous");
|
||||
p.o_part = (float*)o_part_buf->data_ptr();
|
||||
p.ml_part = (float*)ml_part_buf->data_ptr();
|
||||
} else {
|
||||
alloc_split_partials(p);
|
||||
}
|
||||
DISPATCH_HEAD_DIM(p.head_dim, dispatch_decode, p, stream);
|
||||
C10_CUDA_CHECK(cudaGetLastError());
|
||||
return O;
|
||||
}
|
||||
|
||||
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
||||
m.def("attn_decode", &attn_decode,
|
||||
py::arg("q"),
|
||||
py::arg("k"),
|
||||
py::arg("v"),
|
||||
py::arg("mask") = py::none(),
|
||||
py::arg("causal_offset") = -1,
|
||||
py::arg("scale") = 0.0,
|
||||
py::arg("layout") = (int64_t)BHLD,
|
||||
py::arg("o_part_buf") = py::none(),
|
||||
py::arg("ml_part_buf") = py::none(),
|
||||
"GQA decode (tensor-core head-packing on sm_80+, scalar fallback)");
|
||||
}
|
||||
@@ -0,0 +1,144 @@
|
||||
#pragma once
|
||||
#include <cuda_bf16.h>
|
||||
#include <float.h>
|
||||
#include "common.h"
|
||||
#include "layout_policies.cuh"
|
||||
#include "warp_utils.cuh"
|
||||
constexpr int DC_CHUNK = 64;
|
||||
|
||||
// Scalar split-KV decode (fallback for sm < 80, no tensor cores), unified
|
||||
// across contiguous and paged (SGLang flat-pool) K/V via the KV template
|
||||
// parameter. For decode the query is the last token, so its valid range
|
||||
// [0, seq_len) IS the causal range; KV::decode_attend_len expresses that
|
||||
// bound per addressing mode (contig clips to causal_offset, paged = seq_len).
|
||||
template <int HEAD_DIM, typename KV, bool IsCausal, bool HasMask>
|
||||
__global__ void attn_decode_split_kv_kernel(AttentionParams<bf16> p) {
|
||||
int batch = blockIdx.x / p.kv_head;
|
||||
int kv_head = blockIdx.x % p.kv_head;
|
||||
int split = blockIdx.z;
|
||||
int group_size = blockDim.y;
|
||||
int q_head = kv_head * group_size + threadIdx.y;
|
||||
int lane = threadIdx.x;
|
||||
int hd_per_thread = p.head_dim / 32;
|
||||
|
||||
const int seq_len = KV::kv_len(p, batch);
|
||||
const KVContext kctx = KV::template make_ctx<HEAD_DIM>(p, batch, kv_head);
|
||||
|
||||
// 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_d_stride;
|
||||
for (int i = 0; i < hd_per_thread; i++)
|
||||
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;
|
||||
|
||||
float m = -FLT_MAX, d = 0.0f, acc_reg[8] = {0.0f};
|
||||
|
||||
extern __shared__ __align__(16) bf16 smem[];
|
||||
bf16* k_smem = smem;
|
||||
bf16* v_smem = smem + DC_CHUNK * p.head_dim;
|
||||
|
||||
// Split-KV: each split processes a contiguous subset of chunks
|
||||
int chunks_total = (seq_len + DC_CHUNK - 1) / DC_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);
|
||||
|
||||
for (int ci = ch_begin; ci < ch_end; ci++) {
|
||||
int chunk_start = ci * DC_CHUNK;
|
||||
int this_chunk = min(DC_CHUNK, seq_len - chunk_start);
|
||||
|
||||
// Load K and V into shared memory (addressing via KV policy;
|
||||
// paged guards empty slots with zero-fill).
|
||||
int total = this_chunk * p.head_dim;
|
||||
for (int i = threadIdx.y * 32 + lane; i < total;
|
||||
i += blockDim.x * blockDim.y) {
|
||||
int s = i / p.head_dim;
|
||||
int d_dim = i % p.head_dim;
|
||||
int kc = chunk_start + s;
|
||||
KVAddr a = KV::template decode_addr<1>(
|
||||
p, kctx, batch, kv_head, kc, d_dim, true, true);
|
||||
k_smem[i] = a.valid ? *reinterpret_cast<const bf16*>(a.k) : (bf16)0.f;
|
||||
v_smem[i] = a.valid ? *reinterpret_cast<const bf16*>(a.v) : (bf16)0.f;
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
for (int s = 0; s < this_chunk; s++) {
|
||||
float partial = 0.0f;
|
||||
for (int i = 0; i < hd_per_thread; i++)
|
||||
partial += q_reg[i] * __bfloat162float(
|
||||
k_smem[s * p.head_dim + lane * hd_per_thread + i]);
|
||||
partial = warp_reduce_sum(partial) * p.scale;
|
||||
|
||||
int kv_idx = chunk_start + s;
|
||||
if constexpr (HasMask) {
|
||||
if (!p.mask[mask_base + kv_idx])
|
||||
partial = -FLT_MAX;
|
||||
}
|
||||
if constexpr (IsCausal) {
|
||||
if (kv_idx >= KV::decode_attend_len(p, batch))
|
||||
partial = -FLT_MAX;
|
||||
}
|
||||
|
||||
float new_m = fmaxf(m, partial);
|
||||
float alpha = __expf(m - new_m);
|
||||
float beta = __expf(partial - new_m);
|
||||
d = d * alpha + beta;
|
||||
|
||||
for (int i = 0; i < hd_per_thread; i++) {
|
||||
float vv = __bfloat162float(v_smem[s * p.head_dim + lane * hd_per_thread + i]);
|
||||
acc_reg[i] = fmaf(acc_reg[i], alpha, vv * beta);
|
||||
}
|
||||
m = new_m;
|
||||
}
|
||||
__syncthreads();
|
||||
}
|
||||
|
||||
// ---- write UN-normalised partials for this split ----
|
||||
size_t bh = (size_t)batch * p.q_head + q_head;
|
||||
size_t slot = bh * MAX_SPLITS + split;
|
||||
int d0 = lane * hd_per_thread;
|
||||
for (int i = 0; i < hd_per_thread; i++) {
|
||||
int dd = d0 + i;
|
||||
p.o_part[slot * p.head_dim + dd] = acc_reg[i];
|
||||
}
|
||||
if (lane == 0) {
|
||||
p.ml_part[slot * 2] = m;
|
||||
p.ml_part[slot * 2 + 1] = d;
|
||||
}
|
||||
}
|
||||
|
||||
// Split-combine: merges the per-split partials (o_part/ml_part) into the
|
||||
// final normalised O. KV selects the O addressing (contig batch stride vs
|
||||
// paged row stride).
|
||||
template <typename KV>
|
||||
__global__ void attn_decode_combine_kernel(AttentionParams<bf16> p) {
|
||||
int bh = blockIdx.x;
|
||||
int d = threadIdx.x;
|
||||
if (d >= p.head_dim) return;
|
||||
|
||||
int batch = bh / p.q_head;
|
||||
int q_head = bh % p.q_head;
|
||||
|
||||
size_t split_base = (size_t)bh * MAX_SPLITS;
|
||||
const float* mlp = p.ml_part + split_base * 2;
|
||||
const float* op = p.o_part + split_base * p.head_dim;
|
||||
|
||||
float m = -FLT_MAX, l = 0.0f, acc = 0.0f;
|
||||
for (int s = 0; s < p.num_splits; s++) {
|
||||
float mi = mlp[s * 2];
|
||||
if (mi <= -FLT_MAX) continue;
|
||||
float li = mlp[s * 2 + 1];
|
||||
float nm = fmaxf(m, mi);
|
||||
float corr = __expf(m - nm);
|
||||
float e = __expf(mi - nm);
|
||||
acc = fmaf(acc, corr, op[s * p.head_dim + d] * e);
|
||||
l = fmaf(l, corr, li * e);
|
||||
m = nm;
|
||||
}
|
||||
|
||||
float inv = (l > 1e-20f) ? (1.0f / l) : 0.0f;
|
||||
int o_off = KV::q_decode_base(p, batch, q_head) + d * p.q_d_stride;
|
||||
p.o_ptr[o_off] = __float2bfloat16(acc * inv);
|
||||
}
|
||||
@@ -0,0 +1,183 @@
|
||||
#pragma once
|
||||
#include <cfloat>
|
||||
#include <cuda_bf16.h>
|
||||
#include "common.h"
|
||||
#include "layout_policies.cuh"
|
||||
#include "mma_utils.cuh"
|
||||
#include "warp_utils.cuh"
|
||||
|
||||
// Split-K (FlashDecoding) tensor-core decode via GQA head-packing, unified
|
||||
// across contiguous and paged (SGLang flat-pool) K/V via the KV template
|
||||
// parameter. Decode has q_len == 1, so we pack G = q_head/kv_head query
|
||||
// heads into the M=16 rows of mma.sync.m16n8k16, turning G independent GEMVs
|
||||
// into a single GEMM that reuses each loaded K/V tile across all G heads.
|
||||
//
|
||||
// KV = ContigKV (dense tensors) or PagedKV (flat pool + req_to_token).
|
||||
// IsCausal and HasMask are compile-time bools — no runtime branch in the
|
||||
// inner compute loop.
|
||||
//
|
||||
// Traits = KernelTraits<HEAD_DIM, BC=16, WARPS=1, STAGES=2>.
|
||||
template <typename Traits, typename KV, bool IsCausal, bool HasMask>
|
||||
__global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
|
||||
const int lane = threadIdx.x;
|
||||
const int gid = lane >> 2;
|
||||
const int tid4 = lane & 3;
|
||||
|
||||
const int pass = blockIdx.x / p.kv_head;
|
||||
const int kv_head = blockIdx.x % p.kv_head;
|
||||
const int batch = blockIdx.y;
|
||||
const int split = blockIdx.z;
|
||||
|
||||
constexpr int MAX_G = 16;
|
||||
const int G_total = p.q_head / p.kv_head;
|
||||
const int g_begin = pass * MAX_G;
|
||||
const int G = min(MAX_G, G_total - g_begin);
|
||||
const int q_head0 = kv_head * G_total + g_begin;
|
||||
|
||||
// Per-request seq_len (paged reads kv_indptr; contig uses p.kv_len).
|
||||
const int seq_len = KV::kv_len(p, batch);
|
||||
const KVContext kctx = KV::template make_ctx<Traits::HEAD_DIM>(p, batch, kv_head);
|
||||
|
||||
// Double-buffered shared memory for K/V (no sQ needed)
|
||||
__shared__ __align__(16) bf16 sK[Traits::STAGES * Traits::BC * Traits::LD];
|
||||
__shared__ __align__(16) bf16 sV[Traits::STAGES * Traits::BC * Traits::LD];
|
||||
|
||||
// Load Q directly from global into mma A-operand registers.
|
||||
const int q_base = KV::q_decode_base(p, batch, q_head0);
|
||||
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_ptr + q_base, p.q_h_stride, p.q_d_stride,
|
||||
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 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);
|
||||
|
||||
// ---- Load tile lambda: predicated cp.async (addressing via KV policy) ----
|
||||
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 = lane * 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;
|
||||
// All GQA passes consume new K/V directly. Only the first pass
|
||||
// persists it, so no cross-block synchronization is required.
|
||||
KVAddr a = KV::template decode_addr<Traits::VEC>(
|
||||
p, kctx, batch, kv_head, kc, d, valid, pass == 0);
|
||||
int off = r * Traits::LD + swiz_col(d, r, Traits::SWIZ_MASK);
|
||||
cp_async_16_pred(&dK[off], a.k, a.valid);
|
||||
cp_async_16_pred(&dV[off], a.v, a.valid);
|
||||
}
|
||||
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;
|
||||
|
||||
auto process_tile = [&](int it, int buf) {
|
||||
const bf16* bK = sK + buf * Traits::BC * Traits::LD;
|
||||
const bf16* bV = sV + buf * Traits::BC * Traits::LD;
|
||||
int kv0 = (ti_begin + it) * Traits::BC;
|
||||
|
||||
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;
|
||||
|
||||
// Decode: q_len=1, so qrow0=qrow1=0. Paged treats [0, seq_len) as
|
||||
// the causal range (query is the last token); contig clips to the
|
||||
// causal_offset bound. Dead code eliminated when IsCausal == false.
|
||||
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_l_stride,
|
||||
batch, q_head0 + gid, q_head0 + gid + 8,
|
||||
p.mask,
|
||||
va, vb,
|
||||
Sacc, Oacc, m0, m1, l0, l1, lane);
|
||||
|
||||
mma_pv_accumulate<Traits>(Sacc, bV, lane, Oacc);
|
||||
};
|
||||
|
||||
if (ntiles >= STAGES) {
|
||||
#pragma unroll
|
||||
for (int i = 0; i < STAGES; i++)
|
||||
load_tile(ti_begin + i, i);
|
||||
|
||||
for (int it = 0; it < ntiles; it++) {
|
||||
if (it + 1 == ntiles)
|
||||
cp_async_wait_group<0>();
|
||||
else
|
||||
cp_async_wait_group<STAGES - 1>();
|
||||
__syncwarp();
|
||||
process_tile(it, it & (STAGES - 1));
|
||||
__syncwarp();
|
||||
if (it + STAGES < ntiles)
|
||||
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>();
|
||||
__syncwarp();
|
||||
for (int it = 0; it < ntiles; it++)
|
||||
process_tile(it, it);
|
||||
}
|
||||
|
||||
// ---- write UN-normalised partials for this split ----
|
||||
auto split_slot = [&](int h) -> size_t {
|
||||
size_t bh = (size_t)batch * p.q_head + h;
|
||||
return bh * MAX_SPLITS + split;
|
||||
};
|
||||
#pragma unroll
|
||||
for (int dn8 = 0; dn8 < Traits::DN8; dn8++) {
|
||||
int d = dn8 * 8 + 2 * tid4;
|
||||
int r0 = gid, r1 = gid + 8;
|
||||
if (r0 < G) {
|
||||
int h = q_head0 + r0;
|
||||
float* op = p.o_part + split_slot(h) * Traits::HEAD_DIM;
|
||||
op[d] = Oacc[dn8][0];
|
||||
op[d + 1] = Oacc[dn8][1];
|
||||
}
|
||||
if (r1 < G) {
|
||||
int h = q_head0 + r1;
|
||||
float* op = p.o_part + split_slot(h) * Traits::HEAD_DIM;
|
||||
op[d] = Oacc[dn8][2];
|
||||
op[d + 1] = Oacc[dn8][3];
|
||||
}
|
||||
}
|
||||
if (tid4 == 0) {
|
||||
int r0 = gid, r1 = gid + 8;
|
||||
if (r0 < G) {
|
||||
int h = q_head0 + r0;
|
||||
float* mp = p.ml_part + split_slot(h) * 2;
|
||||
mp[0] = m0; mp[1] = l0;
|
||||
}
|
||||
if (r1 < G) {
|
||||
int h = q_head0 + r1;
|
||||
float* mp = p.ml_part + split_slot(h) * 2;
|
||||
mp[0] = m1; mp[1] = l1;
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,233 @@
|
||||
#pragma once
|
||||
// Shared attention dispatchers — used by both production .cu and test .cu.
|
||||
// No torch dependency; pure CUDA.
|
||||
//
|
||||
// The paged and contiguous kernels are unified by the KVSource policy
|
||||
// (ContigKV / PagedKV from layout_policies.cuh), so each launcher struct
|
||||
// below is templated on KV and the paged dispatch is just the same launcher
|
||||
// instantiated with PagedKV. Only the grid/split math differs, and that is
|
||||
// covered by KV::host_q_len / KV::host_kv_len.
|
||||
|
||||
#include <cuda_runtime.h>
|
||||
#include <algorithm>
|
||||
#include "warp_utils.cuh"
|
||||
#include "layout_policies.cuh"
|
||||
#include "prefill_split_q.cuh"
|
||||
#include "decode_split_kv.cuh"
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
#include "prefill_split_q_mma.cuh"
|
||||
#include "decode_split_kv_mma.cuh"
|
||||
#endif
|
||||
|
||||
// Split-KV: compute number of splits to fill all SMs for small-batch decode.
|
||||
// Caps splits so each split processes at least `min_tiles_per_split` tiles,
|
||||
// avoiding excessive loop/prologue overhead when tiles are small.
|
||||
//
|
||||
// Target total grid blocks (`TARGET_BLOCKS`) rather than scaling splits by SM
|
||||
// count. Decode blocks are single-warp (32 threads) and a SM hosts ~11 of
|
||||
// them, so the old `2*sm/base` cap badly undersplit at large batch (B=16 got
|
||||
// 3 splits, optimal ~8). Measured (L20, grid search): bandwidth saturates
|
||||
// near 256-512 total blocks; 512 minimizes worst-case latency across the
|
||||
// B x kv grid; more is pure oversplit overhead.
|
||||
constexpr int DECODE_TARGET_BLOCKS = 512;
|
||||
inline int compute_num_splits(int base_blocks, int tiles_total,
|
||||
int min_tiles_per_split = 1) {
|
||||
int n = (DECODE_TARGET_BLOCKS + base_blocks - 1) / base_blocks;
|
||||
int max_by_work = tiles_total / min_tiles_per_split;
|
||||
return std::max(1, std::min(n, std::min(max_by_work, MAX_SPLITS)));
|
||||
}
|
||||
|
||||
// Dispatch IsCausal × HasMask — eliminates the duplicated 4-way if/else
|
||||
// ladder that appeared in each dispatch_* function. FN must be a function
|
||||
// template <int HEAD_DIM, bool IsCausal, bool HasMask>; HEAD_DIM is forwarded
|
||||
// as the first template argument so callers only spell it once.
|
||||
//
|
||||
// Usage: DISPATCH_CAUSAL_MASK(is_causal, has_mask, launcher<KV>::template launch, HEAD_DIM, p, stream);
|
||||
#define DISPATCH_CAUSAL_MASK(is_causal, has_mask, FN, HEAD_DIM, ...) \
|
||||
do { \
|
||||
if (is_causal) { \
|
||||
if (has_mask) FN<HEAD_DIM, true, true>(__VA_ARGS__); \
|
||||
else FN<HEAD_DIM, true, false>(__VA_ARGS__); \
|
||||
} else { \
|
||||
if (has_mask) FN<HEAD_DIM, false, true>(__VA_ARGS__); \
|
||||
else FN<HEAD_DIM, false, false>(__VA_ARGS__); \
|
||||
} \
|
||||
} while (0)
|
||||
|
||||
// ======================================================================
|
||||
// Prefill launchers (KV selects ContigKV or PagedKV addressing)
|
||||
// ======================================================================
|
||||
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
template <int BC_>
|
||||
struct PrefillKernelConfig {
|
||||
static constexpr int BC = BC_;
|
||||
static constexpr int WARPS = 4;
|
||||
static constexpr int STAGES = 2;
|
||||
};
|
||||
|
||||
// Compile-time configuration map shared by contiguous and paged prefill.
|
||||
// Unsupported head dimensions intentionally have no mapping.
|
||||
template <int HEAD_DIM, bool IsCausal>
|
||||
struct PrefillConfigMap;
|
||||
|
||||
template <> struct PrefillConfigMap<32, false> : PrefillKernelConfig<32> {};
|
||||
template <> struct PrefillConfigMap<32, true> : PrefillKernelConfig<64> {};
|
||||
template <> struct PrefillConfigMap<64, false> : PrefillKernelConfig<32> {};
|
||||
template <> struct PrefillConfigMap<64, true> : PrefillKernelConfig<64> {};
|
||||
template <> struct PrefillConfigMap<128, false> : PrefillKernelConfig<32> {};
|
||||
template <> struct PrefillConfigMap<128, true> : PrefillKernelConfig<32> {};
|
||||
template <> struct PrefillConfigMap<256, false> : PrefillKernelConfig<16> {};
|
||||
template <> struct PrefillConfigMap<256, true> : PrefillKernelConfig<16> {};
|
||||
|
||||
template <typename QSchedule, typename KV>
|
||||
struct PrefillLauncherMMA {
|
||||
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
||||
static void launch(AttentionParams<bf16>& p, cudaStream_t stream) {
|
||||
using Config = PrefillConfigMap<HEAD_DIM, IsCausal>;
|
||||
using Traits = KernelTraits<HEAD_DIM, Config::BC, Config::WARPS, Config::STAGES>;
|
||||
constexpr int ROWS = Traits::BR * Config::WARPS;
|
||||
dim3 grid(QSchedule::host_q_blocks(p, ROWS), p.q_head,
|
||||
QSchedule::host_grid_batch(p));
|
||||
dim3 block(Traits::NUM_THREADS);
|
||||
attn_prefill_split_q_mma_kernel<Traits, QSchedule, KV, IsCausal, HasMask>
|
||||
<<<grid, block, 0, stream>>>(p);
|
||||
}
|
||||
};
|
||||
#endif
|
||||
|
||||
template <typename QSchedule, typename KV>
|
||||
struct PrefillLauncherScalar {
|
||||
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
||||
static void launch(AttentionParams<bf16>& p, cudaStream_t stream) {
|
||||
constexpr int G = (HEAD_DIM == 32) ? 4 : 8, ROWS = 64, P_BC = 32;
|
||||
dim3 grid(QSchedule::host_q_blocks(p, ROWS), p.q_head,
|
||||
QSchedule::host_grid_batch(p));
|
||||
dim3 block(G, ROWS);
|
||||
attn_prefill_split_q_kernel_t<HEAD_DIM, QSchedule, KV, G, ROWS, P_BC,
|
||||
IsCausal, HasMask>
|
||||
<<<grid, block, 0, stream>>>(p);
|
||||
}
|
||||
};
|
||||
|
||||
template <int HEAD_DIM>
|
||||
static inline void dispatch_prefill(AttentionParams<bf16>& p, cudaStream_t stream) {
|
||||
bool is_causal = (p.causal_offset >= 0);
|
||||
bool has_mask = (p.use_mask && p.mask);
|
||||
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
using Launcher = PrefillLauncherMMA<DenseQSchedule, ContigKV>;
|
||||
DISPATCH_CAUSAL_MASK(is_causal, has_mask,
|
||||
Launcher::template launch,
|
||||
HEAD_DIM, p, stream);
|
||||
#else
|
||||
using Launcher = PrefillLauncherScalar<DenseQSchedule, ContigKV>;
|
||||
DISPATCH_CAUSAL_MASK(is_causal, has_mask,
|
||||
Launcher::template launch,
|
||||
HEAD_DIM, p, stream);
|
||||
#endif
|
||||
}
|
||||
|
||||
template <int HEAD_DIM>
|
||||
static inline void dispatch_paged_prefill(AttentionParams<bf16>& p, cudaStream_t stream) {
|
||||
bool is_causal = (p.causal_offset >= 0);
|
||||
bool has_mask = (p.use_mask && p.mask);
|
||||
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
using Launcher = PrefillLauncherMMA<PackedQSchedule, PagedKV>;
|
||||
DISPATCH_CAUSAL_MASK(is_causal, has_mask,
|
||||
Launcher::template launch,
|
||||
HEAD_DIM, p, stream);
|
||||
#else
|
||||
using Launcher = PrefillLauncherScalar<PackedQSchedule, PagedKV>;
|
||||
DISPATCH_CAUSAL_MASK(is_causal, has_mask,
|
||||
Launcher::template launch,
|
||||
HEAD_DIM, p, stream);
|
||||
#endif
|
||||
}
|
||||
|
||||
// ======================================================================
|
||||
// Decode launchers (KV selects ContigKV or PagedKV addressing)
|
||||
// ======================================================================
|
||||
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
// BC=16: halves smem (16KB vs 32KB) → doubles occupancy (6 vs 3 blocks/SM).
|
||||
// For D=256, BC=16 also reduces register pressure (fewer Sacc/PV frags),
|
||||
// enabling STAGES=2 (double-buffer) within the 32KB smem budget — eliminates
|
||||
// the 176-byte spill that STAGES=1+BC=32 suffered.
|
||||
template <typename KV>
|
||||
struct DecodeLauncherMMA {
|
||||
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
||||
static void launch(AttentionParams<bf16>& p, cudaStream_t stream) {
|
||||
int G = p.q_head / p.kv_head;
|
||||
constexpr int MAX_G = 16;
|
||||
int num_passes = (G + MAX_G - 1) / MAX_G;
|
||||
constexpr int BC = 16;
|
||||
int kv_len = KV::host_kv_len(p);
|
||||
int tiles_total = (kv_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);
|
||||
attn_decode_split_kv_mma_kernel<Traits, KV, IsCausal, HasMask>
|
||||
<<<grid, 32, 0, stream>>>(p);
|
||||
}
|
||||
};
|
||||
#endif
|
||||
|
||||
template <typename KV>
|
||||
struct DecodeLauncherScalar {
|
||||
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
||||
static void launch(AttentionParams<bf16>& p, cudaStream_t stream) {
|
||||
int kv_len = KV::host_kv_len(p);
|
||||
int chunks_total = (kv_len + DC_CHUNK - 1) / DC_CHUNK;
|
||||
p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total);
|
||||
size_t smem = 2 * DC_CHUNK * p.head_dim * sizeof(bf16);
|
||||
int group_size = p.q_head / p.kv_head;
|
||||
int g = min(group_size, 32); // cap at 32 to respect 1024-thread limit
|
||||
dim3 grid(p.batch * p.kv_head, 1, p.num_splits);
|
||||
dim3 block(32, g);
|
||||
cudaFuncSetAttribute(
|
||||
attn_decode_split_kv_kernel<HEAD_DIM, KV, IsCausal, HasMask>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
smem);
|
||||
attn_decode_split_kv_kernel<HEAD_DIM, KV, IsCausal, HasMask>
|
||||
<<<grid, block, smem, stream>>>(p);
|
||||
}
|
||||
};
|
||||
|
||||
template <int HEAD_DIM>
|
||||
static inline void dispatch_decode(AttentionParams<bf16>& p, cudaStream_t stream) {
|
||||
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,
|
||||
DecodeLauncherMMA<ContigKV>::template launch,
|
||||
HEAD_DIM, p, stream);
|
||||
#else
|
||||
DISPATCH_CAUSAL_MASK(is_causal, has_mask,
|
||||
DecodeLauncherScalar<ContigKV>::template launch,
|
||||
HEAD_DIM, p, stream);
|
||||
#endif
|
||||
|
||||
attn_decode_combine_kernel<ContigKV><<<p.batch * p.q_head, p.head_dim, 0, stream>>>(p);
|
||||
}
|
||||
|
||||
template <int HEAD_DIM>
|
||||
static inline void dispatch_paged_decode(AttentionParams<bf16>& p, cudaStream_t stream) {
|
||||
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,
|
||||
DecodeLauncherMMA<PagedKV>::template launch,
|
||||
HEAD_DIM, p, stream);
|
||||
#else
|
||||
DISPATCH_CAUSAL_MASK(is_causal, has_mask,
|
||||
DecodeLauncherScalar<PagedKV>::template launch,
|
||||
HEAD_DIM, p, stream);
|
||||
#endif
|
||||
|
||||
attn_decode_combine_kernel<PagedKV><<<p.batch * p.q_head, p.head_dim, 0, stream>>>(p);
|
||||
}
|
||||
@@ -0,0 +1,358 @@
|
||||
#pragma once
|
||||
#include <float.h>
|
||||
#include <torch/extension.h>
|
||||
#include <c10/cuda/CUDAGuard.h>
|
||||
#include "common.h"
|
||||
#include "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_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 ----
|
||||
// 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_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_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_l_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_l_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(p.q_head % p.kv_head == 0,
|
||||
"q_head must be divisible by kv_head");
|
||||
TORCH_CHECK(k.size(3) == p.head_dim, "K/V head_dim must match Q");
|
||||
TORCH_CHECK(q.stride(3) == 1 && k.stride(3) == 1 && v.stride(3) == 1,
|
||||
"Q/K/V head_dim must be contiguous");
|
||||
|
||||
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_ptr = (const T*)q.data_ptr();
|
||||
p.k_ptr = (const T*)k.data_ptr();
|
||||
p.v_ptr = (const T*)v.data_ptr();
|
||||
p.new_k_ptr = nullptr;
|
||||
p.new_v_ptr = nullptr;
|
||||
p.o_ptr = 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,
|
||||
const c10::optional<torch::Tensor>& new_k,
|
||||
const c10::optional<torch::Tensor>& new_v,
|
||||
c10::optional<torch::Tensor> mask,
|
||||
int64_t causal_offset,
|
||||
double scale,
|
||||
AttentionParams<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::kInt32, "req_to_token must be int32");
|
||||
TORCH_CHECK(req_pool_indices.dtype() == torch::kInt32,
|
||||
"req_pool_indices must be int32");
|
||||
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(q.stride(2) == 1 && k_cache.stride(2) == 1 && v_cache.stride(2) == 1,
|
||||
"Q/K/V head_dim must be contiguous");
|
||||
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_l_stride = (int)q.stride(0);
|
||||
p.q_h_stride = (int)q.stride(1);
|
||||
p.q_d_stride = (int)q.stride(2);
|
||||
|
||||
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<int>();
|
||||
p.req_pool_indices = req_pool_indices.data_ptr<int>();
|
||||
p.kv_indptr = kv_indptr.data_ptr<int>();
|
||||
p.qo_indptr = nullptr;
|
||||
p.max_context_len = (int)req_to_token.size(1);
|
||||
|
||||
TORCH_CHECK(new_k.has_value() == new_v.has_value(),
|
||||
"new_k and new_v must be provided together");
|
||||
if (new_k.has_value()) {
|
||||
auto nk = new_k.value();
|
||||
auto nv = new_v.value();
|
||||
TORCH_CHECK(nk.is_cuda() && nv.is_cuda(), "new K/V must be CUDA tensors");
|
||||
TORCH_CHECK(nk.dtype() == torch::kBFloat16 && nv.dtype() == torch::kBFloat16,
|
||||
"new K/V must be bf16");
|
||||
TORCH_CHECK(nk.dim() == 3 && nv.dim() == 3,
|
||||
"new K/V must be 3D [batch, kv_head, head_dim]");
|
||||
TORCH_CHECK(nk.sizes() == nv.sizes(), "new K and V must have identical shapes");
|
||||
TORCH_CHECK(nk.strides() == nv.strides(),
|
||||
"new K and V must have identical strides");
|
||||
TORCH_CHECK(nk.size(0) == p.batch && nk.size(1) == p.kv_head
|
||||
&& nk.size(2) == p.head_dim, "new K/V shape mismatch");
|
||||
TORCH_CHECK(nk.stride(2) == 1 && nv.stride(2) == 1,
|
||||
"new K/V head_dim must be contiguous");
|
||||
p.new_k_ptr = (const T*)nk.data_ptr();
|
||||
p.new_v_ptr = (const T*)nv.data_ptr();
|
||||
p.new_kv_b_stride = (int)nk.stride(0);
|
||||
p.new_kv_h_stride = (int)nk.stride(1);
|
||||
} else {
|
||||
p.new_k_ptr = nullptr;
|
||||
p.new_v_ptr = nullptr;
|
||||
p.new_kv_b_stride = p.new_kv_h_stride = 0;
|
||||
}
|
||||
|
||||
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_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_l_stride = 0;
|
||||
}
|
||||
|
||||
p.o_ptr = 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,
|
||||
torch::Tensor q_tile_to_batch,
|
||||
torch::Tensor q_tile_to_index,
|
||||
c10::optional<torch::Tensor> mask,
|
||||
int64_t causal_offset,
|
||||
double scale,
|
||||
AttentionParams<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_tile_to_batch.is_cuda() && q_tile_to_index.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::kInt32, "req_to_token must be int32");
|
||||
TORCH_CHECK(req_pool_indices.dtype() == torch::kInt32,
|
||||
"req_pool_indices must be int32");
|
||||
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(q_tile_to_batch.dtype() == torch::kInt32,
|
||||
"q_tile_to_batch must be int32");
|
||||
TORCH_CHECK(q_tile_to_index.dtype() == torch::kInt32,
|
||||
"q_tile_to_index 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.q_len = (int)q.size(0);
|
||||
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(q.stride(2) == 1 && k_cache.stride(2) == 1 && v_cache.stride(2) == 1,
|
||||
"Q/K/V head_dim must be contiguous");
|
||||
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]");
|
||||
TORCH_CHECK(q_tile_to_batch.dim() == 1 && q_tile_to_index.dim() == 1,
|
||||
"Q tile mappings must be 1D");
|
||||
TORCH_CHECK(q_tile_to_batch.size(0) == q_tile_to_index.size(0),
|
||||
"Q tile mappings must have equal length");
|
||||
|
||||
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_ptr = (const T*)k_cache.data_ptr();
|
||||
p.v_ptr = (const T*)v_cache.data_ptr();
|
||||
p.new_k_ptr = nullptr;
|
||||
p.new_v_ptr = nullptr;
|
||||
p.q_ptr = (const T*)q.data_ptr();
|
||||
p.req_to_token = req_to_token.data_ptr<int>();
|
||||
p.req_pool_indices = req_pool_indices.data_ptr<int>();
|
||||
p.kv_indptr = kv_indptr.data_ptr<int>();
|
||||
p.qo_indptr = qo_indptr.data_ptr<int>();
|
||||
p.q_tile_to_batch = q_tile_to_batch.data_ptr<int>();
|
||||
p.q_tile_to_index = q_tile_to_index.data_ptr<int>();
|
||||
p.num_q_tiles = (int)q_tile_to_batch.size(0);
|
||||
p.max_context_len = (int)req_to_token.size(1);
|
||||
|
||||
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_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) > 0 && m.size(2) <= p.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_l_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_l_stride = 0;
|
||||
}
|
||||
p.scale = (scale > 0.0) ? (float)scale : 1.0f / sqrtf((float)p.head_dim);
|
||||
|
||||
p.o_ptr = nullptr;
|
||||
p.o_part = nullptr;
|
||||
p.ml_part = nullptr;
|
||||
}
|
||||
@@ -0,0 +1,255 @@
|
||||
#pragma once
|
||||
#include <cuda_bf16.h>
|
||||
#include "common.h"
|
||||
|
||||
// ============================================================================
|
||||
// Attention layout policies keep Q scheduling independent from K/V storage.
|
||||
// DenseQSchedule / PackedQSchedule map blocks to Q tiles; ContigKV / PagedKV
|
||||
// resolve logical K/V positions to physical addresses. This lets the shared
|
||||
// kernels compose Q layout and K/V storage without coupling the two concerns.
|
||||
//
|
||||
// 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_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_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
|
||||
// to kv_addr, so the load loops never redo the hoistable base computation
|
||||
// (e.g. the req_pool_indices global read) element-by-element.
|
||||
// ============================================================================
|
||||
|
||||
#define HOST_FORCEINLINE static __host__ __forceinline__
|
||||
#define DEVICE_FORCEINLINE static __device__ __forceinline__
|
||||
#define HOST_DEV_FORCEINLINE static __host__ __device__ __forceinline__
|
||||
|
||||
using bf16 = __nv_bfloat16;
|
||||
|
||||
// ============================================================================
|
||||
// Q scheduling policies
|
||||
//
|
||||
// Map CUDA blocks to request-local Q tiles independently of K/V storage.
|
||||
// Dense tensors encode the request in blockIdx.z; packed ragged tensors use
|
||||
// a compact precomputed work map indexed by blockIdx.x.
|
||||
// ============================================================================
|
||||
|
||||
struct DenseQSchedule {
|
||||
HOST_FORCEINLINE int host_q_blocks(
|
||||
const AttentionParams<bf16>& p, int rows) {
|
||||
return (p.q_len + rows - 1) / rows;
|
||||
}
|
||||
|
||||
HOST_FORCEINLINE int host_grid_batch(
|
||||
const AttentionParams<bf16>& p) {
|
||||
return p.batch;
|
||||
}
|
||||
|
||||
DEVICE_FORCEINLINE void map_block(
|
||||
const AttentionParams<bf16>&, int& batch, int& q_tile) {
|
||||
batch = blockIdx.z;
|
||||
q_tile = blockIdx.x;
|
||||
}
|
||||
|
||||
DEVICE_FORCEINLINE int q_len(
|
||||
const AttentionParams<bf16>& p, int) {
|
||||
return p.q_len;
|
||||
}
|
||||
|
||||
DEVICE_FORCEINLINE int q_base(
|
||||
const AttentionParams<bf16>& p, int batch, int q_head) {
|
||||
return batch * p.q_b_stride + q_head * p.q_h_stride;
|
||||
}
|
||||
};
|
||||
|
||||
struct PackedQSchedule {
|
||||
HOST_FORCEINLINE int host_q_blocks(
|
||||
const AttentionParams<bf16>& p, int) {
|
||||
return p.num_q_tiles;
|
||||
}
|
||||
|
||||
HOST_FORCEINLINE int host_grid_batch(
|
||||
const AttentionParams<bf16>&) {
|
||||
return 1;
|
||||
}
|
||||
|
||||
DEVICE_FORCEINLINE void map_block(
|
||||
const AttentionParams<bf16>& p, int& batch, int& q_tile) {
|
||||
batch = p.q_tile_to_batch[blockIdx.x];
|
||||
q_tile = p.q_tile_to_index[blockIdx.x];
|
||||
}
|
||||
|
||||
DEVICE_FORCEINLINE int q_len(
|
||||
const AttentionParams<bf16>& p, int batch) {
|
||||
return p.qo_indptr[batch + 1] - p.qo_indptr[batch];
|
||||
}
|
||||
|
||||
DEVICE_FORCEINLINE int q_base(
|
||||
const AttentionParams<bf16>& p, int batch, int q_head) {
|
||||
return p.qo_indptr[batch] * p.q_l_stride + q_head * p.q_h_stride;
|
||||
}
|
||||
};
|
||||
|
||||
// Hoisted per-(batch, kv_head) addressing context.
|
||||
struct KVContext {
|
||||
int kv_base; // contig: batch*kv_b_stride + kv_head*kv_h_stride
|
||||
int req_idx; // paged: req_pool_indices[batch]
|
||||
int64_t rtt_stride; // paged: max_context_len
|
||||
int64_t pool_stride; // paged: kv_head * HEAD_DIM
|
||||
int64_t head_off; // paged: kv_head * HEAD_DIM
|
||||
};
|
||||
|
||||
// Per-element K/V global addresses for one (kc, d) position of a K/V tile.
|
||||
// The pointers are ALWAYS the computed addresses (never nullptr) — callers
|
||||
// gate on `valid` (cp.async src_size=0, or a guarded scalar deref). `valid`
|
||||
// starts as "within the request's seq_len"; the paged policy further degrades
|
||||
// it when req_to_token maps the position to a negative slot (empty padding).
|
||||
// This matches the original hand-rolled load loops, where the address was
|
||||
// always formed and the predicate decided whether anything was read.
|
||||
struct KVAddr {
|
||||
const void* k;
|
||||
const void* v;
|
||||
bool valid;
|
||||
};
|
||||
|
||||
// ---- Contiguous K/V ----
|
||||
struct ContigKV {
|
||||
static constexpr bool kPaged = false;
|
||||
|
||||
HOST_FORCEINLINE int host_kv_len(const AttentionParams<bf16>& p) {
|
||||
return p.kv_len;
|
||||
}
|
||||
|
||||
// decode: same offset (q_len == 1, so there is no row stride component)
|
||||
DEVICE_FORCEINLINE int q_decode_base(
|
||||
const AttentionParams<bf16>& p, int batch, int q_head) {
|
||||
return batch * p.q_b_stride + q_head * p.q_h_stride;
|
||||
}
|
||||
|
||||
DEVICE_FORCEINLINE int kv_len(const AttentionParams<bf16>& p, int) {
|
||||
return p.kv_len;
|
||||
}
|
||||
DEVICE_FORCEINLINE int causal_offset(
|
||||
const AttentionParams<bf16>& p, int, int) {
|
||||
return p.causal_offset;
|
||||
}
|
||||
// decode: exclusive bound of the single query's attend range
|
||||
DEVICE_FORCEINLINE int decode_attend_len(const AttentionParams<bf16>& p, int) {
|
||||
return (p.kv_len < p.causal_offset + 1) ? p.kv_len : (p.causal_offset + 1);
|
||||
}
|
||||
|
||||
template <int HEAD_DIM>
|
||||
DEVICE_FORCEINLINE KVContext make_ctx(
|
||||
const AttentionParams<bf16>& p, int batch, int kv_head) {
|
||||
KVContext c = {};
|
||||
c.kv_base = batch * p.kv_b_stride + kv_head * p.kv_h_stride;
|
||||
return c;
|
||||
}
|
||||
DEVICE_FORCEINLINE int resolve_token(
|
||||
const AttentionParams<bf16>& p, const KVContext& c, int kc, bool valid) {
|
||||
return valid ? kc : -1;
|
||||
}
|
||||
DEVICE_FORCEINLINE KVAddr kv_addr_from_token(
|
||||
const AttentionParams<bf16>& p, const KVContext& c, int token, int d) {
|
||||
const bool valid = token >= 0;
|
||||
const int safe_token = valid ? token : 0;
|
||||
const int64_t gmem_off = (int64_t)c.kv_base
|
||||
+ (int64_t)safe_token * p.kv_l_stride
|
||||
+ (int64_t)d * p.kv_d_stride;
|
||||
return {&p.k_ptr[gmem_off], &p.v_ptr[gmem_off], valid};
|
||||
}
|
||||
|
||||
template <int VEC>
|
||||
DEVICE_FORCEINLINE KVAddr decode_addr(
|
||||
const AttentionParams<bf16>& p, const KVContext& c,
|
||||
int, int, int kc, int d, bool valid, bool) {
|
||||
int token = resolve_token(p, c, kc, valid);
|
||||
return kv_addr_from_token(p, c, token, d);
|
||||
}
|
||||
};
|
||||
|
||||
// ---- Paged (SGLang-style flat pool) K/V ----
|
||||
struct PagedKV {
|
||||
static constexpr bool kPaged = true;
|
||||
|
||||
HOST_FORCEINLINE int host_kv_len(const AttentionParams<bf16>& p) {
|
||||
return p.max_context_len;
|
||||
}
|
||||
|
||||
// decode: Q is [batch, q_head, head_dim], so batch is the outer row
|
||||
DEVICE_FORCEINLINE int q_decode_base(
|
||||
const AttentionParams<bf16>& p, int batch, int q_head) {
|
||||
return batch * p.q_l_stride + q_head * p.q_h_stride;
|
||||
}
|
||||
|
||||
DEVICE_FORCEINLINE int kv_len(const AttentionParams<bf16>& p, int batch) {
|
||||
return p.kv_indptr[batch + 1] - p.kv_indptr[batch];
|
||||
}
|
||||
DEVICE_FORCEINLINE int causal_offset(
|
||||
const AttentionParams<bf16>& p, int batch, int q_len) {
|
||||
return kv_len(p, batch) - q_len;
|
||||
}
|
||||
// decode: the query is the last token, so [0, seq_len) IS its causal range
|
||||
DEVICE_FORCEINLINE int decode_attend_len(const AttentionParams<bf16>& p, int batch) {
|
||||
return kv_len(p, batch);
|
||||
}
|
||||
|
||||
template <int HEAD_DIM>
|
||||
DEVICE_FORCEINLINE KVContext make_ctx(
|
||||
const AttentionParams<bf16>& p, int batch, int kv_head) {
|
||||
KVContext c = {};
|
||||
c.req_idx = p.req_pool_indices[batch];
|
||||
c.rtt_stride = (int64_t)p.max_context_len;
|
||||
c.pool_stride = (int64_t)p.kv_head * HEAD_DIM;
|
||||
c.head_off = (int64_t)kv_head * HEAD_DIM;
|
||||
return c;
|
||||
}
|
||||
DEVICE_FORCEINLINE int resolve_token(
|
||||
const AttentionParams<bf16>& p, const KVContext& c, int kc, bool valid) {
|
||||
return valid ? p.req_to_token[c.req_idx * c.rtt_stride + kc] : -1;
|
||||
}
|
||||
DEVICE_FORCEINLINE KVAddr kv_addr_from_token(
|
||||
const AttentionParams<bf16>& p, const KVContext& c, int slot, int d) {
|
||||
const bool valid = slot >= 0;
|
||||
const int safe_slot = valid ? slot : 0;
|
||||
const int64_t gmem_off = (int64_t)safe_slot * c.pool_stride + c.head_off + d;
|
||||
return {&p.k_ptr[gmem_off], &p.v_ptr[gmem_off], valid};
|
||||
}
|
||||
|
||||
DEVICE_FORCEINLINE KVAddr new_kv_addr(
|
||||
const AttentionParams<bf16>& p, int batch, int kv_head, int d) {
|
||||
const int64_t off = (int64_t)batch * p.new_kv_b_stride
|
||||
+ (int64_t)kv_head * p.new_kv_h_stride + d;
|
||||
return {&p.new_k_ptr[off], &p.new_v_ptr[off], true};
|
||||
}
|
||||
|
||||
DEVICE_FORCEINLINE void store_new_kv(
|
||||
const AttentionParams<bf16>& p, const KVContext& c,
|
||||
int seq_len, int d, const KVAddr& src) {
|
||||
int slot = resolve_token(p, c, seq_len - 1, true);
|
||||
const int64_t off = (int64_t)slot * c.pool_stride + c.head_off + d;
|
||||
const_cast<bf16*>(p.k_ptr)[off] = *reinterpret_cast<const bf16*>(src.k);
|
||||
const_cast<bf16*>(p.v_ptr)[off] = *reinterpret_cast<const bf16*>(src.v);
|
||||
}
|
||||
|
||||
template <int VEC>
|
||||
DEVICE_FORCEINLINE KVAddr decode_addr(
|
||||
const AttentionParams<bf16>& p, const KVContext& c,
|
||||
int batch, int kv_head, int kc, int d, bool valid, bool persist) {
|
||||
if (p.new_k_ptr && valid && kc == kv_len(p, batch) - 1) {
|
||||
KVAddr src = new_kv_addr(p, batch, kv_head, d);
|
||||
if (persist) {
|
||||
#pragma unroll
|
||||
for (int j = 0; j < VEC; j++) {
|
||||
KVAddr value = new_kv_addr(p, batch, kv_head, d + j);
|
||||
store_new_kv(p, c, kc + 1, d + j, value);
|
||||
}
|
||||
}
|
||||
return src;
|
||||
}
|
||||
int token = resolve_token(p, c, kc, valid);
|
||||
return kv_addr_from_token(p, c, token, d);
|
||||
}
|
||||
};
|
||||
@@ -0,0 +1,274 @@
|
||||
#pragma once
|
||||
#include <cfloat>
|
||||
#include <cuda_fp16.h>
|
||||
#include <cuda_runtime.h>
|
||||
|
||||
#include "../common/mma.cuh"
|
||||
|
||||
// Predicated cp.async (4-operand form) requires CUDA 11.2+.
|
||||
// bf16 mma.sync requires sm_80+ (guarded at build time by ASTRAI_NO_MMA).
|
||||
#if CUDART_VERSION < 11020
|
||||
#error "AstrAI CUDA kernels require CUDA 11.2 or later (CUDART_VERSION >= 11020)."
|
||||
#endif
|
||||
|
||||
// ============================================================================
|
||||
// KernelTraits — FlashAttention-v2 style compile-time configuration bundle.
|
||||
//
|
||||
// Bundles all dimension-dependent constants so device functions only need a
|
||||
// single Traits template parameter rather than scattered <KD, NC8, KT2, ...>.
|
||||
// ============================================================================
|
||||
template <int HEAD_DIM_, int BC_, int WARPS_, int STAGES_>
|
||||
struct KernelTraits {
|
||||
static constexpr int HEAD_DIM = HEAD_DIM_;
|
||||
static constexpr int BC = BC_; // K/V tile size along seq dim
|
||||
static constexpr int WARPS = WARPS_; // warps per block
|
||||
static constexpr int STAGES = STAGES_; // double-buffer stages (1 or 2)
|
||||
|
||||
static constexpr int BR = 16; // Q rows per warp (mma M=16)
|
||||
|
||||
// Derived: mma tile counts from the shared mma_shape (m16n8k16 for bf16)
|
||||
static constexpr int KD = HEAD_DIM / astrai::mma_shape<bf16>::k; // Q/K k-slides
|
||||
static constexpr int NC8 = BC / 8; // S n-tiles (N=8)
|
||||
static constexpr int KT2 = BC / astrai::mma_shape<bf16>::k; // P k-tiles (K=16)
|
||||
static constexpr int DN8 = HEAD_DIM / 8; // O n-tiles (N=8)
|
||||
|
||||
static constexpr int LD = HEAD_DIM; // smem leading dim
|
||||
|
||||
// XOR swizzle chunk bits for ldmatrix bank-conflict avoidance.
|
||||
// mask = log2(LD/8) bits, clamped to stay within LD.
|
||||
static constexpr int SWIZ_MASK = (HEAD_DIM >= 64) ? 7 : (HEAD_DIM / 8 - 1);
|
||||
|
||||
static constexpr int NUM_THREADS = WARPS * 32;
|
||||
static constexpr int VEC = 8; // bf16 per cp.async unit (16 bytes)
|
||||
static constexpr int TOTAL = BC * HEAD_DIM; // total elements per tile
|
||||
};
|
||||
|
||||
// ---- PTX wrappers ----
|
||||
using bf16 = __nv_bfloat16;
|
||||
// bf16 mma.sync lives in the shared astrai::mma_sync template (common/mma.cuh).
|
||||
|
||||
// read two adjacent bf16 from smem as one packed .b32 (elem0 low, elem1 high)
|
||||
__device__ __forceinline__ unsigned ld2(const bf16* p) {
|
||||
return *reinterpret_cast<const unsigned*>(p);
|
||||
}
|
||||
|
||||
// pack two floats into one bf16x2 as .b32
|
||||
__device__ __forceinline__ unsigned pk2(float a, float b) {
|
||||
__nv_bfloat162 v = __floats2bfloat162_rn(a, b);
|
||||
return *reinterpret_cast<unsigned*>(&v);
|
||||
}
|
||||
|
||||
// pack two (non-contiguous) bf16 into one .b32
|
||||
__device__ __forceinline__ unsigned pkb(bf16 a, bf16 b) {
|
||||
__nv_bfloat162 v;
|
||||
v.x = a;
|
||||
v.y = b;
|
||||
return *reinterpret_cast<unsigned*>(&v);
|
||||
}
|
||||
|
||||
// ldmatrix lives in the shared template (common/mma.cuh):
|
||||
// `astrai::ldmatrix_x2<bf16>` / `<bf16, /*Trans=*/true>` load the K/V
|
||||
// fragments with the exact register layout mma expects.
|
||||
|
||||
// XOR swizzle for shared-memory column at 8-bf16 chunk granularity.
|
||||
__device__ __forceinline__ int swiz_col(int d, int r, int mask = 7) {
|
||||
return ((d >> 3) ^ (r & mask)) << 3 | (d & 7);
|
||||
}
|
||||
|
||||
// Predicated cp.async: copy 16 bytes when `pred`, otherwise zero-fill.
|
||||
// BypassL1 defaults to .cg (L2 only); false selects .ca (L1 + L2).
|
||||
// src_size=0 means no bytes are read, so an out-of-bounds address is safe.
|
||||
template <bool BypassL1 = true>
|
||||
__device__ __forceinline__ void cp_async_16_pred(bf16* smem_ptr,
|
||||
const void* gmem_ptr,
|
||||
bool pred) {
|
||||
unsigned smem_addr = __cvta_generic_to_shared(smem_ptr);
|
||||
int src_size = pred ? 16 : 0;
|
||||
if constexpr (BypassL1) {
|
||||
asm volatile("cp.async.cg.shared.global [%0], [%1], 16, %2;"
|
||||
:: "r"(smem_addr), "l"(gmem_ptr), "r"(src_size));
|
||||
} else {
|
||||
asm volatile("cp.async.ca.shared.global [%0], [%1], 16, %2;"
|
||||
:: "r"(smem_addr), "l"(gmem_ptr), "r"(src_size));
|
||||
}
|
||||
}
|
||||
|
||||
__device__ __forceinline__ void cp_async_commit() {
|
||||
asm volatile("cp.async.commit_group;");
|
||||
}
|
||||
|
||||
__device__ __forceinline__ void cp_async_wait_all() {
|
||||
asm volatile("cp.async.wait_all;");
|
||||
}
|
||||
|
||||
template <int N>
|
||||
__device__ __forceinline__ void cp_async_wait_group() {
|
||||
asm volatile("cp.async.wait_group %0;" :: "n"(N));
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// 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_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(
|
||||
const bf16* __restrict__ q,
|
||||
int stride_row,
|
||||
int stride_d,
|
||||
int qra, int qrb,
|
||||
bool va, bool vb,
|
||||
int tid4,
|
||||
unsigned Qa[KD][4])
|
||||
{
|
||||
#pragma unroll
|
||||
for (int kt = 0; kt < KD; kt++) {
|
||||
int c = kt * 16 + tid4 * 2;
|
||||
const unsigned* pau = reinterpret_cast<const unsigned*>(
|
||||
&q[qra * stride_row + c * stride_d]);
|
||||
const unsigned* pbu = reinterpret_cast<const unsigned*>(
|
||||
&q[qrb * stride_row + c * stride_d]);
|
||||
Qa[kt][0] = va ? pau[0] : 0u;
|
||||
Qa[kt][1] = vb ? pbu[0] : 0u;
|
||||
Qa[kt][2] = va ? pau[4] : 0u;
|
||||
Qa[kt][3] = vb ? pbu[4] : 0u;
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// S = Q @ K^T (Qa pre-loaded by the caller; scale applied post-mma in the
|
||||
// caller to avoid bf16 precision loss).
|
||||
// Traits provides KD, NC8, LD, and SWIZ_MASK.
|
||||
// ---------------------------------------------------------------------------
|
||||
template <typename Traits>
|
||||
__device__ inline void mma_compute_scores(
|
||||
const unsigned Qa[Traits::KD][4],
|
||||
const bf16* __restrict__ sK,
|
||||
int lane,
|
||||
float Sacc[Traits::NC8][4])
|
||||
{
|
||||
#pragma unroll
|
||||
for (int n8 = 0; n8 < Traits::NC8; n8++) {
|
||||
Sacc[n8][0] = Sacc[n8][1] = Sacc[n8][2] = Sacc[n8][3] = 0.0f;
|
||||
int krow_l = n8 * 8 + (lane & 7);
|
||||
int kcol_h = (lane & 8) ? 8 : 0;
|
||||
#pragma unroll
|
||||
for (int kt = 0; kt < Traits::KD; kt++) {
|
||||
unsigned b[2];
|
||||
astrai::ldmatrix_x2<bf16>(b, &sK[krow_l * Traits::LD
|
||||
+ swiz_col(kt * 16 + kcol_h, krow_l, Traits::SWIZ_MASK)]);
|
||||
astrai::mma_sync<bf16>(Sacc[n8], Qa[kt], b, Sacc[n8]);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Online softmax + Oacc rescale for one K/V tile.
|
||||
//
|
||||
// HasMask is a compile-time template bool: when false, the mask branch is
|
||||
// entirely dead-code-eliminated from the inner unrolled loop.
|
||||
// ---------------------------------------------------------------------------
|
||||
template <typename Traits, bool HasMask>
|
||||
__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_l_stride,
|
||||
int mask_batch, int mask_head0, int mask_head1,
|
||||
const bool* __restrict__ mask,
|
||||
bool valid0, bool valid1,
|
||||
float Sacc[Traits::NC8][4],
|
||||
float Oacc[Traits::DN8][4],
|
||||
float& m0, float& m1,
|
||||
float& l0, float& l1,
|
||||
int lane)
|
||||
{
|
||||
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_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;
|
||||
int c1 = cc + 1;
|
||||
bool b0 = !valid0 || (cc >= maxc0) || (HasMask && !mask[mask_base0 + cc]);
|
||||
bool b1 = !valid0 || (c1 >= maxc0) || (HasMask && !mask[mask_base0 + c1]);
|
||||
bool b2 = !valid1 || (cc >= maxc1) || (HasMask && !mask[mask_base1 + cc]);
|
||||
bool b3 = !valid1 || (c1 >= maxc1) || (HasMask && !mask[mask_base1 + c1]);
|
||||
float s0 = b0 ? -FLT_MAX : Sacc[n8][0];
|
||||
float s1 = b1 ? -FLT_MAX : Sacc[n8][1];
|
||||
float s2 = b2 ? -FLT_MAX : Sacc[n8][2];
|
||||
float s3 = b3 ? -FLT_MAX : Sacc[n8][3];
|
||||
Sacc[n8][0] = s0; Sacc[n8][1] = s1;
|
||||
Sacc[n8][2] = s2; Sacc[n8][3] = s3;
|
||||
rmax0 = fmaxf(rmax0, fmaxf(s0, s1));
|
||||
rmax1 = fmaxf(rmax1, fmaxf(s2, s3));
|
||||
}
|
||||
rmax0 = fmaxf(rmax0, __shfl_xor_sync(0xFFFFFFFF, rmax0, 1));
|
||||
rmax0 = fmaxf(rmax0, __shfl_xor_sync(0xFFFFFFFF, rmax0, 2));
|
||||
rmax1 = fmaxf(rmax1, __shfl_xor_sync(0xFFFFFFFF, rmax1, 1));
|
||||
rmax1 = fmaxf(rmax1, __shfl_xor_sync(0xFFFFFFFF, rmax1, 2));
|
||||
|
||||
float nm0 = fmaxf(m0, rmax0), nm1 = fmaxf(m1, rmax1);
|
||||
float corr0 = __expf(m0 - nm0);
|
||||
float corr1 = __expf(m1 - nm1);
|
||||
float pn0 = (nm0 == -FLT_MAX) ? 0.0f : 1.0f;
|
||||
float pn1 = (nm1 == -FLT_MAX) ? 0.0f : 1.0f;
|
||||
|
||||
float rsum0 = 0.0f, rsum1 = 0.0f;
|
||||
#pragma unroll
|
||||
for (int n8 = 0; n8 < Traits::NC8; n8++) {
|
||||
float p0 = pn0 * __expf(Sacc[n8][0] - nm0);
|
||||
float p1 = pn0 * __expf(Sacc[n8][1] - nm0);
|
||||
float p2 = pn1 * __expf(Sacc[n8][2] - nm1);
|
||||
float p3 = pn1 * __expf(Sacc[n8][3] - nm1);
|
||||
Sacc[n8][0] = p0; Sacc[n8][1] = p1;
|
||||
Sacc[n8][2] = p2; Sacc[n8][3] = p3;
|
||||
rsum0 += p0 + p1;
|
||||
rsum1 += p2 + p3;
|
||||
}
|
||||
rsum0 += __shfl_xor_sync(0xFFFFFFFF, rsum0, 1);
|
||||
rsum0 += __shfl_xor_sync(0xFFFFFFFF, rsum0, 2);
|
||||
rsum1 += __shfl_xor_sync(0xFFFFFFFF, rsum1, 1);
|
||||
rsum1 += __shfl_xor_sync(0xFFFFFFFF, rsum1, 2);
|
||||
l0 = l0 * corr0 + rsum0;
|
||||
l1 = l1 * corr1 + rsum1;
|
||||
m0 = nm0; m1 = nm1;
|
||||
|
||||
#pragma unroll
|
||||
for (int j = 0; j < Traits::DN8; j++) {
|
||||
Oacc[j][0] *= corr0; Oacc[j][1] *= corr0;
|
||||
Oacc[j][2] *= corr1; Oacc[j][3] *= corr1;
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// O += P @ V (Sacc must contain P = attention weights after softmax).
|
||||
// Traits provides DN8, KT2, LD, and SWIZ_MASK.
|
||||
// ---------------------------------------------------------------------------
|
||||
template <typename Traits>
|
||||
__device__ inline void mma_pv_accumulate(
|
||||
float Sacc[][4],
|
||||
const bf16* __restrict__ sV,
|
||||
int lane,
|
||||
float Oacc[Traits::DN8][4])
|
||||
{
|
||||
#pragma unroll
|
||||
for (int kt2 = 0; kt2 < Traits::KT2; kt2++) {
|
||||
unsigned Pa[4];
|
||||
Pa[0] = pk2(Sacc[kt2 * 2][0], Sacc[kt2 * 2][1]);
|
||||
Pa[1] = pk2(Sacc[kt2 * 2][2], Sacc[kt2 * 2][3]);
|
||||
Pa[2] = pk2(Sacc[kt2 * 2 + 1][0], Sacc[kt2 * 2 + 1][1]);
|
||||
Pa[3] = pk2(Sacc[kt2 * 2 + 1][2], Sacc[kt2 * 2 + 1][3]);
|
||||
int vrow_l = kt2 * 16 + (lane & 15);
|
||||
#pragma unroll
|
||||
for (int dn8 = 0; dn8 < Traits::DN8; dn8++) {
|
||||
unsigned b[2];
|
||||
astrai::ldmatrix_x2<bf16, true>(b, &sV[vrow_l * Traits::LD
|
||||
+ swiz_col(dn8 * 8, vrow_l, Traits::SWIZ_MASK)]);
|
||||
astrai::mma_sync<bf16>(Oacc[dn8], Pa, b, Oacc[dn8]);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,86 @@
|
||||
#include "dispatchers.cuh"
|
||||
#include "entry_utils.cuh"
|
||||
|
||||
torch::Tensor attn_paged_decode(
|
||||
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,
|
||||
c10::optional<torch::Tensor> new_k,
|
||||
c10::optional<torch::Tensor> new_v,
|
||||
c10::optional<torch::Tensor> mask,
|
||||
int64_t causal_offset,
|
||||
double scale,
|
||||
c10::optional<torch::Tensor> o_part_buf,
|
||||
c10::optional<torch::Tensor> ml_part_buf,
|
||||
c10::optional<torch::Tensor> out_buf
|
||||
) {
|
||||
const at::cuda::OptionalCUDAGuard device_guard(device_of(q));
|
||||
auto stream = at::cuda::getCurrentCUDAStream();
|
||||
|
||||
AttentionParams<bf16> p;
|
||||
attn_pack_paged_decode_params(q, k_cache, v_cache,
|
||||
req_to_token, req_pool_indices, kv_indptr,
|
||||
new_k, new_v,
|
||||
mask, causal_offset, scale, p);
|
||||
|
||||
torch::Tensor O;
|
||||
if (out_buf.has_value() && out_buf->defined()) {
|
||||
TORCH_CHECK(out_buf->dtype() == q.dtype(), "out_buf dtype must match q");
|
||||
TORCH_CHECK(out_buf->is_cuda() && out_buf->is_contiguous(),
|
||||
"out_buf must be a contiguous CUDA tensor");
|
||||
TORCH_CHECK(out_buf->size(0) >= q.size(0), "out_buf batch too small");
|
||||
TORCH_CHECK(out_buf->size(1) == q.size(1), "out_buf heads must match q");
|
||||
TORCH_CHECK(out_buf->size(2) == q.size(2), "out_buf head_dim must match q");
|
||||
TORCH_CHECK(q.is_contiguous(),
|
||||
"q must be contiguous when out_buf is provided");
|
||||
O = out_buf.value().slice(0, 0, q.size(0));
|
||||
} else {
|
||||
O = torch::empty({q.size(0), q.size(1), q.size(2)}, q.options());
|
||||
}
|
||||
p.o_ptr = (bf16*)O.data_ptr();
|
||||
|
||||
if (o_part_buf.has_value() && ml_part_buf.has_value()
|
||||
&& o_part_buf->defined() && ml_part_buf->defined()) {
|
||||
TORCH_CHECK(o_part_buf->scalar_type() == torch::kFloat32, "o_part_buf must be f32");
|
||||
TORCH_CHECK(ml_part_buf->scalar_type() == torch::kFloat32, "ml_part_buf must be f32");
|
||||
int64_t o_needed = (int64_t)p.batch * p.q_head * MAX_SPLITS * p.head_dim;
|
||||
int64_t ml_needed = (int64_t)p.batch * p.q_head * MAX_SPLITS * 2;
|
||||
TORCH_CHECK(o_part_buf->numel() >= o_needed,
|
||||
"o_part_buf too small: need ", o_needed, " got ", o_part_buf->numel());
|
||||
TORCH_CHECK(ml_part_buf->numel() >= ml_needed,
|
||||
"ml_part_buf too small: need ", ml_needed, " got ", ml_part_buf->numel());
|
||||
TORCH_CHECK(o_part_buf->is_cuda() && ml_part_buf->is_cuda(),
|
||||
"split buffers must be CUDA tensors");
|
||||
TORCH_CHECK(o_part_buf->is_contiguous() && ml_part_buf->is_contiguous(),
|
||||
"split buffers must be contiguous");
|
||||
p.o_part = (float*)o_part_buf->data_ptr();
|
||||
p.ml_part = (float*)ml_part_buf->data_ptr();
|
||||
} else {
|
||||
alloc_split_partials(p);
|
||||
}
|
||||
DISPATCH_HEAD_DIM(p.head_dim, dispatch_paged_decode, p, stream);
|
||||
C10_CUDA_CHECK(cudaGetLastError());
|
||||
return O;
|
||||
}
|
||||
|
||||
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
||||
m.def("attn_paged_decode", &attn_paged_decode,
|
||||
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("new_k") = py::none(),
|
||||
py::arg("new_v") = py::none(),
|
||||
py::arg("mask") = py::none(),
|
||||
py::arg("causal_offset") = -1,
|
||||
py::arg("scale") = 0.0,
|
||||
py::arg("o_part_buf") = py::none(),
|
||||
py::arg("ml_part_buf") = py::none(),
|
||||
py::arg("out_buf") = py::none(),
|
||||
"SGLang-style paged decode: flat KV pool + req_to_token + kv_indptr.");
|
||||
}
|
||||
@@ -0,0 +1,51 @@
|
||||
#include "dispatchers.cuh"
|
||||
#include "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,
|
||||
torch::Tensor q_tile_to_batch,
|
||||
torch::Tensor q_tile_to_index,
|
||||
c10::optional<torch::Tensor> mask,
|
||||
int64_t causal_offset,
|
||||
double scale
|
||||
) {
|
||||
const at::cuda::OptionalCUDAGuard device_guard(device_of(q));
|
||||
auto stream = at::cuda::getCurrentCUDAStream();
|
||||
|
||||
AttentionParams<bf16> p;
|
||||
attn_pack_paged_prefill_params(q, k_cache, v_cache,
|
||||
req_to_token, req_pool_indices,
|
||||
kv_indptr, qo_indptr,
|
||||
q_tile_to_batch, q_tile_to_index, mask,
|
||||
causal_offset, scale, p);
|
||||
|
||||
auto O = torch::empty({q.size(0), q.size(1), q.size(2)}, q.options());
|
||||
p.o_ptr = (bf16*)O.data_ptr();
|
||||
|
||||
DISPATCH_HEAD_DIM(p.head_dim, dispatch_paged_prefill, p, stream);
|
||||
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("q_tile_to_batch"),
|
||||
py::arg("q_tile_to_index"),
|
||||
py::arg("mask") = py::none(),
|
||||
py::arg("causal_offset") = -1,
|
||||
py::arg("scale") = 0.0,
|
||||
"SGLang-style paged prefill: flat KV pool + ragged batch.");
|
||||
}
|
||||
@@ -0,0 +1,39 @@
|
||||
#include "dispatchers.cuh"
|
||||
#include "entry_utils.cuh"
|
||||
|
||||
torch::Tensor attn_prefill(
|
||||
torch::Tensor q,
|
||||
torch::Tensor k,
|
||||
torch::Tensor v,
|
||||
c10::optional<torch::Tensor> mask,
|
||||
int64_t causal_offset,
|
||||
double scale,
|
||||
int64_t layout
|
||||
) {
|
||||
const at::cuda::OptionalCUDAGuard device_guard(device_of(q));
|
||||
auto stream = at::cuda::getCurrentCUDAStream();
|
||||
|
||||
AttentionParams<bf16> p;
|
||||
attn_pack_params(q, k, v, mask, causal_offset, scale, layout, p);
|
||||
TORCH_CHECK(p.head_dim % 16 == 0, "head_dim must be multiple of 16");
|
||||
|
||||
auto O = torch::empty_strided(q.sizes(), q.strides(), q.options());
|
||||
auto O_view = (layout == BLHD) ? O.transpose(1, 2) : O;
|
||||
p.o_ptr = (bf16*)O_view.data_ptr();
|
||||
|
||||
DISPATCH_HEAD_DIM(p.head_dim, dispatch_prefill, p, stream);
|
||||
C10_CUDA_CHECK(cudaGetLastError());
|
||||
return O;
|
||||
}
|
||||
|
||||
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
||||
m.def("attn_prefill", &attn_prefill,
|
||||
py::arg("q"),
|
||||
py::arg("k"),
|
||||
py::arg("v"),
|
||||
py::arg("mask") = py::none(),
|
||||
py::arg("causal_offset") = -1,
|
||||
py::arg("scale") = 0.0,
|
||||
py::arg("layout") = (int64_t)BHLD,
|
||||
"GQA prefill (tensor-core mma on sm_80+, scalar fallback)");
|
||||
}
|
||||
@@ -0,0 +1,157 @@
|
||||
#pragma once
|
||||
#include <cfloat>
|
||||
#include <cuda_bf16.h>
|
||||
#include "common.h"
|
||||
#include "layout_policies.cuh"
|
||||
|
||||
using bf16 = __nv_bfloat16;
|
||||
|
||||
// v9: group-split register blocking. G threads cooperate on one query row,
|
||||
// each owning HEAD_DIM/G dims of qreg[]/acc[]. IsCausal and HasMask are
|
||||
// compile-time bools — the compiler eliminates dead branches.
|
||||
// Unified across contiguous and paged (SGLang flat-pool) K/V via KV.
|
||||
// Templated on <HEAD_DIM, KV, G, ROWS, P_BC, IsCausal, HasMask>.
|
||||
|
||||
template <int G>
|
||||
__device__ __forceinline__ float group_reduce_sum(float v, unsigned mask) {
|
||||
#pragma unroll
|
||||
for (int o = G / 2; o > 0; o >>= 1)
|
||||
v += __shfl_xor_sync(mask, v, o);
|
||||
return v;
|
||||
}
|
||||
|
||||
// load 8 contiguous bf16 from (16-byte aligned) smem as one float4
|
||||
__device__ __forceinline__ void ld8(const bf16* p, float* o) {
|
||||
float4 raw = *reinterpret_cast<const float4*>(p);
|
||||
const __nv_bfloat162* h = reinterpret_cast<const __nv_bfloat162*>(&raw);
|
||||
#pragma unroll
|
||||
for (int j = 0; j < 4; j++) {
|
||||
float2 f = __bfloat1622float2(h[j]);
|
||||
o[2 * j] = f.x;
|
||||
o[2 * j + 1] = f.y;
|
||||
}
|
||||
}
|
||||
|
||||
template <int HEAD_DIM, typename QSchedule, typename KV, int G, int ROWS, int P_BC,
|
||||
bool IsCausal, bool HasMask>
|
||||
__global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
|
||||
constexpr int DPT = HEAD_DIM / G;
|
||||
|
||||
int batch, q_tile;
|
||||
QSchedule::map_block(p, batch, q_tile);
|
||||
|
||||
int q_head = blockIdx.y;
|
||||
int gpos = threadIdx.x; // 0..G-1 (which d-chunk)
|
||||
int row = threadIdx.y; // 0..ROWS-1
|
||||
int q_row = q_tile * ROWS + row;
|
||||
|
||||
// Per-request dims (from KV policy — paged reads kv_indptr/qo_indptr).
|
||||
const int seq_len = KV::kv_len(p, batch);
|
||||
const int q_len = QSchedule::q_len(p, batch);
|
||||
const int causal_off = KV::causal_offset(p, batch, q_len);
|
||||
const int kv_head = q_head / (p.q_head / p.kv_head);
|
||||
const KVContext kctx = KV::template make_ctx<HEAD_DIM>(p, batch, kv_head);
|
||||
|
||||
__shared__ __align__(16) bf16 sK[P_BC * HEAD_DIM];
|
||||
__shared__ __align__(16) bf16 sV[P_BC * HEAD_DIM];
|
||||
|
||||
// Q: stride-based load [batch, q_head, q_len, head_dim]
|
||||
const int q_base = QSchedule::q_base(p, batch, q_head);
|
||||
float qreg[DPT];
|
||||
if (q_row < q_len) {
|
||||
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_ptr[q_off + i * p.q_d_stride]);
|
||||
}
|
||||
|
||||
float m = -FLT_MAX, l = 0.0f;
|
||||
float acc[DPT];
|
||||
#pragma unroll
|
||||
for (int i = 0; i < DPT; i++)
|
||||
acc[i] = 0.0f;
|
||||
|
||||
int mask_batch_base = batch * p.mask_b_stride + q_head * p.mask_h_stride;
|
||||
int tiles = (seq_len + P_BC - 1) / P_BC;
|
||||
int tt = G * ROWS;
|
||||
int lid = row * G + gpos;
|
||||
|
||||
int lane_in_warp = lid & 31;
|
||||
unsigned gmask = (G == 32) ? 0xFFFFFFFFu
|
||||
: (((1u << G) - 1u) << (lane_in_warp & ~(G - 1)));
|
||||
|
||||
for (int ti = 0; ti < tiles; ti++) {
|
||||
int kv0 = ti * P_BC;
|
||||
int tlen = min(P_BC, seq_len - kv0);
|
||||
|
||||
// Load K/V into shared memory (addressing via KV policy; paged
|
||||
// guards empty slots with zero-fill).
|
||||
for (int i = lid; i < tlen * HEAD_DIM; i += tt) {
|
||||
int s = i / HEAD_DIM;
|
||||
int d_dim = i % HEAD_DIM;
|
||||
int kc = kv0 + s;
|
||||
int token = KV::resolve_token(p, kctx, kc, true);
|
||||
KVAddr a = KV::kv_addr_from_token(p, kctx, token, d_dim);
|
||||
sK[i] = a.valid ? *reinterpret_cast<const bf16*>(a.k) : (bf16)0.f;
|
||||
sV[i] = a.valid ? *reinterpret_cast<const bf16*>(a.v) : (bf16)0.f;
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
int lim = tlen;
|
||||
if constexpr (IsCausal) {
|
||||
if (q_row < q_len) {
|
||||
int ep = causal_off + q_row + 1;
|
||||
if (kv0 >= ep)
|
||||
lim = 0;
|
||||
else if (kv0 + tlen > ep)
|
||||
lim = ep - kv0;
|
||||
}
|
||||
}
|
||||
|
||||
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;
|
||||
#pragma unroll
|
||||
for (int i = 0; i < DPT; i += 8) {
|
||||
float k8[8];
|
||||
ld8(kr + i, k8);
|
||||
#pragma unroll
|
||||
for (int j = 0; j < 8; j++)
|
||||
part = fmaf(qreg[i + j], k8[j], part);
|
||||
}
|
||||
float dot = group_reduce_sum<G>(part, gmask) * p.scale;
|
||||
|
||||
int kv_idx = kv0 + s;
|
||||
if constexpr (HasMask) {
|
||||
if (!p.mask[mask_row_base + kv_idx])
|
||||
dot = -FLT_MAX;
|
||||
}
|
||||
|
||||
float nm = fmaxf(m, dot);
|
||||
float al = __expf(m - nm);
|
||||
float be = __expf(dot - nm);
|
||||
l = l * al + be;
|
||||
|
||||
const bf16* vr = sV + s * HEAD_DIM + gpos * DPT;
|
||||
#pragma unroll
|
||||
for (int i = 0; i < DPT; i += 8) {
|
||||
float v8[8];
|
||||
ld8(vr + i, v8);
|
||||
#pragma unroll
|
||||
for (int j = 0; j < 8; j++)
|
||||
acc[i + j] = fmaf(v8[j], be, acc[i + j] * al);
|
||||
}
|
||||
m = nm;
|
||||
}
|
||||
__syncthreads();
|
||||
}
|
||||
|
||||
if (q_row < q_len) {
|
||||
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_ptr[o_off + i * p.q_d_stride] = __float2bfloat16(acc[i] * rl);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,157 @@
|
||||
#pragma once
|
||||
#include <cfloat>
|
||||
#include <cuda_bf16.h>
|
||||
#include "common.h"
|
||||
#include "layout_policies.cuh"
|
||||
#include "mma_utils.cuh"
|
||||
|
||||
// Tensor-core prefill flash attention (raw mma.sync PTX), unified across
|
||||
// contiguous and paged (SGLang flat-pool) K/V via the KV template parameter.
|
||||
// One warp owns BR=16 query rows. S = Q@K^T and O = P@V run on bf16 tensor
|
||||
// cores via mma.sync.m16n8k16 (f32 accumulate).
|
||||
//
|
||||
// KV = ContigKV (dense [batch, kv_head, kv_len, head_dim]) or PagedKV
|
||||
// (flat pool + req_to_token, ragged batches via qo_indptr/kv_indptr).
|
||||
// IsCausal and HasMask are compile-time bools — the compiler eliminates all
|
||||
// dead branches in the inner compute loop (FA2-style).
|
||||
//
|
||||
// Traits = KernelTraits<HEAD_DIM, BC, WARPS=4, STAGES=2>.
|
||||
template <typename Traits, typename QSchedule, typename KV, bool IsCausal, bool HasMask>
|
||||
__global__ void attn_prefill_split_q_mma_kernel(AttentionParams<bf16> p) {
|
||||
const int warp = threadIdx.x / 32;
|
||||
const int lane = threadIdx.x % 32;
|
||||
const int gid = lane >> 2; // 0..7
|
||||
const int tid4 = lane & 3; // 0..3
|
||||
|
||||
const int q_head = blockIdx.y;
|
||||
int batch, q_tile;
|
||||
QSchedule::map_block(p, batch, q_tile);
|
||||
const int kv_head = q_head / (p.q_head / p.kv_head);
|
||||
const int qrow0 = (q_tile * Traits::WARPS + warp) * Traits::BR;
|
||||
|
||||
// Per-request dims (from KV policy — paged reads kv_indptr/qo_indptr).
|
||||
const int seq_len = KV::kv_len(p, batch);
|
||||
const int q_len = QSchedule::q_len(p, batch);
|
||||
const int causal_off = KV::causal_offset(p, batch, q_len);
|
||||
const KVContext kctx = KV::template make_ctx<Traits::HEAD_DIM>(p, batch, kv_head);
|
||||
|
||||
// Static shared memory: double-buffered K/V (no sQ — Q goes direct
|
||||
// to registers in mma A-operand layout).
|
||||
__shared__ __align__(16) bf16 sK[Traits::STAGES * Traits::BC * Traits::LD];
|
||||
__shared__ __align__(16) bf16 sV[Traits::STAGES * Traits::BC * Traits::LD];
|
||||
|
||||
// Load Q fragments straight from global into mma A-operand layout.
|
||||
const int q_base = QSchedule::q_base(p, batch, q_head);
|
||||
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_ptr + q_base, p.q_l_stride, p.q_d_stride,
|
||||
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 int tiles = (seq_len + Traits::BC - 1) / Traits::BC;
|
||||
const int qr0 = qrow0 + gid;
|
||||
const int qr1 = qrow0 + gid + 8;
|
||||
|
||||
// Causal tile-skip bounds (dead code when IsCausal == false)
|
||||
const int max_kv = qrow0 + Traits::BR - 1 + causal_off;
|
||||
const int block_max_kv =
|
||||
q_tile * 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: predicated cp.async (addressing via KV policy) ----
|
||||
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;
|
||||
int token = KV::resolve_token(p, kctx, kc, valid);
|
||||
KVAddr a = KV::kv_addr_from_token(p, kctx, token, d);
|
||||
int off = r * Traits::LD + swiz_col(d, r, Traits::SWIZ_MASK);
|
||||
cp_async_16_pred(&dK[off], a.k, a.valid);
|
||||
cp_async_16_pred(&dV[off], a.v, a.valid);
|
||||
}
|
||||
cp_async_commit();
|
||||
};
|
||||
|
||||
// ---- Prologue: issue first tile load ----
|
||||
load_tile(0, 0);
|
||||
|
||||
for (int ti = 0; ti <= t_end; ti++) {
|
||||
int buf = ti & 1;
|
||||
|
||||
// Wait for current tile, then publish cross-warp + guard buffer reuse.
|
||||
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;
|
||||
|
||||
// Warp-level causal skip (dead branch eliminated when IsCausal == false)
|
||||
if (!IsCausal || kv0 <= max_kv) {
|
||||
|
||||
float Sacc[Traits::NC8][4];
|
||||
mma_compute_scores<Traits>(Qa, bK, lane, Sacc);
|
||||
|
||||
// Post-multiply scale in float (no bf16 precision loss)
|
||||
#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;
|
||||
mma_softmax_tile<Traits, HasMask>(kv0, maxc0, maxc1,
|
||||
qr0, qr1,
|
||||
p.mask_b_stride, p.mask_h_stride, p.mask_l_stride,
|
||||
batch, q_head, q_head,
|
||||
p.mask,
|
||||
va, vb,
|
||||
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 = QSchedule::q_base(p, batch, q_head);
|
||||
#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_ptr[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_ptr[o_base + qr1 * p.q_l_stride + d * p.q_d_stride]) = v;
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,13 @@
|
||||
#pragma once
|
||||
#include <cuda_bf16.h>
|
||||
|
||||
using bf16 = __nv_bfloat16;
|
||||
|
||||
static constexpr int MAX_SPLITS = 32;
|
||||
|
||||
__device__ inline float warp_reduce_sum(float val) {
|
||||
#pragma unroll
|
||||
for (int offset = 16; offset > 0; offset >>= 1)
|
||||
val += __shfl_xor_sync(0xFFFFFFFF, val, offset);
|
||||
return val;
|
||||
}
|
||||
Reference in New Issue
Block a user