KernelBench hard · RTX PRO 6000
Paged Attention DeepSeek V4.1 Flash
36.2%geomean peak fraction across shapes
Hand-written CUDA split-KV paged decode through load_inline: one launch, four warps merging (m, l, acc), an atomic ticket electing the last CTA to combine the chunk partials. What makes it interesting is the ceiling: the agent derived a 1205 GB/s bandwidth roof from torch.Tensor.sum() and stopped, 26% short of what this GPU delivers.
harnessdeepseek-claudeagent session2h 22mtotal wall2h 25mcheck35sbenchmark4soutput tokens224,852cost$16.10gpu-lock wait1h 21mgpu-lock held20mregimememory
Per-shape vs governing ceilingeach shape graded against whichever binds — bf16 compute or HBM bandwidth
8×32×8×128×1024×160.059 ms31.7%0.57 TB/s · 32% of 1.8 TB/s HBM · also 2 TFLOPS (0% of compute)
32×32×8×128×2048×160.261 ms57.3%1.03 TB/s · 57% of 1.8 TB/s HBM · also 4 TFLOPS (1% of compute)
4×64×8×128×4096×160.103 ms36.2%0.65 TB/s · 36% of 1.8 TB/s HBM · also 5 TFLOPS (1% of compute)
16×32×8×128×1535×160.116 ms48.5%0.87 TB/s · 48% of 1.8 TB/s HBM · also 3 TFLOPS (1% of compute)
8×16×4×64×2000×160.047 ms19.6%0.35 TB/s · 20% of 1.8 TB/s HBM · also 1 TFLOPS (0% of compute)
compute-bound memory-bound · bar + right column = official fraction of the ceiling (the geomean input)
geomean(31.7% · 57.3% · 36.2% · 48.5% · 19.6%) = 36.2%
Kernel source (redacted)
"""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 <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <cuda_runtime.h>
#include <cuda_bf16.h>
#include <math_constants.h>
#include <unordered_map>
#include <string>
#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<const uint4*>(p);
__nv_bfloat162 h0 = *reinterpret_cast<const __nv_bfloat162*>(&u.x);
__nv_bfloat162 h1 = *reinterpret_cast<const __nv_bfloat162*>(&u.y);
__nv_bfloat162 h2 = *reinterpret_cast<const __nv_bfloat162*>(&u.z);
__nv_bfloat162 h3 = *reinterpret_cast<const __nv_bfloat162*>(&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<unsigned*>(&h0);
u.y = *reinterpret_cast<unsigned*>(&h1);
u.z = *reinterpret_cast<unsigned*>(&h2);
u.w = *reinterpret_cast<unsigned*>(&h3);
*reinterpret_cast<uint4*>(p) = u;
}
// VEC floats = VEC/8 uint4 loads
template <int VEC>
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 <int VEC>
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 <int VEC>
DEVI void stf4(float* __restrict__ p, const float* o) {
#pragma unroll
for (int q = 0; q < VEC / 4; ++q)
*reinterpret_cast<float4*>(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 <int VEC>
DEVI void ldf4(const float* __restrict__ p, float* o) {
#pragma unroll
for (int q = 0; q < VEC / 4; ++q) {
float4 a = __ldcg(reinterpret_cast<const float4*>(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 <int VEC, int G, int D>
__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<VEC>(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<VEC>(row + off, kv);
ldv<VEC>(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<VEC>(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<const float4*>(src + d0);
float4 t1 = *reinterpret_cast<const float4*>(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<float4*>(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<std::string, Cfg> 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 <int VEC, int G, int D>
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<VEC, G, D><<<grid, 128, 0, stream>>>(
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<const bf16*>(q.data_ptr<at::BFloat16>());
const bf16* kp = reinterpret_cast<const bf16*>(kvc.data_ptr<at::BFloat16>());
const int* btp = bt.data_ptr<int>();
const int* slp = sl.data_ptr<int>();
float* pp = cfg.partial.data_ptr<float>();
int* cp = cfg.counter.data_ptr<int>();
bf16* op = reinterpret_cast<bf16*>(out.data_ptr<at::BFloat16>());
#define DISPATCH(V, GG, DD) \
if (D == DD && G == GG) { \
launch<V, GG, DD>(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
20260910_202124_deepseek-claude_deepseek-flash_03_paged_attention