KernelBench cuda · B200
DeepSeek NSA Claude Fable 5
2.07%geomean peak fraction across shapes
manually audited: clean
harnessor-fableagent session1h 60mtotal wall2h 0mcheck5sbenchmark3soutput tokens—regimecompute
Per-shape vs governing ceilingeach shape graded against whichever binds — bf16 compute or HBM bandwidth
1×16×2048×640.614 ms5.6%28 TFLOPS · 1% of 2,250 TF bf16 peak · also 0.03 TB/s (0% of HBM)
1×16×4127×641.278 ms10.9%55 TFLOPS · 2% of 2,250 TF bf16 peak · also 0.03 TB/s (0% of HBM)
1×8×8192×641.373 ms20.0%100 TFLOPS · 4% of 2,250 TF bf16 peak · also 0.02 TB/s (0% of HBM)
1×8×8191×1282.657 ms20.7%103 TFLOPS · 5% of 2,250 TF bf16 peak · also 0.03 TB/s (0% of HBM)
4×8×1024×640.532 ms3.2%16 TFLOPS · 1% of 2,250 TF bf16 peak · also 0.03 TB/s (0% of HBM)
2×8×3000×640.922 ms8.0%40 TFLOPS · 2% of 2,250 TF bf16 peak · also 0.03 TB/s (0% of HBM)
compute-bound memory-bound · bar + right column = official fraction of the ceiling (the geomean input)
geomean(5.6% · 10.9% · 20.0% · 20.7% · 3.2% · 8.0%) = 9.3% · published 2.1% (lower of repeated isolated re-benchmark passes)
Kernel source (redacted)
"""NSA-style block-select sparse attention — hand-written CUDA (SM100/SM120).
Semantics (matches reference.nsa_attend):
per query t: block importance = mean over causal keys j<=t in each 64-key
block of (q·k)/sqrt(D); select top-8 blocks (ties -> larger block index,
matching Python's descending tuple sort), union with the last-64-token
sliding window, fp32 softmax over exactly that key set.
Structural facts exploited:
* Every non-diagonal causal block is FULL (64 keys), so its importance is
exactly q · mean(k_block) * scale -> precompute block key-means once
(stored transposed [D, NB] so importance loads coalesce across lanes).
* The diagonal block's causal keys [64*(t//64), t] are always a subset of
the sliding window [t-63, t], so selected blocks only contribute keys
strictly below the window start; the window supplies the rest.
Kernel shape (tuned on B200 via ncu + A/B sweeps; see run transcript):
one warp per query, 8 warps/CTA. Scores: 8 (D=64) or 16 (D=128) lanes
cooperate per key row -> 4 cache lines per load instruction instead of 32,
with q hoisted to registers (the shared-mem q reads otherwise re-issue per
iteration and bank-conflict). Scores + key indices live packed as float2
in shared memory, read back as float4 pairs in the P·V pass (halves the
shared-load count). P·V keeps 4 independent coalesced V-row loads in
flight per lane. __launch_bounds__ min-blocks 6 (D=64) / 4 (D=128) buys
occupancy that the natural register allocation leaves on the table.
The Python wrapper caches a CUDA graph per distinct (data_ptr, shape,
stride) input set: the eval harness times individual model calls, and graph
replay removes the pybind + double-launch dispatch overhead. Replay
recomputes from live input memory (verified by in-place mutation tests);
new/rescaled input tensors get their own capture.
"""
from __future__ import annotations
import os
# This box's PATH nvcc is a broken wrapper; point at a real toolkit.
_c = os.environ.get("CUDA_HOME", "")
if not _c or not os.path.exists(os.path.join(_c, "bin", "nvcc")):
for _cand in ("/usr/local/cuda-12.8", "/usr/local/cuda-13", "/usr/local/cuda"):
if os.path.exists(os.path.join(_cand, "bin", "nvcc")):
os.environ["CUDA_HOME"] = _cand
break
import torch
import torch.nn as nn
if "TORCH_CUDA_ARCH_LIST" not in os.environ and torch.cuda.is_available():
_maj, _min = torch.cuda.get_device_capability(0)
os.environ["TORCH_CUDA_ARCH_LIST"] = f"{_maj}.{_min}"
_CPP_SRC = "torch::Tensor nsa_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v);"
_CUDA_SRC = r"""
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAException.h>
#include <cuda_bf16.h>
#include <math.h>
namespace {
constexpr int BLK = 64; // key block size
constexpr int TOPN = 8; // top-n blocks
constexpr int WIN = 64; // sliding window
constexpr int MAXK = TOPN * BLK + WIN; // 576 max selected keys
constexpr int WARPS = 8; // queries per CTA
using bf16 = __nv_bfloat16;
using bf162 = __nv_bfloat162;
__device__ __forceinline__ float bf2f(bf16 x) { return __bfloat162float(x); }
// monotonic float -> uint mapping so uint compare == float compare
__device__ __forceinline__ unsigned f2ord(float f) {
unsigned u = __float_as_uint(f);
return (u & 0x80000000u) ? ~u : (u | 0x80000000u);
}
// per-(b,h,block) key mean, fp32, stored TRANSPOSED as [BH, D, NB] so the
// importance pass reads coalesced across block indices.
__global__ void kmean_kernel(const bf16* __restrict__ k, float* __restrict__ kmean,
int S, int D, int NB) {
const int idx = blockIdx.x; // bh * NB + bi
const int bi = idx % NB;
const int d = threadIdx.x;
if (d >= D) return;
const int s0 = bi * BLK;
const int cnt = min(BLK, S - s0);
const long long bh = idx / NB;
const bf16* kp = k + (bh * S + s0) * (long long)D + d;
float acc = 0.f;
for (int j = 0; j < cnt; ++j) acc += bf2f(kp[(long long)j * D]);
kmean[(bh * D + d) * (long long)NB + bi] = acc / (float)cnt;
}
// main fused kernel: one warp per query
template <int D>
__global__ void __launch_bounds__(WARPS * 32, D == 64 ? 6 : 4)
nsa_kernel(const bf16* __restrict__ q, const bf16* __restrict__ k,
const bf16* __restrict__ v, const float* __restrict__ kmean,
bf16* __restrict__ o, const int S, const int NB, const long long total) {
constexpr int DV = D / 32; // output dims per lane
const int lane = threadIdx.x & 31;
const int wid = threadIdx.x >> 5;
const long long qid = (long long)blockIdx.x * WARPS + wid;
if (qid >= total) return;
const int t = (int)(qid % S);
const long long bh = qid / S;
const float scale = rsqrtf((float)D);
// per-warp shared memory carve-up (host computes the same sizes)
extern __shared__ char smem_raw[];
const int imp_bytes = ((NB * 4 + 15) / 16) * 16;
const int per_warp = D * 4 + imp_bytes + MAXK * 8 + 96;
char* wb = smem_raw + (long long)wid * per_warp;
float* q_s = reinterpret_cast<float*>(wb);
float* imp = reinterpret_cast<float*>(wb + D * 4);
float2* pk = reinterpret_cast<float2*>(wb + D * 4 + imp_bytes);
int* seg = reinterpret_cast<int*>(wb + D * 4 + imp_bytes + MAXK * 8);
int* seg_s = seg; // 9 segment starts
int* seg_c = seg + (TOPN + 1); // 10 cumulative counts; seg[19] = count
const long long base = bh * S;
const bf16* qrow = q + (base + t) * (long long)D;
for (int d = lane; d < D; d += 32) q_s[d] = bf2f(qrow[d]);
__syncwarp();
const int bi_t = t >> 6; // diagonal block
const int nbc = bi_t + 1; // causal block count
const int w0 = max(0, t - (WIN - 1));
// full-block importances from transposed key means: lanes sweep 32
// consecutive block indices per chunk, coalesced along NB
for (int cb = 0; cb < bi_t; cb += 32) {
const int bi = cb + lane;
const bool act = bi < bi_t;
float acc = 0.f;
const float* kmt = kmean + bh * D * (long long)NB;
if (act) {
#pragma unroll 8
for (int d = 0; d < D; ++d)
acc += q_s[d] * kmt[(long long)d * NB + bi];
}
if (act) imp[bi] = acc * scale;
}
// diagonal block: mean over the causal part [bi_t*64, t] only
{
const int s0 = bi_t << 6;
const int cnt = t + 1 - s0;
float ks[DV];
#pragma unroll
for (int i = 0; i < DV; ++i) ks[i] = 0.f;
const bf16* kp = k + (base + s0) * (long long)D + lane * DV;
#pragma unroll 4
for (int j = 0; j < cnt; ++j) {
if constexpr (DV == 2) {
bf162 kk = *reinterpret_cast<const bf162*>(kp + (long long)j * D);
float2 f = __bfloat1622float2(kk);
ks[0] += f.x; ks[1] += f.y;
} else {
uint2 raw = *reinterpret_cast<const uint2*>(kp + (long long)j * D);
float2 f0 = __bfloat1622float2(*reinterpret_cast<bf162*>(&raw.x));
float2 f1 = __bfloat1622float2(*reinterpret_cast<bf162*>(&raw.y));
ks[0] += f0.x; ks[1] += f0.y; ks[2] += f1.x; ks[3] += f1.y;
}
}
float part = 0.f;
#pragma unroll
for (int i = 0; i < DV; ++i) part += ks[i] * q_s[lane * DV + i];
#pragma unroll
for (int off = 16; off; off >>= 1)
part += __shfl_xor_sync(0xffffffffu, part, off);
if (lane == 0) imp[bi_t] = part * scale / (float)cnt;
}
__syncwarp();
// top-8 blocks; ties resolved toward larger block index (reference sorts
// (imp, bi) tuples descending, so equal importances favor higher bi)
int n_sel;
int sel[TOPN];
if (nbc <= TOPN) {
n_sel = nbc;
#pragma unroll
for (int r = 0; r < TOPN; ++r)
if (r < nbc) sel[r] = r;
} else {
n_sel = TOPN;
for (int r = 0; r < TOPN; ++r) {
unsigned long long best = 0ull;
for (int bi = lane; bi < nbc; bi += 32) {
unsigned long long key =
((unsigned long long)f2ord(imp[bi]) << 32) | (unsigned)bi;
if (key > best) best = key;
}
#pragma unroll
for (int off = 16; off; off >>= 1) {
unsigned long long oth = __shfl_xor_sync(0xffffffffu, best, off);
if (oth > best) best = oth;
}
const int chosen = (int)(best & 0xffffffffu);
sel[r] = chosen;
__syncwarp();
if (lane == 0) imp[chosen] = -INFINITY;
__syncwarp();
}
}
// selected key ranges: blocks clipped below the window start, then window
if (lane == 0) {
int nseg = 0, cum = 0;
for (int r = 0; r < n_sel; ++r) {
const int bi = sel[r];
if (bi == bi_t) continue; // fully inside the window
const int s0 = bi << 6;
const int e = min(s0 + BLK, w0);
if (e > s0) { seg_s[nseg] = s0; seg_c[nseg] = cum; cum += e - s0; ++nseg; }
}
seg_s[nseg] = w0; seg_c[nseg] = cum; cum += t + 1 - w0; ++nseg;
seg_c[nseg] = cum;
seg[19] = nseg;
}
__syncwarp();
const int nseg = seg[19];
const int M = seg_c[nseg];
// scores: GL lanes cooperate per key row (4 cache lines per load instr
// instead of 32); q chunk hoisted to registers; packed (score, index)
float lmax = -INFINITY;
{
constexpr int GL = (D == 64) ? 8 : 16; // lanes per key row
constexpr int GK = 32 / GL; // keys per iteration
const int g = lane / GL;
const int sub = lane % GL;
float qreg[8];
#pragma unroll
for (int i = 0; i < 8; ++i) qreg[i] = q_s[sub * 8 + i];
int r = 0;
for (int sb = 0; sb < M; sb += GK) {
const int s = sb + g;
const bool act = s < M;
int j = 0;
if (act) {
while (r + 1 < nseg && s >= seg_c[r + 1]) ++r;
j = seg_s[r] + (s - seg_c[r]);
}
float acc = 0.f;
if (act) {
const uint4 raw = *reinterpret_cast<const uint4*>(
k + (base + j) * (long long)D + sub * 8);
float2 f0 = __bfloat1622float2(*reinterpret_cast<const bf162*>(&raw.x));
float2 f1 = __bfloat1622float2(*reinterpret_cast<const bf162*>(&raw.y));
float2 f2 = __bfloat1622float2(*reinterpret_cast<const bf162*>(&raw.z));
float2 f3 = __bfloat1622float2(*reinterpret_cast<const bf162*>(&raw.w));
acc = f0.x * qreg[0] + f0.y * qreg[1]
+ f1.x * qreg[2] + f1.y * qreg[3]
+ f2.x * qreg[4] + f2.y * qreg[5]
+ f3.x * qreg[6] + f3.y * qreg[7];
}
acc += __shfl_xor_sync(0xffffffffu, acc, 1);
acc += __shfl_xor_sync(0xffffffffu, acc, 2);
acc += __shfl_xor_sync(0xffffffffu, acc, 4);
if (GL == 16) acc += __shfl_xor_sync(0xffffffffu, acc, 8);
if (act) {
const float sc = acc * scale;
if (sub == 0) pk[s] = make_float2(sc, __int_as_float(j));
lmax = fmaxf(lmax, sc);
}
}
}
#pragma unroll
for (int off = 16; off; off >>= 1)
lmax = fmaxf(lmax, __shfl_xor_sync(0xffffffffu, lmax, off));
float lsum = 0.f;
{
float* pkf = reinterpret_cast<float*>(pk);
for (int s = lane; s < M; s += 32) {
const float e = __expf(pkf[2 * s] - lmax);
pkf[2 * s] = e;
lsum += e;
}
}
#pragma unroll
for (int off = 16; off; off >>= 1)
lsum += __shfl_xor_sync(0xffffffffu, lsum, off);
__syncwarp();
// P·V: lane owns DV output dims; packed pairs are read two slots at a
// time; the 4-wide unroll keeps 4 independent coalesced V loads in flight
float acc[DV], acc1[DV], acc2[DV], acc3[DV];
#pragma unroll
for (int i = 0; i < DV; ++i) acc[i] = acc1[i] = acc2[i] = acc3[i] = 0.f;
const bf16* vb = v + base * (long long)D + lane * DV;
const float4* pk4 = reinterpret_cast<const float4*>(pk);
int s = 0;
for (; s + 4 <= M; s += 4) {
const float4 ab = pk4[s >> 1];
const float4 cd = pk4[(s >> 1) + 1];
const float p0 = ab.x, p1 = ab.z, p2 = cd.x, p3 = cd.z;
const long long j0 = __float_as_int(ab.y);
const long long j1 = __float_as_int(ab.w);
const long long j2 = __float_as_int(cd.y);
const long long j3 = __float_as_int(cd.w);
if constexpr (DV == 2) {
float2 f0 = __bfloat1622float2(*reinterpret_cast<const bf162*>(vb + j0 * D));
float2 f1 = __bfloat1622float2(*reinterpret_cast<const bf162*>(vb + j1 * D));
float2 f2 = __bfloat1622float2(*reinterpret_cast<const bf162*>(vb + j2 * D));
float2 f3 = __bfloat1622float2(*reinterpret_cast<const bf162*>(vb + j3 * D));
acc[0] += p0 * f0.x; acc[1] += p0 * f0.y;
acc1[0] += p1 * f1.x; acc1[1] += p1 * f1.y;
acc2[0] += p2 * f2.x; acc2[1] += p2 * f2.y;
acc3[0] += p3 * f3.x; acc3[1] += p3 * f3.y;
} else {
uint2 r0 = *reinterpret_cast<const uint2*>(vb + j0 * D);
uint2 r1 = *reinterpret_cast<const uint2*>(vb + j1 * D);
uint2 r2 = *reinterpret_cast<const uint2*>(vb + j2 * D);
uint2 r3 = *reinterpret_cast<const uint2*>(vb + j3 * D);
float2 f0a = __bfloat1622float2(*reinterpret_cast<bf162*>(&r0.x));
float2 f0b = __bfloat1622float2(*reinterpret_cast<bf162*>(&r0.y));
float2 f1a = __bfloat1622float2(*reinterpret_cast<bf162*>(&r1.x));
float2 f1b = __bfloat1622float2(*reinterpret_cast<bf162*>(&r1.y));
float2 f2a = __bfloat1622float2(*reinterpret_cast<bf162*>(&r2.x));
float2 f2b = __bfloat1622float2(*reinterpret_cast<bf162*>(&r2.y));
float2 f3a = __bfloat1622float2(*reinterpret_cast<bf162*>(&r3.x));
float2 f3b = __bfloat1622float2(*reinterpret_cast<bf162*>(&r3.y));
acc[0] += p0 * f0a.x; acc[1] += p0 * f0a.y;
acc[2] += p0 * f0b.x; acc[3] += p0 * f0b.y;
acc1[0] += p1 * f1a.x; acc1[1] += p1 * f1a.y;
acc1[2] += p1 * f1b.x; acc1[3] += p1 * f1b.y;
acc2[0] += p2 * f2a.x; acc2[1] += p2 * f2a.y;
acc2[2] += p2 * f2b.x; acc2[3] += p2 * f2b.y;
acc3[0] += p3 * f3a.x; acc3[1] += p3 * f3a.y;
acc3[2] += p3 * f3b.x; acc3[3] += p3 * f3b.y;
}
}
for (; s < M; ++s) {
const float2 pj = pk[s];
const long long j = __float_as_int(pj.y);
if constexpr (DV == 2) {
float2 f = __bfloat1622float2(*reinterpret_cast<const bf162*>(vb + j * D));
acc[0] += pj.x * f.x; acc[1] += pj.x * f.y;
} else {
uint2 raw = *reinterpret_cast<const uint2*>(vb + j * D);
float2 f0 = __bfloat1622float2(*reinterpret_cast<bf162*>(&raw.x));
float2 f1 = __bfloat1622float2(*reinterpret_cast<bf162*>(&raw.y));
acc[0] += pj.x * f0.x; acc[1] += pj.x * f0.y;
acc[2] += pj.x * f1.x; acc[3] += pj.x * f1.y;
}
}
#pragma unroll
for (int i = 0; i < DV; ++i) acc[i] += acc1[i] + (acc2[i] + acc3[i]);
const float inv = 1.f / lsum;
bf16* orow = o + (base + t) * (long long)D + lane * DV;
#pragma unroll
for (int i = 0; i < DV; ++i) orow[i] = __float2bfloat16(acc[i] * inv);
}
} // namespace
torch::Tensor nsa_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v) {
TORCH_CHECK(q.is_cuda() && k.is_cuda() && v.is_cuda(), "cuda tensors required");
TORCH_CHECK(q.scalar_type() == torch::kBFloat16 &&
k.scalar_type() == torch::kBFloat16 &&
v.scalar_type() == torch::kBFloat16, "bf16 tensors required");
TORCH_CHECK(q.dim() == 4, "expected (B,H,S,D)");
auto qc = q.contiguous();
auto kc = k.contiguous();
auto vc = v.contiguous();
const int B = qc.size(0), H = qc.size(1), S = qc.size(2), Dh = qc.size(3);
TORCH_CHECK(Dh == 64 || Dh == 128, "D must be 64 or 128");
TORCH_CHECK(S <= 65535, "S exceeds supported range");
const int NB = (S + BLK - 1) / BLK;
auto stream = at::cuda::getCurrentCUDAStream();
auto kmean = torch::empty(
{(long long)B * H, Dh, NB},
torch::TensorOptions().dtype(torch::kFloat32).device(q.device()));
const bf16* qp = reinterpret_cast<const bf16*>(qc.data_ptr());
const bf16* kp = reinterpret_cast<const bf16*>(kc.data_ptr());
const bf16* vp = reinterpret_cast<const bf16*>(vc.data_ptr());
kmean_kernel<<<B * H * NB, Dh, 0, stream>>>(kp, kmean.data_ptr<float>(), S, Dh, NB);
C10_CUDA_KERNEL_LAUNCH_CHECK();
auto o = torch::empty_like(qc);
bf16* op = reinterpret_cast<bf16*>(o.data_ptr());
const long long total = (long long)B * H * S;
const int imp_bytes = ((NB * 4 + 15) / 16) * 16;
const int per_warp = Dh * 4 + imp_bytes + MAXK * 8 + 96;
const int smem = per_warp * WARPS;
const int grid = (int)((total + WARPS - 1) / WARPS);
if (Dh == 64) {
static bool init64 = false;
if (!init64) {
cudaFuncSetAttribute(nsa_kernel<64>,
cudaFuncAttributeMaxDynamicSharedMemorySize, 200 * 1024);
init64 = true;
}
nsa_kernel<64><<<grid, WARPS * 32, smem, stream>>>(
qp, kp, vp, kmean.data_ptr<float>(), op, S, NB, total);
} else {
static bool init128 = false;
if (!init128) {
cudaFuncSetAttribute(nsa_kernel<128>,
cudaFuncAttributeMaxDynamicSharedMemorySize, 200 * 1024);
init128 = true;
}
nsa_kernel<128><<<grid, WARPS * 32, smem, stream>>>(
qp, kp, vp, kmean.data_ptr<float>(), op, S, NB, total);
}
C10_CUDA_KERNEL_LAUNCH_CHECK();
return o;
}
"""
_ext = None
def _get_ext():
global _ext
if _ext is None:
from torch.utils.cpp_extension import load_inline
_ext = load_inline(
name="nsa_sparse_attn_v6",
cpp_sources=[_CPP_SRC],
cuda_sources=[_CUDA_SRC],
functions=["nsa_forward"],
extra_cuda_cflags=["-O3", "--use_fast_math"],
verbose=False,
)
return _ext
class _GraphCache:
"""CUDA-graph replay cache keyed on the exact input tensors.
Replay reads the live contents of the captured input pointers, so
mutating an input in place and replaying recomputes honestly; any new
tensor (different pointer / shape / stride) triggers a fresh capture.
"""
MAX_ENTRIES = 24
def __init__(self):
self._cache = {}
@staticmethod
def _key(q, k, v):
return (
q.data_ptr(), k.data_ptr(), v.data_ptr(),
tuple(q.shape), tuple(q.stride()),
tuple(k.stride()), tuple(v.stride()),
)
def __call__(self, ext, q, k, v):
if os.environ.get("KBH_NSA_NO_GRAPH") == "1":
return ext.nsa_forward(q, k, v)
key = self._key(q, k, v)
ent = self._cache.get(key)
if ent is None:
if len(self._cache) >= self.MAX_ENTRIES:
self._cache.clear()
try:
side = torch.cuda.Stream()
side.wait_stream(torch.cuda.current_stream())
with torch.cuda.stream(side):
ext.nsa_forward(q, k, v) # warm up allocator/module state
torch.cuda.current_stream().wait_stream(side)
graph = torch.cuda.CUDAGraph()
with torch.cuda.graph(graph):
out = ext.nsa_forward(q, k, v)
except Exception:
return ext.nsa_forward(q, k, v)
ent = (graph, out)
self._cache[key] = ent
graph, out = ent
graph.replay()
return out
_graphs = _GraphCache()
class Model(nn.Module):
def __init__(self, B: int, H: int, S: int, D: int):
super().__init__()
self.B, self.H, self.S, self.D = B, H, S, D
self.register_buffer("_dummy", torch.zeros(1, dtype=torch.bfloat16))
_get_ext()
def forward(self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor) -> torch.Tensor:
return _graphs(_get_ext(), q, k, v)
def get_init_inputs():
return [1, 16, 1024, 64]
def get_inputs():
q = torch.randn(1, 16, 1024, 64, dtype=torch.bfloat16)
k = torch.randn(1, 16, 1024, 64, dtype=torch.bfloat16)
v = torch.randn(1, 16, 1024, 64, dtype=torch.bfloat16)
return [q, k, v]
20260719_092604_or-fable_anthropic_claude-fable-5_02_deepseek_nsa