KernelBench cuda · RTX PRO 6000
DeepSeek NSA DeepSeek V4.1 Flash
0.196 msgeomean latency across six shapes · lower is better
manually audited: clean
DeepSeek V4.1 Flash hand-wrote an SM120 NSA kernel in inline PTX: mma.sync bf16, ldmatrix, xor-swizzled cp.async, and an fp32 block-scoring top-8 prologue fused into the attention kernel. It reproduces the reference's exact tie-break and then over-computes, executing about 78% of the dense causal block triangle where the semantics need 14%.
harnessdeepseek-claudeagent session8h 55mtotal wall8h 57mcheck6mbenchmark1soutput tokens1,260,278cost$62.90gpu-lock wait1h 19mgpu-lock held4h 40mregimecompute
Per-shape latencymilliseconds are the measured result; the dense-equivalent roofline is not a useful ceiling
1×16×2048×640.090 ms
1×16×4127×640.239 ms
1×8×8192×640.408 ms
1×8×8191×1280.736 ms
4×8×1024×640.059 ms
2×8×3000×640.150 ms
geomean(0.090 ms · 0.239 ms · 0.408 ms · 0.736 ms · 0.059 ms · 0.150 ms) = 0.196 ms
Kernel source (redacted)
"""DeepSeek NSA-inspired sparse attention -- CUDA implementation (SM120).
Semantics match reference.nsa_attend exactly (bench-faithful simplification of
Native Sparse Attention: block scoring -> top-n block select -> union with a
local sliding window -> softmax over the selected keys only).
Kernel plan
A1 blocksum : BM[b,h,bi,:] = sum_{j in block bi} K[b,h,j,:] (fp64 accum)
A2 select : block importance = q . BM / len (fp32; for the query's own
diagonal block the sum is causally truncated via a Kahan prefix
scan) followed by a running top-8 per query.
The identity mean_j(q.k_j) = q . (sum_j k_j) / L removes the
O(S^2) score pass entirely: scoring is O(B H S D) FLOPs.
B attend : tiled causal online-softmax attention (m16n8k16 mma, 4 warps,
16 rows/warp, double-buffered cp.async) over every causal block;
NSA selection and the sliding window enter purely as a
per-element mask built from a per-CTA bitmap of the selected
blocks. Dense-tile masking is the right trade: the per-query
selection sets are near-uniform for the benchmark distributions,
so a gathered variant would move far more bytes than it saves.
"""
from __future__ import annotations
import os
import sys
# load_inline shells out to `ninja`; make sure the interpreter's bin dir (which
# may hold it) is reachable even when PATH was not set up by an activate script.
_bindir = os.path.dirname(os.path.abspath(sys.executable))
_parts = os.environ.get("PATH", "").split(os.pathsep)
for _p in (_bindir, "/[REDACTED-home]/kb-cuda/.venv/bin"):
if _p and os.path.isdir(_p) and _p not in _parts:
_parts.append(_p)
os.environ["PATH"] = os.pathsep.join(_parts)
import torch
import torch.nn as nn
from torch.utils.cpp_extension import load_inline
BLOCK_SIZE = 64
TOP_N_BLOCKS = 8
SLIDING_WINDOW = 64
_CUDA_SRC = r"""
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAException.h>
#include <cuda_runtime.h>
#include <cuda_bf16.h>
#include <math_constants.h>
#include <cstdint>
#include <cstring>
#include <cmath>
using bf16 = __nv_bfloat16;
#define BS 64
#define TOPN 8
#define DEV __device__ __forceinline__
DEV float exp2_approx(float x) {
float r;
asm("ex2.approx.f32 %0, %1;\n" : "=f"(r) : "f"(x));
return r;
}
DEV void mma16816(float c[4], const uint32_t a[4], uint32_t b0, uint32_t b1) {
asm volatile(
"mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32 "
"{%0,%1,%2,%3}, {%4,%5,%6,%7}, {%8,%9}, {%0,%1,%2,%3};\n"
: "+f"(c[0]), "+f"(c[1]), "+f"(c[2]), "+f"(c[3])
: "r"(a[0]), "r"(a[1]), "r"(a[2]), "r"(a[3]), "r"(b0), "r"(b1));
}
DEV void ldm_x4(uint32_t r[4], const void* p) {
uint32_t a = (uint32_t)__cvta_generic_to_shared(p);
asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0,%1,%2,%3}, [%4];\n"
: "=r"(r[0]), "=r"(r[1]), "=r"(r[2]), "=r"(r[3])
: "r"(a));
}
DEV void ldm_x4_t(uint32_t r[4], const void* p) {
uint32_t a = (uint32_t)__cvta_generic_to_shared(p);
asm volatile("ldmatrix.sync.aligned.m8n8.x4.trans.shared.b16 {%0,%1,%2,%3}, [%4];\n"
: "=r"(r[0]), "=r"(r[1]), "=r"(r[2]), "=r"(r[3])
: "r"(a));
}
DEV void cp_async16(void* smem, const void* gmem, int srcbytes) {
uint32_t s = (uint32_t)__cvta_generic_to_shared(smem);
asm volatile("cp.async.cg.shared.global [%0], [%1], 16, %2;\n" ::"r"(s), "l"(gmem),
"r"(srcbytes));
}
DEV void cp_commit() { asm volatile("cp.async.commit_group;\n"); }
template <int N>
DEV void cp_wait() {
asm volatile("cp.async.wait_group %0;\n" ::"n"(N));
}
DEV uint32_t pack2(float x, float y) {
__nv_bfloat162 p = __floats2bfloat162_rn(x, y);
uint32_t r;
memcpy(&r, &p, 4);
return r;
}
// =========================================================== A1: block sums
// Fallback for head dims the vectorised kernel cannot cover.
__global__ void blocksum_scalar_kernel(const bf16* __restrict__ K, float* __restrict__ BM,
int S, int D, int NB) {
const int bi = blockIdx.x;
const long long bh = blockIdx.y;
const bf16* Kp = K + bh * (long long)S * D;
const int s0 = bi * BS;
const int s1 = min(s0 + BS, S);
for (int d = threadIdx.x; d < D; d += blockDim.x) {
float acc = 0.f, comp = 0.f;
for (int j = s0; j < s1; ++j) {
float x = __bfloat162float(Kp[(long long)j * D + d]);
float y = x - comp;
float t = acc + y;
comp = (t - acc) - y;
acc = t;
}
BM[(bh * NB + bi) * (long long)D + d] = acc;
}
}
// One 8-wide column strip per thread, whole block in registers, one smem
// reduction at the end -- the rows-per-block sweep is what costs time here.
__global__ void blocksum_kernel(const bf16* __restrict__ K, float* __restrict__ BM,
int S, int D, int NB) {
constexpr int VEC = 8;
__shared__ float sh[256 * VEC];
const int bi = blockIdx.x;
const long long bh = blockIdx.y;
const bf16* Kp = K + bh * (long long)S * D;
const int s0 = bi * BS;
const int nrow = min(BS, S - s0);
const int CH = D >> 3;
const int RG = 256 / CH;
const int tid = threadIdx.x;
const int c = tid % CH, g = tid / CH;
float acc[VEC];
#pragma unroll
for (int i = 0; i < VEC; ++i) acc[i] = 0.f;
for (int j = g; j < nrow; j += RG) {
const uint4 raw = *(const uint4*)(Kp + (long long)(s0 + j) * D + c * VEC);
const __nv_bfloat162* h = (const __nv_bfloat162*)&raw;
#pragma unroll
for (int i = 0; i < VEC / 2; ++i) {
const float2 f = __bfloat1622float2(h[i]);
acc[2 * i] += f.x;
acc[2 * i + 1] += f.y;
}
}
float* dst = sh + g * D + c * VEC;
#pragma unroll
for (int i = 0; i < VEC; ++i) dst[i] = acc[i];
__syncthreads();
if (tid < D) {
float s = 0.f;
for (int r = 0; r < RG; ++r) s += sh[r * D + tid];
BM[(bh * NB + bi) * (long long)D + tid] = s;
}
}
// =========================================================== A2: top-n select
// Per query row, keep a running top-TOPN of the block importances. Rank is
// exactly the order the python reference sorts with -- (importance, block
// index) descending -- so ties hand the slot to the larger block index. The
// set is kept unsorted; each candidate simply evicts the current worst slot
// (located with a branch-free 8-way scan), which is all the caller needs since
// the selected blocks are consumed as a set.
// Running top-TOPN, kept sorted descending in 8 register pairs (no indexed
// array access, no branches: the whole insert is one compare-and-shift per
// slot). Candidates must be offered in non-decreasing block index order, which
// makes `nv >= v[k]` the exact tie-break the reference gets from
// `sorted(block_imp, reverse=True)` on (importance, index) tuples: an equally
// scoring later block owns the larger index and wins the slot.
DEV void top8_init(float v[TOPN], int ix[TOPN]) {
#pragma unroll
for (int k = 0; k < TOPN; ++k) {
v[k] = -CUDART_INF_F;
ix[k] = -1;
}
}
// b[k] = "nv belongs at or above slot k" is monotone in k, so the slot that
// actually receives nv is the one where b flips (b[k] && !b[k-1]); every slot
// below it shifts down by one and the previous last place drops out.
// Slots are visited bottom-up so that v[k-1] is still the pre-insert value
// when slot k needs it; bcur carries b[k] down from the slot above.
DEV void top8_insert(float v[TOPN], int ix[TOPN], float nv, int ni) {
bool bcur = (nv >= v[TOPN - 1]);
#pragma unroll
for (int k = TOPN - 1; k >= 0; --k) {
const bool take = (k > 0) && (nv >= v[k - 1]);
const float sv = take ? v[k - 1] : nv;
const int si = take ? ix[k - 1] : ni;
v[k] = bcur ? sv : v[k];
ix[k] = bcur ? si : ix[k];
bcur = take;
}
}
// v[] ends up sorted descending, so a candidate below the last slot can only
// shift itself -- nothing above it moves. Testing the last slot first is one
// compare against ~50 instructions of shift, and only top_n/candidates of the
// candidates ever clear it.
DEV void top8_try(float v[TOPN], int ix[TOPN], float nv, int ni) {
if (nv >= v[TOPN - 1]) top8_insert(v, ix, nv, ni);
}
DEV void ld_q16(float qv[16], const bf16* p) {
float4 a = *(const float4*)p;
float4 b = *(const float4*)(p + 8);
bf16 ha[8], hb[8];
memcpy(ha, &a, sizeof(ha));
memcpy(hb, &b, sizeof(hb));
#pragma unroll
for (int i = 0; i < 8; ++i) {
qv[i] = __bfloat162float(ha[i]);
qv[i + 8] = __bfloat162float(hb[i]);
}
}
// q . row, 16 head dims held in registers against one contiguous fp32 row.
// LDG selects the read-only global path; the diagonal block's row lives in
// shared memory, where __ldg is illegal.
template <bool LDG>
DEV float dot16(const float qv[16], const float* __restrict__ bm) {
const float4* p = (const float4*)bm;
float4 r0 = LDG ? __ldg(p) : *p;
float4 r1 = LDG ? __ldg(p + 1) : *(p + 1);
float4 r2 = LDG ? __ldg(p + 2) : *(p + 2);
float4 r3 = LDG ? __ldg(p + 3) : *(p + 3);
float bv[16];
memcpy(bv + 0, &r0, sizeof(r0));
memcpy(bv + 4, &r1, sizeof(r1));
memcpy(bv + 8, &r2, sizeof(r2));
memcpy(bv + 12, &r3, sizeof(r3));
float s0 = 0.f, s1 = 0.f, s2 = 0.f, s3 = 0.f;
#pragma unroll
for (int i = 0; i < 16; i += 4) {
s0 = fmaf(qv[i + 0], bv[i + 0], s0);
s1 = fmaf(qv[i + 1], bv[i + 1], s1);
s2 = fmaf(qv[i + 2], bv[i + 2], s2);
s3 = fmaf(qv[i + 3], bv[i + 3], s3);
}
return (s0 + s1) + (s2 + s3);
}
// Two rows per thread: the BM row is read once and both dots are formed from
// the same registers, which halves the LDG count the loop is bound on.
template <bool LDG>
DEV void dot16x2(const float qa[16], const float qb_[16], const float* __restrict__ bm,
float& oa, float& ob) {
const float4* p = (const float4*)bm;
float4 r0 = LDG ? __ldg(p) : *p;
float4 r1 = LDG ? __ldg(p + 1) : *(p + 1);
float4 r2 = LDG ? __ldg(p + 2) : *(p + 2);
float4 r3 = LDG ? __ldg(p + 3) : *(p + 3);
float bv[16];
memcpy(bv + 0, &r0, sizeof(r0));
memcpy(bv + 4, &r1, sizeof(r1));
memcpy(bv + 8, &r2, sizeof(r2));
memcpy(bv + 12, &r3, sizeof(r3));
float a0 = 0.f, a1 = 0.f, a2 = 0.f, a3 = 0.f;
float b0 = 0.f, b1 = 0.f, b2 = 0.f, b3 = 0.f;
#pragma unroll
for (int i = 0; i < 16; i += 4) {
a0 = fmaf(qa[i + 0], bv[i + 0], a0);
a1 = fmaf(qa[i + 1], bv[i + 1], a1);
a2 = fmaf(qa[i + 2], bv[i + 2], a2);
a3 = fmaf(qa[i + 3], bv[i + 3], a3);
b0 = fmaf(qb_[i + 0], bv[i + 0], b0);
b1 = fmaf(qb_[i + 1], bv[i + 1], b1);
b2 = fmaf(qb_[i + 2], bv[i + 2], b2);
b3 = fmaf(qb_[i + 3], bv[i + 3], b3);
}
oa = (a0 + a1) + (a2 + a3);
ob = (b0 + b1) + (b2 + b3);
}
template <int RPT, bool LDG>
DEV void dot16xN(const float qv[RPT][16], const float* __restrict__ bm, float o[RPT]) {
const float4* p = (const float4*)bm;
float4 r0 = LDG ? __ldg(p) : *p;
float4 r1 = LDG ? __ldg(p + 1) : *(p + 1);
float4 r2 = LDG ? __ldg(p + 2) : *(p + 2);
float4 r3 = LDG ? __ldg(p + 3) : *(p + 3);
float bv[16];
memcpy(bv + 0, &r0, sizeof(r0));
memcpy(bv + 4, &r1, sizeof(r1));
memcpy(bv + 8, &r2, sizeof(r2));
memcpy(bv + 12, &r3, sizeof(r3));
float a[RPT][4];
#pragma unroll
for (int r = 0; r < RPT; ++r)
#pragma unroll
for (int i = 0; i < 4; ++i) a[r][i] = 0.f;
#pragma unroll
for (int i = 0; i < 16; i += 4) {
#pragma unroll
for (int r = 0; r < RPT; ++r) {
a[r][0] = fmaf(qv[r][i + 0], bv[i + 0], a[r][0]);
a[r][1] = fmaf(qv[r][i + 1], bv[i + 1], a[r][1]);
a[r][2] = fmaf(qv[r][i + 2], bv[i + 2], a[r][2]);
a[r][3] = fmaf(qv[r][i + 3], bv[i + 3], a[r][3]);
}
}
#pragma unroll
for (int r = 0; r < RPT; ++r) o[r] = (a[r][0] + a[r][1]) + (a[r][2] + a[r][3]);
}
// =========================================================== B: sparse attend
template <int DT, int HS>
__global__ void __launch_bounds__(128) attend_kernel(const bf16* __restrict__ Q,
const bf16* __restrict__ K,
const bf16* __restrict__ V,
int* __restrict__ SEL,
bf16* __restrict__ O, int S, int NB,
float qscale, const float* __restrict__ BM,
int NBS, float scale) {
constexpr int KSTEPS = DT / 16; // head-dim steps
constexpr int NT = DT / 8; // accumulator n-tiles (8 head dims each)
constexpr int SUB = BS / HS; // pipeline stages per key block
constexpr int NTK = HS / 8; // score n-tiles per stage
extern __shared__ char smem[];
const int qb = blockIdx.x;
const long long bh = blockIdx.y;
const int s0 = qb * BS;
const int rem = min(BS, S - s0);
// K and V are staged HS keys at a time. The tile is what costs shared
// memory: at DT=128 the full 64-key double buffer is 64 KB, exactly one CTA
// per SM, and halving it restores three -- most of the difference between
// this kernel running at 45% and at 60% of the mma rate. At DT=64 the tile
// is half the size to begin with, so the whole block is staged at once and
// the extra stages would only add barriers.
bf16* Ks = (bf16*)smem; // 2*HS*DT
bf16* Vs = Ks + 2 * HS * DT; // 2*HS*DT
uint32_t* selbits = (uint32_t*)(Vs + 2 * HS * DT); // NB*2
const int tid = threadIdx.x;
const int warp = tid >> 5, lane = tid & 31;
const int row0 = warp * 16;
const int grp = lane >> 3, r8 = lane & 7;
const int row_a = row0 + (lane >> 2);
const int row_b = row_a + 8;
const int col0 = (lane & 3) * 2;
constexpr int VECPER = DT >> 3;
const bf16* Qb = Q + bh * (long long)S * DT;
const bf16* Kb = K + bh * (long long)S * DT;
const bf16* Vb = V + bh * (long long)S * DT;
// ---- fused select prologue -------------------------------------------------
// Block scores for this CTA's own 64 query rows, top-n, and the union bitmap.
// PS aliases the K/V staging tile exactly (BS*DT floats == 4*HS*DT bf16), so
// this costs no occupancy.
{
float* PS = (float*)Ks;
const bf16* Kp = K + bh * (long long)S * DT;
const float* BMp = BM + bh * (long long)NB * DT;
for (int i = tid; i < NB * 2; i += 128) selbits[i] = 0u;
__syncthreads();
// running key sums down each column, straight from global memory
if (tid < DT) {
const int d = tid;
float acc = 0.f, comp = 0.f;
for (int r = 0; r < BS; ++r) {
float x = 0.f;
if (s0 + r < S) x = __bfloat162float(Kp[(long long)(s0 + r) * DT + d]);
float y = x - comp;
float t = acc + y;
comp = (t - acc) - y;
acc = t;
PS[r * DT + d] = acc;
}
}
__syncthreads();
constexpr int TPQ = (DT == 64) ? 4 : 8;
constexpr int RPT = (DT == 64) ? 2 : 4;
const int p = tid / TPQ, c = tid % TPQ;
const int q0 = RPT * p;
const int zoff = c * 16;
float qv[RPT][16];
float tv[RPT][TOPN];
int ti[RPT][TOPN];
float diag[RPT];
bool ok[RPT];
#pragma unroll
for (int r = 0; r < RPT; ++r) {
const int q = q0 + r;
ok[r] = (s0 + q) < S;
if (ok[r])
ld_q16(qv[r], Q + bh * (long long)S * DT + (long long)(s0 + q) * DT + zoff);
else
#pragma unroll
for (int i = 0; i < 16; ++i) qv[r][i] = 0.f;
top8_init(tv[r], ti[r]);
}
{
float a[RPT];
#pragma unroll
for (int r = 0; r < RPT; ++r) a[r] = dot16<false>(qv[r], PS + (q0 + r) * DT + zoff);
#pragma unroll
for (int o = 1; o < TPQ; o <<= 1)
#pragma unroll
for (int r = 0; r < RPT; ++r) a[r] += __shfl_xor_sync(0xffffffffu, a[r], o);
#pragma unroll
for (int r = 0; r < RPT; ++r) diag[r] = a[r] * scale / (float)(q0 + r + 1);
}
{
const float invs = scale * 0.015625f;
const float* bmp = BMp + zoff;
#pragma unroll 4
for (int bb = 0; bb < qb; ++bb, bmp += DT) {
float v[RPT];
dot16xN<RPT, true>(qv, bmp, v);
#pragma unroll
for (int o = 1; o < TPQ; o <<= 1)
#pragma unroll
for (int r = 0; r < RPT; ++r) v[r] += __shfl_xor_sync(0xffffffffu, v[r], o);
#pragma unroll
for (int r = 0; r < RPT; ++r) top8_try(tv[r], ti[r], v[r] * invs, bb);
}
}
#pragma unroll
for (int r = 0; r < RPT; ++r)
if (ok[r]) top8_insert(tv[r], ti[r], diag[r], qb);
if (c == 0) {
#pragma unroll
for (int r = 0; r < RPT; ++r) {
const int rq = q0 + r;
int* sp = SEL + ((bh * (long long)NBS + s0 + rq) * TOPN);
if (ok[r]) {
#pragma unroll
for (int k = 0; k < TOPN; ++k) {
sp[k] = ti[r][k];
const int id = ti[r][k];
if (id >= 0 && id < NB) atomicOr(&selbits[id * 2 + (rq >> 5)], 1u << (rq & 31));
}
} else {
#pragma unroll
for (int k = 0; k < TOPN; ++k) sp[k] = -1;
}
}
}
// PS dies here; the staging tile is free to be overwritten.
__syncthreads();
}
// Q A-fragments: constant for the whole kernel. a0/a1 live in rows r/r+8 at
// the low half of the 16-dim step, a2/a3 at the high half -- exactly the
// ldmatrix.x4 packing the B fragments below use.
uint32_t qa[KSTEPS][4];
{
const int g = lane >> 2, tt = (lane & 3) * 2;
const int ra = s0 + row0 + g;
const bool ok0 = (ra < S), ok1 = (ra + 8 < S);
#pragma unroll
for (int kk = 0; kk < KSTEPS; ++kk) {
const bf16* p = Qb + (long long)ra * DT + kk * 16 + tt;
qa[kk][0] = ok0 ? *(const uint32_t*)p : 0u;
qa[kk][1] = ok1 ? *(const uint32_t*)(p + 8 * DT) : 0u;
qa[kk][2] = ok0 ? *(const uint32_t*)(p + 8) : 0u;
qa[kk][3] = ok1 ? *(const uint32_t*)(p + 8 * DT + 8) : 0u;
}
}
float of[NT][4];
#pragma unroll
for (int i = 0; i < NT; ++i)
#pragma unroll
for (int j = 0; j < 4; ++j) of[i][j] = 0.f;
float mrow[2] = {-1e30f, -1e30f};
float lrow[2] = {0.f, 0.f};
const int vecper = DT >> 3;
const int nchunk = HS * vecper;
// Stage s carries key block kb = s/2, second half when s is odd.
auto issue = [&](int s, int buf) {
const int kb = s / SUB, half = s - kb * SUB;
bf16* kd = Ks + buf * HS * DT;
bf16* vd = Vs + buf * HS * DT;
const int j0 = kb * BS + half * HS;
if (j0 + HS <= S) {
// No row of this stage runs past the end of the sequence, so every
// address is affine in the lane index: with nchunk a compile-time
// constant over a fixed 128-lane stride the loop unrolls into plain
// cp.asyncs with precomputed offsets. Only the ragged final block
// needs the general (variable-address) form below.
const bf16* gp = Kb + (long long)j0 * DT;
const bf16* gq = Vb + (long long)j0 * DT;
#pragma unroll
for (int i = tid; i < nchunk; i += 128) {
int row = i / vecper, c8 = i - row * vecper;
const int sm = row * DT + ((c8 ^ (row & 7)) << 3);
const long long go = (long long)row * DT + c8 * 8;
cp_async16(kd + sm, gp + go, 16);
cp_async16(vd + sm, gq + go, 16);
}
} else {
for (int i = tid; i < nchunk; i += 128) {
int row = i / vecper, c8 = i - row * vecper;
int gj = j0 + row;
int ok = (gj < S) ? 16 : 0;
int gjc = min(gj, S - 1);
long long off = (long long)gjc * DT + c8 * 8;
cp_async16(kd + row * DT + ((c8 ^ (row & 7)) << 3), Kb + off, ok);
cp_async16(vd + row * DT + ((c8 ^ (row & 7)) << 3), Vb + off, ok);
}
}
cp_commit();
};
const int nstage = SUB * (qb + 1);
issue(0, 0);
for (int s = 0; s < nstage; ++s) {
// One barrier per stage. The wait is <0> because exactly one group is
// ever in flight; the barrier then does double duty -- it publishes the
// landed cp.async data and it separates the reads of buffer (s+1)&1 at
// stage s-1 from the issue below, which overwrites that same buffer.
cp_wait<0>();
__syncthreads();
if (s + 1 < nstage) issue(s + 1, (s + 1) & 1);
const int kb = s / SUB, half = s - kb * SUB;
const bf16* kd = Ks + (s & 1) * HS * DT;
const bf16* vd = Vs + (s & 1) * HS * DT;
const uint32_t sela0 = (selbits[kb * 2 + (row_a >> 5)] >> (row_a & 31)) & 1u;
const uint32_t selb0 = (selbits[kb * 2 + (row_b >> 5)] >> (row_b & 31)) & 1u;
// A tile strictly older than the previous one reaches a query row *only*
// through the top-n selection. If no row in this warp selected this block
// the entire tile contributes nothing, so skip it -- loads, mma, softmax.
if (!(kb < qb - 1 && !__any_sync(0xffffffffu, sela0 | selb0))) {
// ---------------- S = Q K^T ----------------
float sf[NTK][4];
#pragma unroll
for (int n = 0; n < NTK; ++n)
#pragma unroll
for (int j = 0; j < 4; ++j) sf[n][j] = 0.f;
#pragma unroll
for (int kk = 0; kk < KSTEPS; ++kk) {
#pragma unroll
for (int np = 0; np < NTK / 2; ++np) {
uint32_t bb[4];
const int nk = np * 16 + r8 + ((grp & 1) ? 8 : 0);
const int nc = kk * 2 + ((grp & 2) ? 1 : 0);
ldm_x4(bb, kd + nk * DT + ((nc ^ (nk & 7)) << 3));
mma16816(sf[2 * np], qa[kk], bb[0], bb[2]);
mma16816(sf[2 * np + 1], qa[kk], bb[1], bb[3]);
}
}
// ---------------- scale + NSA mask ----------------
{
const uint32_t sela = sela0, selb = selb0;
const float NEG = -CUDART_INF_F;
if (kb < qb - 1) {
// every key of the block is causal; only the selection bit applies
const float ma = sela ? 0.f : NEG;
const float mb = selb ? 0.f : NEG;
#pragma unroll
for (int n = 0; n < NTK; ++n) {
sf[n][0] = sf[n][0] * qscale + ma;
sf[n][1] = sf[n][1] * qscale + ma;
sf[n][2] = sf[n][2] * qscale + mb;
sf[n][3] = sf[n][3] * qscale + mb;
}
} else if (kb == qb - 1) {
// sliding window reaches into this block: keep col >= row+1
const int wa = row_a + 1, wb = row_b + 1;
#pragma unroll
for (int n = 0; n < NTK; ++n) {
const int cc = half * HS + n * 8 + col0;
sf[n][0] = sf[n][0] * qscale + ((sela || cc >= wa) ? 0.f : NEG);
sf[n][1] = sf[n][1] * qscale + ((sela || cc + 1 >= wa) ? 0.f : NEG);
sf[n][2] = sf[n][2] * qscale + ((selb || cc >= wb) ? 0.f : NEG);
sf[n][3] = sf[n][3] * qscale + ((selb || cc + 1 >= wb) ? 0.f : NEG);
}
} else {
// diagonal block: causal, window, and the S tail
const int ra = row_a, rb = row_b;
#pragma unroll
for (int n = 0; n < NTK; ++n) {
const int cc = half * HS + n * 8 + col0;
sf[n][0] = sf[n][0] * qscale +
((cc <= ra && (sela || cc > ra - BS) && cc < rem) ? 0.f : NEG);
sf[n][1] = sf[n][1] * qscale +
((cc + 1 <= ra && (sela || cc + 1 > ra - BS) && cc + 1 < rem) ? 0.f : NEG);
sf[n][2] = sf[n][2] * qscale +
((cc <= rb && (selb || cc > rb - BS) && cc < rem) ? 0.f : NEG);
sf[n][3] = sf[n][3] * qscale +
((cc + 1 <= rb && (selb || cc + 1 > rb - BS) && cc + 1 < rem) ? 0.f
: NEG);
}
}
}
// ---------------- online softmax (base 2) ----------------
float m0, m1;
if constexpr (NTK == 8) {
float a0 = fmaxf(sf[0][0], sf[0][1]), b0 = fmaxf(sf[1][0], sf[1][1]);
float c0 = fmaxf(sf[2][0], sf[2][1]), d0 = fmaxf(sf[3][0], sf[3][1]);
float e0 = fmaxf(sf[4][0], sf[4][1]), f0 = fmaxf(sf[5][0], sf[5][1]);
float g0 = fmaxf(sf[6][0], sf[6][1]), h0 = fmaxf(sf[7][0], sf[7][1]);
float a1 = fmaxf(sf[0][2], sf[0][3]), b1 = fmaxf(sf[1][2], sf[1][3]);
float c1 = fmaxf(sf[2][2], sf[2][3]), d1 = fmaxf(sf[3][2], sf[3][3]);
float e1 = fmaxf(sf[4][2], sf[4][3]), f1 = fmaxf(sf[5][2], sf[5][3]);
float g1 = fmaxf(sf[6][2], sf[6][3]), h1 = fmaxf(sf[7][2], sf[7][3]);
m0 = fmaxf(fmaxf(fmaxf(a0, b0), fmaxf(c0, d0)), fmaxf(fmaxf(e0, f0), fmaxf(g0, h0)));
m1 = fmaxf(fmaxf(fmaxf(a1, b1), fmaxf(c1, d1)), fmaxf(fmaxf(e1, f1), fmaxf(g1, h1)));
} else {
float a0 = fmaxf(sf[0][0], sf[0][1]), b0 = fmaxf(sf[1][0], sf[1][1]);
float c0 = fmaxf(sf[2][0], sf[2][1]), d0 = fmaxf(sf[3][0], sf[3][1]);
float a1 = fmaxf(sf[0][2], sf[0][3]), b1 = fmaxf(sf[1][2], sf[1][3]);
float c1 = fmaxf(sf[2][2], sf[2][3]), d1 = fmaxf(sf[3][2], sf[3][3]);
m0 = fmaxf(fmaxf(a0, b0), fmaxf(c0, d0));
m1 = fmaxf(fmaxf(a1, b1), fmaxf(c1, d1));
}
m0 = fmaxf(m0, __shfl_xor_sync(0xffffffffu, m0, 1));
m0 = fmaxf(m0, __shfl_xor_sync(0xffffffffu, m0, 2));
m1 = fmaxf(m1, __shfl_xor_sync(0xffffffffu, m1, 1));
m1 = fmaxf(m1, __shfl_xor_sync(0xffffffffu, m1, 2));
const float nm0 = fmaxf(mrow[0], m0);
const float nm1 = fmaxf(mrow[1], m1);
const float al0 = exp2_approx(mrow[0] - nm0);
const float al1 = exp2_approx(mrow[1] - nm1);
#pragma unroll
for (int n = 0; n < NTK; ++n) {
sf[n][0] = exp2_approx(sf[n][0] - nm0);
sf[n][1] = exp2_approx(sf[n][1] - nm0);
sf[n][2] = exp2_approx(sf[n][2] - nm1);
sf[n][3] = exp2_approx(sf[n][3] - nm1);
}
float r0, r1;
if constexpr (NTK == 8) {
r0 = (((sf[0][0] + sf[0][1]) + (sf[1][0] + sf[1][1])) +
((sf[2][0] + sf[2][1]) + (sf[3][0] + sf[3][1]))) +
(((sf[4][0] + sf[4][1]) + (sf[5][0] + sf[5][1])) +
((sf[6][0] + sf[6][1]) + (sf[7][0] + sf[7][1])));
r1 = (((sf[0][2] + sf[0][3]) + (sf[1][2] + sf[1][3])) +
((sf[2][2] + sf[2][3]) + (sf[3][2] + sf[3][3]))) +
(((sf[4][2] + sf[4][3]) + (sf[5][2] + sf[5][3])) +
((sf[6][2] + sf[6][3]) + (sf[7][2] + sf[7][3])));
} else {
r0 = ((sf[0][0] + sf[0][1]) + (sf[1][0] + sf[1][1])) +
((sf[2][0] + sf[2][1]) + (sf[3][0] + sf[3][1]));
r1 = ((sf[0][2] + sf[0][3]) + (sf[1][2] + sf[1][3])) +
((sf[2][2] + sf[2][3]) + (sf[3][2] + sf[3][3]));
}
r0 += __shfl_xor_sync(0xffffffffu, r0, 1);
r0 += __shfl_xor_sync(0xffffffffu, r0, 2);
r1 += __shfl_xor_sync(0xffffffffu, r1, 1);
r1 += __shfl_xor_sync(0xffffffffu, r1, 2);
lrow[0] = lrow[0] * al0 + r0;
lrow[1] = lrow[1] * al1 + r1;
// al == exp2(m_old - m_new) is exactly 1 whenever this block did not raise
// the running max of any row in the warp, and then the whole NT*4 rescale
// of the PV accumulator is a multiply by one. The vote keeps it off the
// serial path between the QK and PV halves of the iteration.
if (__any_sync(0xffffffffu, (al0 != 1.f) | (al1 != 1.f))) {
#pragma unroll
for (int i = 0; i < NT; ++i) {
of[i][0] *= al0;
of[i][1] *= al0;
of[i][2] *= al1;
of[i][3] *= al1;
}
}
mrow[0] = nm0;
mrow[1] = nm1;
// ---------------- O += P V ----------------
// reduction runs over the staged keys, so there are HS/16 k-steps here --
// distinct from the DT/16 head-dim steps of the QK^T pass above.
#pragma unroll
for (int kk = 0; kk < NTK / 2; ++kk) {
uint32_t pa[4];
pa[0] = pack2(sf[2 * kk][0], sf[2 * kk][1]);
pa[1] = pack2(sf[2 * kk][2], sf[2 * kk][3]);
pa[2] = pack2(sf[2 * kk + 1][0], sf[2 * kk + 1][1]);
pa[3] = pack2(sf[2 * kk + 1][2], sf[2 * kk + 1][3]);
#pragma unroll
for (int j = 0; j < NT / 2; ++j) {
uint32_t vb[4];
const int vk = kk * 16 + r8 + ((grp & 1) ? 8 : 0);
const int vc = j * 2 + ((grp & 2) ? 1 : 0);
ldm_x4_t(vb, vd + vk * DT + ((vc ^ (vk & 7)) << 3));
mma16816(of[2 * j], pa, vb[0], vb[1]);
mma16816(of[2 * j + 1], pa, vb[2], vb[3]);
}
}
}
}
// ---------------- epilogue ----------------
{
bf16* Op = O + bh * (long long)S * DT;
const float inv0 = 1.f / lrow[0];
const float inv1 = 1.f / lrow[1];
#pragma unroll
for (int i = 0; i < NT; ++i) {
if (row_a < rem)
*(__nv_bfloat162*)(Op + (long long)(s0 + row_a) * DT + i * 8 + col0) =
__floats2bfloat162_rn(of[i][0] * inv0, of[i][1] * inv0);
if (row_b < rem)
*(__nv_bfloat162*)(Op + (long long)(s0 + row_b) * DT + i * 8 + col0) =
__floats2bfloat162_rn(of[i][2] * inv1, of[i][3] * inv1);
}
}
}
// =========================================================== host
torch::Tensor nsa_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v,
torch::Tensor BM, torch::Tensor SEL) {
TORCH_CHECK(q.is_cuda() && k.is_cuda() && v.is_cuda(), "inputs must be CUDA");
TORCH_CHECK(q.dim() == 4, "q must be (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), D = qc.size(3);
TORCH_CHECK(D <= 128, "D must be <= 128");
const int DT = (D <= 64) ? 64 : 128;
const long long BH = (long long)B * H;
const int NB = (S + BS - 1) / BS;
const int NBS = NB * BS;
const float scale = 1.0f / sqrtf((float)D);
auto opts = qc.options();
torch::Tensor qp, kp, vp;
if (D == DT) {
qp = qc;
kp = kc;
vp = vc;
} else {
qp = torch::zeros({B, H, S, DT}, opts);
kp = torch::zeros({B, H, S, DT}, opts);
vp = torch::zeros({B, H, S, DT}, opts);
qp.narrow(3, 0, D).copy_(qc);
kp.narrow(3, 0, D).copy_(kc);
vp.narrow(3, 0, D).copy_(vc);
}
// BM / SEL are pure scratch, rewritten in full every call, and are owned by
// the caller so they survive across invocations. O is written in full by the
// attention kernel (every row below S, every column) so it needs no clear.
TORCH_CHECK(BM.numel() >= BH * (long long)NB * DT && BM.is_cuda() &&
SEL.numel() >= BH * (long long)NBS * TOPN && SEL.is_cuda(),
"scratch tensors too small");
auto O = torch::empty({B, H, S, DT}, opts);
const bf16* qptr = (const bf16*)qp.data_ptr();
const bf16* kptr = (const bf16*)kp.data_ptr();
const bf16* vptr = (const bf16*)vp.data_ptr();
const dim3 g1(NB, (unsigned)BH);
auto stream = at::cuda::getCurrentCUDAStream();
const int ch = DT >> 3;
if ((DT & 7) == 0 && ch * 8 == DT && ch <= 256 && (256 % ch) == 0)
blocksum_kernel<<<g1, 256, 0, stream>>>(kptr, BM.data_ptr<float>(), S, DT, NB);
else
blocksum_scalar_kernel<<<g1, 256, 0, stream>>>(kptr, BM.data_ptr<float>(), S, DT, NB);
const float qscale = scale * 1.4426950408889634f; // log2(e)
auto launch = [&](auto kern, size_t sh) {
cudaFuncSetAttribute(kern, cudaFuncAttributeMaxDynamicSharedMemorySize, (int)sh);
kern<<<g1, 128, sh, stream>>>(qptr, kptr, vptr, SEL.data_ptr<int>(),
(bf16*)O.data_ptr(), S, NB, qscale, BM.data_ptr<float>(),
NBS, scale);
};
if (DT == 64) {
const size_t sh3 = (size_t)(4 * 64 * DT) * 2 + (size_t)NB * 8;
launch(attend_kernel<64, 64>, sh3);
} else {
const size_t sh3 = (size_t)(4 * 32 * DT) * 2 + (size_t)NB * 8;
launch(attend_kernel<128, 32>, sh3);
}
auto out = O.narrow(3, 0, D).contiguous();
return out;
}
"""
def _load():
return load_inline(
name="nsa_sm120",
cpp_sources="torch::Tensor nsa_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor BM, torch::Tensor SEL);",
cuda_sources=_CUDA_SRC,
functions=["nsa_forward"],
extra_cuda_cflags=["-O3", "-arch=sm_120a", "--expt-relaxed-constexpr"],
verbose=False,
)
_MOD = None
def _get():
global _MOD
if _MOD is None:
_MOD = _load()
return _MOD
def _slot_refcount():
"""Refcount of a graph-slot output that no other code holds, measured in the
same `for g, out in slots` context that Model.forward uses to test it."""
probe = torch.zeros(1, dtype=torch.bfloat16)
holder = [(None, probe)]
for _g, out in holder:
return sys.getrefcount(out)
return 3
class Model(nn.Module):
"""Same interface as reference.Model.
The three kernels are launched every call; a CUDA-graph replay is used when
the caller re-uses the same input tensors (as the timing loop does) because
the per-launch host latency of the three-kernel sequence is a large share of
wall time on short sequences.
"""
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))
self._scratch = None
self._graphs = {}
self._graphs_off = False
self._ref_base = None
def _get_scratch(self, q):
B, H, S, D = q.shape
key = (B, H, S, D)
ws = self._scratch
if ws is not None and ws[0] == key:
return ws[1], ws[2]
DT = 64 if D <= 64 else 128
nb = (S + BLOCK_SIZE - 1) // BLOCK_SIZE
bm = torch.empty((B * H, nb, DT), dtype=torch.float32, device=q.device)
sel = torch.empty((B * H, nb * BLOCK_SIZE, TOP_N_BLOCKS),
dtype=torch.int32, device=q.device)
self._scratch = (key, bm, sel)
return bm, sel
def _key(self, q, k, v):
return (
q.data_ptr(), k.data_ptr(), v.data_ptr(),
tuple(q.shape), tuple(k.shape), tuple(v.shape),
tuple(q.stride()), tuple(k.stride()), tuple(v.stride()),
)
def _slot_free(self, slots, idx):
# A replay overwrites the captured output in place, so it is only safe
# when the caller kept no reference to the previous result. The
# baseline (refcount of a slot nobody holds) is measured in this exact
# loop context rather than assumed, so the test stays correct if the
# interpreter accounts references differently.
if self._ref_base is None:
self._ref_base = _slot_refcount()
out = slots[idx][1]
return sys.getrefcount(out) <= self._ref_base
def _plain_forward(self, q, k, v, bm, sel):
o = _get().nsa_forward(q, k, v, bm, sel)
return o.to(torch.bfloat16)
def _capture(self, q, k, v, key, ref):
try:
s = torch.cuda.Stream()
s.wait_stream(torch.cuda.current_stream())
with torch.cuda.stream(s):
for _ in range(2):
self._plain_forward(q, k, v, *self._get_scratch(q))
torch.cuda.current_stream().wait_stream(s)
g = torch.cuda.CUDAGraph()
with torch.cuda.graph(g):
out = self._plain_forward(q, k, v, *self._get_scratch(q))
# Validate: an incomplete capture would replay stale memory.
g.replay()
if not torch.equal(out, ref):
self._graphs_off = True
return
self._graphs.setdefault(key, []).append((g, out))
except Exception:
# Any capture failure just disables the replay fast path.
self._graphs_off = True
def forward(self, q, k, v):
bm, sel = self._get_scratch(q)
key = self._key(q, k, v)
slots = self._graphs.get(key)
if slots:
for i in range(len(slots)):
if self._slot_free(slots, i):
g, out = slots[i]
g.replay()
return out
o = self._plain_forward(q, k, v, bm, sel)
# First call for this key: build one slot. Later calls whose result is
# still held by the caller (so no slot is reusable) build a second one,
# which lets a `y = model(...)` loop alternate between two buffers.
if not self._graphs_off and key not in self._graphs:
self._capture(q, k, v, key, o)
elif (
not self._graphs_off
and slots is not None
and len(slots) < 2
and len(self._graphs) <= 8
):
self._capture(q, k, v, key, o)
return o
20260910_150236_deepseek-claude_deepseek-flash_02_deepseek_nsa