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