"""Paged-attention decode kernel (SM120 Blackwell) — custom CUDA. Single-query decode with a paged KV cache packed as [K | V] along the last dim. Design notes ------------ The KV cache pool is (num_blocks, page_size, num_kv_heads, 2*head_dim) bf16, so one (token, kv_head) row is 2*head_dim*2 contiguous bytes: K in the first half, V in the second. A warp of 32 lanes with 16-byte vector loads covers exactly 128 elements = one K row (or one V row). Work decomposition (memory-bound: the kernel is bandwidth-limited, not compute-limited, so the design maximises outstanding 16-byte loads per warp): * one CTA = 128 threads = 4 warps owns `chunk_tokens` consecutive KV tokens of one (batch, kv_head) pair; each warp takes a `chunk_tokens/4` slice, so the CTA is the unit that publishes a partial and the split-K fan-out is 4x cheaper than a warp-per-chunk schedule. * a head_dim row is covered by 8 lanes x VEC dims (VEC = head_dim/8), so one warp-load of a KV row feeds 4 query heads at once and the QK^T reduction is a 3-step butterfly (masks 1,2,4) inside each 8-lane group. The butterflies of all rounds are interleaved for ILP. * the 4 warps merge their (m, l, acc) partials through shared memory, then one global partial per CTA is stored. * split-K combining is fused into the same launch: after a __threadfence(), thread 0 of each CTA takes a ticket from a per-(batch, kv_head) counter; the CTA drawing the last ticket combines the (L2-resident) CTA partials and writes the bf16 output. The partials are read with __ldcg so they come from L2, which is the coherence point between SMs. """ import math import torch import torch.nn as nn from torch.utils.cpp_extension import load_inline _CUDA_SRC = r""" #include #include #include #include #include #include #include #define DEVI __device__ __forceinline__ #define LOG2E 1.4426950408889634f using bf16 = __nv_bfloat16; // ---- 16-byte vector load/store helpers (bf16 <-> float) ------------------- DEVI void ld8(const bf16* __restrict__ p, float* o) { uint4 u = *reinterpret_cast(p); __nv_bfloat162 h0 = *reinterpret_cast(&u.x); __nv_bfloat162 h1 = *reinterpret_cast(&u.y); __nv_bfloat162 h2 = *reinterpret_cast(&u.z); __nv_bfloat162 h3 = *reinterpret_cast(&u.w); float2 f0 = __bfloat1622float2(h0); float2 f1 = __bfloat1622float2(h1); float2 f2 = __bfloat1622float2(h2); float2 f3 = __bfloat1622float2(h3); o[0] = f0.x; o[1] = f0.y; o[2] = f1.x; o[3] = f1.y; o[4] = f2.x; o[5] = f2.y; o[6] = f3.x; o[7] = f3.y; } DEVI void st8(bf16* __restrict__ p, const float* o) { __nv_bfloat162 h0 = __floats2bfloat162_rn(o[0], o[1]); __nv_bfloat162 h1 = __floats2bfloat162_rn(o[2], o[3]); __nv_bfloat162 h2 = __floats2bfloat162_rn(o[4], o[5]); __nv_bfloat162 h3 = __floats2bfloat162_rn(o[6], o[7]); uint4 u; u.x = *reinterpret_cast(&h0); u.y = *reinterpret_cast(&h1); u.z = *reinterpret_cast(&h2); u.w = *reinterpret_cast(&h3); *reinterpret_cast(p) = u; } // VEC floats = VEC/8 uint4 loads template DEVI void ldv(const bf16* __restrict__ p, float* o) { #pragma unroll for (int q = 0; q < VEC / 8; ++q) ld8(p + 8 * q, o + 8 * q); } template DEVI void stv(bf16* __restrict__ p, const float* o) { #pragma unroll for (int q = 0; q < VEC / 8; ++q) st8(p + 8 * q, o + 8 * q); } // VEC floats as float4 stores template DEVI void stf4(float* __restrict__ p, const float* o) { #pragma unroll for (int q = 0; q < VEC / 4; ++q) *reinterpret_cast(p + 4 * q) = make_float4(o[4 * q], o[4 * q + 1], o[4 * q + 2], o[4 * q + 3]); } // L2-only float4 loads (partials are produced by other SMs, L1 is not coherent) template DEVI void ldf4(const float* __restrict__ p, float* o) { #pragma unroll for (int q = 0; q < VEC / 4; ++q) { float4 a = __ldcg(reinterpret_cast(p + 4 * q)); o[4 * q] = a.x; o[4 * q + 1] = a.y; o[4 * q + 2] = a.z; o[4 * q + 3] = a.w; } } // --------------------------------------------------------------------------- // Kernel layout // * one CTA = 128 threads = 4 warps, owning `chunk_tokens` consecutive KV // tokens of one (batch, kv_head) pair; each warp takes a `chunk_tokens/4` // token slice. // * a head_dim row is covered by 8 lanes x VEC dims (VEC = head_dim/8), so a // warp serves 4 query heads at once and the QK^T reduction is a 3-step // butterfly (masks 1,2,4) inside each 8-lane group. // * the 4 warps merge their (m, l, acc) partials in shared memory, so only // one global partial per CTA is published. // * split-K: a per-(batch,kv_head) ticket counter; the CTA that draws the // last ticket combines the (L2-resident) CTA partials and writes the // bf16 output. // --------------------------------------------------------------------------- template __global__ void __launch_bounds__(128) pa_decode_kernel( const bf16* __restrict__ q, // (B, H, D) const bf16* __restrict__ kvc, // (NB, P, Hkv, 2D) packed [K|V] const int* __restrict__ block_table, const int* __restrict__ seq_lens, float* __restrict__ partial, // (B*Hkv, num_chunks, G, D+4) int* __restrict__ counter, // (B*Hkv) bf16* __restrict__ out, // (B, H, D) int Hkv, int max_blocks, int log2_page, int chunk_tokens, int num_chunks, float scale) { constexpr int GRP = 4; // head groups per warp (32 lanes / 8 lanes) constexpr int R = G / GRP; // rounds constexpr int PAD = D + 4; // 16-byte-aligned row stride constexpr int NBLK = G * D / 8; // active threads in merge / reduce phases static_assert(G % GRP == 0, "G must be a multiple of 4"); static_assert(D % 8 == 0, "head_dim must be a multiple of 8"); __shared__ float sp[4 * G * PAD]; // per-warp partials __shared__ int sflag; const int tid = threadIdx.x; const int lane = tid & 31; const int warp = tid >> 5; const int grp = lane >> 3; // 8-lane head group const int dj = lane & 7; // dim block inside the head const int off = dj * VEC; const int pair = blockIdx.y; const int b = pair / Hkv; const int kvh = pair - b * Hkv; const int H = G * Hkv; const int D2 = 2 * D; const size_t row_stride = (size_t)Hkv * D2; const int L = seq_lens[b]; const int wt = chunk_tokens >> 2; // tokens per warp const int t0 = blockIdx.x * chunk_tokens + warp * wt; int t1 = t0 + wt; if (t1 > L) t1 = L; float qr[R][VEC]; #pragma unroll for (int r = 0; r < R; ++r) { const int g = r * GRP + grp; ldv(q + ((size_t)b * H + kvh * G + g) * D + off, qr[r]); } float m[R], l[R], acc[R][VEC]; #pragma unroll for (int r = 0; r < R; ++r) { m[r] = -CUDART_INF_F; l[r] = 0.f; #pragma unroll for (int i = 0; i < VEC; ++i) acc[r][i] = 0.f; } if (t0 < L) { const int slot_mask = (1 << log2_page) - 1; const int npage = 1 << log2_page; const int* bt = block_table + (size_t)b * max_blocks; const int p0 = t0 >> log2_page; const int p1 = (t1 + slot_mask) >> log2_page; for (int p = p0; p < p1; ++p) { const int page = __ldg(bt + p); const bf16* base = kvc + ((size_t)page << log2_page) * row_stride + (size_t)kvh * D2; const int first = (p == p0) ? (t0 & slot_mask) : 0; const int last = (p == p1 - 1) ? ((t1 - 1) & slot_mask) + 1 : npage; #pragma unroll 2 for (int s = first; s < last; ++s) { const bf16* row = base + (size_t)s * row_stride; float kv[VEC], vv[VEC]; ldv(row + off, kv); ldv(row + D + off, vv); float sc[R]; #pragma unroll for (int r = 0; r < R; ++r) { float t = 0.f; #pragma unroll for (int i = 0; i < VEC; ++i) t = fmaf(qr[r][i], kv[i], t); sc[r] = t; } // interleave the independent butterflies for ILP #pragma unroll for (int st = 1; st <= 4; st <<= 1) { #pragma unroll for (int r = 0; r < R; ++r) sc[r] += __shfl_xor_sync(0xffffffffu, sc[r], st); } #pragma unroll for (int r = 0; r < R; ++r) { const float s = sc[r] * scale; const float mn = fmaxf(m[r], s); if (mn > m[r]) { const float a = exp2f((m[r] - mn) * LOG2E); l[r] *= a; #pragma unroll for (int i = 0; i < VEC; ++i) acc[r][i] *= a; m[r] = mn; } const float pv = exp2f((s - m[r]) * LOG2E); l[r] += pv; #pragma unroll for (int i = 0; i < VEC; ++i) acc[r][i] = fmaf(pv, vv[i], acc[r][i]); } } } } // ---- publish this warp's partial into shared memory ------------------- { float* wp = sp + warp * (G * PAD); #pragma unroll for (int r = 0; r < R; ++r) { const int g = r * GRP + grp; float* dst = wp + g * PAD; if (dj == 0) { dst[D] = m[r]; dst[D + 1] = l[r]; } stf4(dst + off, acc[r]); } } __syncthreads(); // ---- merge the 4 warp partials; publish one CTA partial ---------------- if (tid < NBLK) { const int g = (tid * 8) / D; const int d0 = (tid * 8) % D; float mm = -CUDART_INF_F; #pragma unroll for (int w = 0; w < 4; ++w) mm = fmaxf(mm, sp[(w * G + g) * PAD + D]); float o[8]; #pragma unroll for (int i = 0; i < 8; ++i) o[i] = 0.f; float ls = 0.f; #pragma unroll for (int w = 0; w < 4; ++w) { const float* src = sp + (w * G + g) * PAD; const float lw = src[D + 1]; if (lw > 0.f) { const float a = exp2f((src[D] - mm) * LOG2E); ls += lw * a; float4 t0 = *reinterpret_cast(src + d0); float4 t1 = *reinterpret_cast(src + d0 + 4); o[0] += a * t0.x; o[1] += a * t0.y; o[2] += a * t0.z; o[3] += a * t0.w; o[4] += a * t1.x; o[5] += a * t1.y; o[6] += a * t1.z; o[7] += a * t1.w; } } float* dst = partial + (((size_t)pair * num_chunks + blockIdx.x) * G + g) * PAD + d0; #pragma unroll for (int i = 0; i < 8; i += 4) *reinterpret_cast(dst + i) = make_float4(o[i], o[i + 1], o[i + 2], o[i + 3]); if (d0 == 0) { dst[D] = mm; dst[D + 1] = ls; } } __threadfence(); __syncthreads(); if (tid == 0) sflag = atomicAdd(&counter[pair], 1); __syncthreads(); if (sflag != num_chunks - 1) return; // ---- last CTA for this pair: combine the chunk partials ---------------- __threadfence(); if (tid < NBLK) { const int g = (tid * 8) / D; const int d0 = (tid * 8) % D; const size_t cstride = (size_t)G * PAD; const float* cbase = partial + ((size_t)pair * num_chunks * G + g) * PAD; float mx = -CUDART_INF_F; for (int c = 0; c < num_chunks; ++c) mx = fmaxf(mx, __ldcg(cbase + c * cstride + D)); float o[8]; #pragma unroll for (int i = 0; i < 8; ++i) o[i] = 0.f; float ls = 0.f; for (int c = 0; c < num_chunks; ++c) { const float* ck = cbase + c * cstride; const float lc = __ldcg(ck + D + 1); if (lc > 0.f) { const float a = exp2f((__ldcg(ck + D) - mx) * LOG2E); ls += lc * a; float tmp[8]; ldf4<8>(ck + d0, tmp); #pragma unroll for (int i = 0; i < 8; ++i) o[i] = fmaf(a, tmp[i], o[i]); } } if (ls > 0.f) { const float inv = 1.f / ls; #pragma unroll for (int i = 0; i < 8; ++i) o[i] *= inv; } stv<8>(out + ((size_t)b * H + kvh * G + g) * D + d0, o); if (tid == 0) counter[pair] = 0; } } // --------------------------------------------------------------------------- // Launcher: cached scratch buffers + a tiny dispatch. // --------------------------------------------------------------------------- struct Cfg { torch::Tensor partial; torch::Tensor counter; int num_chunks = 0; }; static std::unordered_map g_cfg; static int g_chunk_tokens = -1; static int chunk_tokens_env() { if (g_chunk_tokens < 0) { const char* e = getenv("PA_CT"); g_chunk_tokens = e ? atoi(e) : 128; if (g_chunk_tokens < 64) g_chunk_tokens = 64; } return g_chunk_tokens; } template static void launch(const bf16* q, const bf16* kvc, const int* bt, const int* sl, float* partial, int* counter, bf16* out, int B, int Hkv, int MB, int log2_page, int chunk_tokens, int num_chunks, float scale, cudaStream_t stream) { dim3 grid(num_chunks, B * Hkv); pa_decode_kernel<<>>( q, kvc, bt, sl, partial, counter, out, Hkv, MB, log2_page, chunk_tokens, num_chunks, scale); } torch::Tensor paged_decode(torch::Tensor q, torch::Tensor kvc, torch::Tensor bt, torch::Tensor sl) { const int B = q.size(0); const int H = q.size(1); const int D = q.size(2); const int P = kvc.size(1); const int Hkv = kvc.size(2); const int MB = bt.size(1); const int G = H / Hkv; TORCH_CHECK(D == 128 || D == 64, "unsupported head_dim ", D); TORCH_CHECK(G == 4 || G == 8, "unsupported GQA group ", G); TORCH_CHECK((P & (P - 1)) == 0, "page_size must be a power of two"); const int log2_page = __builtin_ctz(P); int chunk_tokens = chunk_tokens_env(); if (chunk_tokens % (4 * P) != 0) chunk_tokens = 4 * P; // warp slices stay page aligned // Size the schedule from the full configured sequence length so the cached // buffers never need to grow when seq_lens shrink. const int max_len = MB << log2_page; const int num_chunks = (max_len + chunk_tokens - 1) / chunk_tokens; std::string key = std::to_string(B) + ":" + std::to_string(H) + ":" + std::to_string(Hkv) + ":" + std::to_string(D) + ":" + std::to_string(num_chunks); auto it = g_cfg.find(key); if (it == g_cfg.end()) { Cfg c; c.partial = torch::zeros({(long)(B * Hkv) * num_chunks * G * (D + 4)}, q.options().dtype(torch::kFloat32)); c.counter = torch::zeros({B * Hkv}, q.options().dtype(torch::kInt32)); c.num_chunks = num_chunks; it = g_cfg.emplace(key, std::move(c)).first; } Cfg& cfg = it->second; auto out = torch::empty({B, H, D}, q.options()); cudaStream_t stream = at::cuda::getCurrentCUDAStream(); const float scale = 1.0f / sqrtf((float)D); const bf16* qp = reinterpret_cast(q.data_ptr()); const bf16* kp = reinterpret_cast(kvc.data_ptr()); const int* btp = bt.data_ptr(); const int* slp = sl.data_ptr(); float* pp = cfg.partial.data_ptr(); int* cp = cfg.counter.data_ptr(); bf16* op = reinterpret_cast(out.data_ptr()); #define DISPATCH(V, GG, DD) \ if (D == DD && G == GG) { \ launch(qp, kp, btp, slp, pp, cp, op, B, Hkv, MB, log2_page, \ chunk_tokens, num_chunks, scale, stream); \ } else DISPATCH(16, 4, 128) DISPATCH(16, 8, 128) DISPATCH(8, 4, 64) DISPATCH(8, 8, 64) { TORCH_CHECK(false, "unreachable dispatch"); } #undef DISPATCH return out; } """ _CPP_SRC = "torch::Tensor paged_decode(torch::Tensor q, torch::Tensor kvc, torch::Tensor bt, torch::Tensor sl);" _ext = None def _get_ext(): global _ext if _ext is None: import os import shutil import sys # The interpreter's own bin dir usually holds the `ninja` wheel entry # point; torch shells out to a bare `ninja`, so make sure it is on PATH. bindir = os.path.dirname(os.path.abspath(sys.executable)) if shutil.which("ninja") is None and os.path.exists(os.path.join(bindir, "ninja")): os.environ["PATH"] = bindir + os.pathsep + os.environ.get("PATH", "") os.environ.setdefault("TORCH_CUDA_ARCH_LIST", "12.0") _ext = load_inline( name="pa_decode_ext", cpp_sources=_CPP_SRC, cuda_sources=_CUDA_SRC, functions=["paged_decode"], extra_cuda_cflags=["-O3"], verbose=False, ) return _ext class Model(nn.Module): """Single-query paged attention decode.""" def __init__( self, batch: int, num_heads: int, num_kv_heads: int, head_dim: int, seq_len: int, page_size: int, ): super().__init__() assert num_heads % num_kv_heads == 0, "num_heads must be a multiple of num_kv_heads (GQA)" self.batch = batch self.num_heads = num_heads self.num_kv_heads = num_kv_heads self.head_dim = head_dim self.seq_len = seq_len self.page_size = page_size self.group_size = num_heads // num_kv_heads self.scale = 1.0 / math.sqrt(head_dim) self._ext = _get_ext() def forward(self, query, kv_cache, block_table, seq_lens): return self._ext.paged_decode(query, kv_cache, block_table, seq_lens) def get_inputs(): """Build random paged inputs for the current module-level shape knobs.""" B = BATCH H = NUM_HEADS Hkv = NUM_KV_HEADS D = HEAD_DIM L = SEQ_LEN P = PAGE_SIZE pages_per_seq = (L + P - 1) // P total_pages = max(B * pages_per_seq + 8, 64) query = torch.randn(B, H, D, dtype=torch.bfloat16) * 0.1 kv_cache = torch.randn(total_pages, P, Hkv, 2 * D, dtype=torch.bfloat16) * 0.1 perm = torch.randperm(total_pages)[: B * pages_per_seq].reshape(B, pages_per_seq).int() block_table = perm.contiguous() seq_lens = torch.full((B,), L, dtype=torch.int32) return [query, kv_cache, block_table, seq_lens] def get_init_inputs(): return [BATCH, NUM_HEADS, NUM_KV_HEADS, HEAD_DIM, SEQ_LEN, PAGE_SIZE] # --- Shape knobs (overridden by check.py / benchmark.py from shapes.py) ---- BATCH = 8 NUM_HEADS = 32 NUM_KV_HEADS = 8 HEAD_DIM = 128 SEQ_LEN = 1024 PAGE_SIZE = 16