KernelBench cuda · RTX PRO 6000
DeepSeek NSA Grok 4.7
0.984 msgeomean latency across six shapes · lower is better
manually audited: clean
Grok 4.7 hand-wrote an SM120 NSA kernel in CUDA C++ with nvcuda::wmma bf16 m16n16k16 fragments, no cuBLAS and no torch op in the timed path: a block-sum kernel, a fused fp32 block-scoring plus top-8 kernel, a WMMA tile attention kernel for D=64 short sequences and a register-resident warp-per-query kernel for the long shapes and D=128. It computes the real op on every graded shape, verified against the reference at S=8192 and S=8191. It is also slow: 0.1002 of dense bf16 peak, against 0.5019 for the DeepSeek V4.1 Flash PTX kernel on the same problem.
harnessgrokagent session58mtotal wall59mcheck6mbenchmark2soutput tokens—gpu-lock wait0sgpu-lock held6mregimecompute
Per-shape latencymilliseconds are the measured result; the dense-equivalent roofline is not a useful ceiling
1×16×2048×640.508 ms
1×16×4127×641.595 ms
1×8×8192×641.781 ms
1×8×8191×1282.308 ms
4×8×1024×640.289 ms
2×8×3000×640.945 ms
geomean(0.508 ms · 1.595 ms · 1.781 ms · 2.308 ms · 0.289 ms · 0.945 ms) = 0.984 ms
Kernel source (redacted)
"""DeepSeek NSA-inspired sparse attention, CUDA SM120.
Block importance is q · sum(K_block) / (count * sqrt(D)), with smaller block
index winning ties. Top-8 is unioned with the causal sliding window. Attention
over a query tile uses WMMA bf16 tensor cores (m16n16k16) and online softmax.
"""
from __future__ import annotations
import os
import torch
import torch.nn as nn
from torch.utils.cpp_extension import load_inline
_BLOCK = 64
_CPP = r"""
#include <torch/extension.h>
void nsa_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v,
torch::Tensor block_sum, torch::Tensor topk, torch::Tensor out);
"""
_CUDA = r"""
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <cuda_bf16.h>
#include <mma.h>
#include <cfloat>
#include <algorithm>
#include <cstdint>
using namespace nvcuda::wmma;
static inline __device__ float warp_sum(float v) {
#pragma unroll
for (int off = 16; off > 0; off >>= 1)
v += __shfl_xor_sync(0xffffffff, v, off);
return v;
}
#define CONSIDER(score, bi) \
do { \
float _s = (score); \
int _b = (bi); \
if (!(_s < ts7 || (_s == ts7 && _b > tb7))) { \
if (_s > ts0 || (_s == ts0 && _b < tb0)) { \
ts7 = ts6; tb7 = tb6; ts6 = ts5; tb6 = tb5; ts5 = ts4; tb5 = tb4; \
ts4 = ts3; tb4 = tb3; ts3 = ts2; tb3 = tb2; ts2 = ts1; tb2 = tb1; \
ts1 = ts0; tb1 = tb0; ts0 = _s; tb0 = _b; \
} else if (_s > ts1 || (_s == ts1 && _b < tb1)) { \
ts7 = ts6; tb7 = tb6; ts6 = ts5; tb6 = tb5; ts5 = ts4; tb5 = tb4; \
ts4 = ts3; tb4 = tb3; ts3 = ts2; tb3 = tb2; ts2 = ts1; tb2 = tb1; \
ts1 = _s; tb1 = _b; \
} else if (_s > ts2 || (_s == ts2 && _b < tb2)) { \
ts7 = ts6; tb7 = tb6; ts6 = ts5; tb6 = tb5; ts5 = ts4; tb5 = tb4; \
ts4 = ts3; tb4 = tb3; ts3 = ts2; tb3 = tb2; ts2 = _s; tb2 = _b; \
} else if (_s > ts3 || (_s == ts3 && _b < tb3)) { \
ts7 = ts6; tb7 = tb6; ts6 = ts5; tb6 = tb5; ts5 = ts4; tb5 = tb4; \
ts4 = ts3; tb4 = tb3; ts3 = _s; tb3 = _b; \
} else if (_s > ts4 || (_s == ts4 && _b < tb4)) { \
ts7 = ts6; tb7 = tb6; ts6 = ts5; tb6 = tb5; ts5 = ts4; tb5 = tb4; \
ts4 = _s; tb4 = _b; \
} else if (_s > ts5 || (_s == ts5 && _b < tb5)) { \
ts7 = ts6; tb7 = tb6; ts6 = ts5; tb6 = tb5; ts5 = _s; tb5 = _b; \
} else if (_s > ts6 || (_s == ts6 && _b < tb6)) { \
ts7 = ts6; tb7 = tb6; ts6 = _s; tb6 = _b; \
} else { \
ts7 = _s; tb7 = _b; \
} \
} \
} while (0)
template <int D>
__global__ void __launch_bounds__(32, 8)
block_sum_kernel(const __nv_bfloat16* __restrict__ k, float* __restrict__ bsum,
int S, int n_blocks) {
constexpr int V = D / 32;
const int bi = blockIdx.x;
const int bh = blockIdx.y;
const int lane = threadIdx.x;
const int s0 = bi * 64;
const int slen = min(64, S - s0);
const __nv_bfloat16* base = k + (static_cast<long long>(bh) * S + s0) * D;
float accv[V];
#pragma unroll
for (int i = 0; i < V; ++i) accv[i] = 0.f;
for (int r = 0; r < slen; ++r) {
const __nv_bfloat16* row = base + static_cast<long long>(r) * D + lane * V;
#pragma unroll
for (int i = 0; i < V; i += 2) {
__nv_bfloat162 kv = *reinterpret_cast<const __nv_bfloat162*>(row + i);
float2 f = __bfloat1622float2(kv);
accv[i] += f.x;
accv[i + 1] += f.y;
}
}
float* dst = bsum + (static_cast<long long>(bh) * n_blocks + bi) * D + lane * V;
#pragma unroll
for (int i = 0; i < V; i += 2)
*reinterpret_cast<float2*>(dst + i) = make_float2(accv[i], accv[i + 1]);
}
template <int D>
__device__ __forceinline__ float dot_bf16(const __nv_bfloat16* row, const float* q, int lane) {
constexpr int V = D / 32;
float partial = 0.f;
const __nv_bfloat16* p = row + lane * V;
#pragma unroll
for (int i = 0; i < V; i += 2) {
__nv_bfloat162 kv = *reinterpret_cast<const __nv_bfloat162*>(p + i);
float2 f = __bfloat1622float2(kv);
partial = fmaf(q[i], f.x, partial);
partial = fmaf(q[i + 1], f.y, partial);
}
return warp_sum(partial);
}
template <int D>
__device__ __forceinline__ float dot_f32(const float* row, const float* q, int lane) {
constexpr int V = D / 32;
float partial = 0.f;
const float* p = row + lane * V;
#pragma unroll
for (int i = 0; i < V; i += 2) {
float2 f = *reinterpret_cast<const float2*>(p + i);
partial = fmaf(q[i], f.x, partial);
partial = fmaf(q[i + 1], f.y, partial);
}
return warp_sum(partial);
}
template <int D>
__global__ void __launch_bounds__(128, 8)
score_kernel(const __nv_bfloat16* __restrict__ q, const __nv_bfloat16* __restrict__ k,
const float* __restrict__ block_sum, short* __restrict__ topk,
int S, int n_blocks, float scale) {
constexpr int V = D / 32;
const int lane = threadIdx.x & 31;
const int warp = threadIdx.x >> 5;
const int warps = blockDim.x >> 5;
const int bh = blockIdx.y;
const int t = blockIdx.x * warps + warp;
if (t >= S) return;
const long long head = static_cast<long long>(bh) * S;
const __nv_bfloat16* qrow = q + (head + t) * D;
const __nv_bfloat16* kbase = k + head * D;
const float* bs = block_sum + static_cast<long long>(bh) * n_blocks * D;
float qreg[V];
#pragma unroll
for (int i = 0; i < V; i += 2) {
__nv_bfloat162 qq = *reinterpret_cast<const __nv_bfloat162*>(qrow + lane * V + i);
float2 f = __bfloat1622float2(qq);
qreg[i] = f.x;
qreg[i + 1] = f.y;
}
float ts0 = -INFINITY, ts1 = -INFINITY, ts2 = -INFINITY, ts3 = -INFINITY;
float ts4 = -INFINITY, ts5 = -INFINITY, ts6 = -INFINITY, ts7 = -INFINITY;
int tb0 = 32767, tb1 = 32767, tb2 = 32767, tb3 = 32767;
int tb4 = 32767, tb5 = 32767, tb6 = 32767, tb7 = 32767;
const int cb = t >> 6;
const float inv_block = 1.0f / 64.0f;
for (int bi = 0; bi < cb; ++bi) {
float dot = dot_f32<D>(bs + static_cast<long long>(bi) * D, qreg, lane);
CONSIDER(dot * scale * inv_block, bi);
}
float psum = 0.f;
const int block_start = cb << 6;
const int count = t - block_start + 1;
for (int j = block_start; j <= t; ++j)
psum += dot_bf16<D>(kbase + static_cast<long long>(j) * D, qreg, lane);
CONSIDER(psum * scale / static_cast<float>(count), cb);
if (lane == 0) {
short* dst = topk + (head + t) * 8;
dst[0] = (short)tb0; dst[1] = (short)tb1; dst[2] = (short)tb2; dst[3] = (short)tb3;
dst[4] = (short)tb4; dst[5] = (short)tb5; dst[6] = (short)tb6; dst[7] = (short)tb7;
}
}
// smem: ksm | vsm | qsm | acc | ml | per-warp scores | per-warp P
template <int D, int NW>
struct TCLay {
int ksm, vsm, qsm, acc, ml, sc, p, bytes;
};
template <int D, int NW>
__host__ __device__ TCLay<D, NW> tc_layout() {
constexpr int TILE = NW * 16;
TCLay<D, NW> o{};
int p = 0;
auto align16 = [](int x) { return (x + 15) & ~15; };
o.ksm = p; p = align16(p + 64 * D * 2);
o.vsm = p; p = align16(p + 64 * D * 2);
o.qsm = p; p = align16(p + TILE * D * 2);
o.acc = p; p = align16(p + TILE * D * 4);
o.ml = p; p = align16(p + TILE * 2 * 4);
// scores and PV-tmp alias (a warp finishes scores before PV)
int sc_bytes = 16 * 64 * 4;
int pv_bytes = 16 * D * 4;
int scratch = sc_bytes > pv_bytes ? sc_bytes : pv_bytes;
o.sc = p; p = align16(p + NW * scratch);
o.p = p; p = align16(p + NW * 16 * 64 * 2);
o.bytes = p;
return o;
}
// One CTA owns TILE consecutive queries and loops key blocks. Only queries that
// actually touch the block are packed into WMMA groups of 16, so long-context
// top-8 does not pay for the queries that skipped the block.
template <int D, int TILE>
__global__ void __launch_bounds__(128, 1)
tc_compact_kernel(const __nv_bfloat16* __restrict__ q,
const __nv_bfloat16* __restrict__ k,
const __nv_bfloat16* __restrict__ v,
const short* __restrict__ topk,
__nv_bfloat16* __restrict__ o,
int S, int n_blocks, float scale) {
constexpr int THREADS = 128;
constexpr int BLOCK = 64;
constexpr int WINDOW = 64;
const int lane = threadIdx.x & 31;
const int warp = threadIdx.x >> 5;
const int bh = blockIdx.y;
const int q0 = blockIdx.x * TILE;
if (q0 >= S) return;
const int tile_q = min(TILE, S - q0);
const int tile_max_t = q0 + tile_q - 1;
extern __shared__ __align__(16) char smem[];
auto align16 = [](int x) { return (x + 15) & ~15; };
int off = 0;
__nv_bfloat16* ksm = reinterpret_cast<__nv_bfloat16*>(smem + off);
off = align16(off + 64 * D * 2);
__nv_bfloat16* vsm = reinterpret_cast<__nv_bfloat16*>(smem + off);
off = align16(off + 64 * D * 2);
__nv_bfloat16* qsm = reinterpret_cast<__nv_bfloat16*>(smem + off);
off = align16(off + TILE * D * 2);
float* accs = reinterpret_cast<float*>(smem + off);
off = align16(off + TILE * D * 4);
float* mls = reinterpret_cast<float*>(smem + off);
off = align16(off + TILE * 2 * 4);
int* act = reinterpret_cast<int*>(smem + off);
off += TILE * 4;
int* ar0 = reinterpret_cast<int*>(smem + off);
off += TILE * 4;
int* ar1 = reinterpret_cast<int*>(smem + off);
off = align16(off + TILE * 4);
constexpr int SCORE_N = (D > 64 ? D : 64);
float* scores = reinterpret_cast<float*>(smem + off);
off = align16(off + 16 * SCORE_N * 4);
__nv_bfloat16* ps = reinterpret_cast<__nv_bfloat16*>(smem + off);
off = align16(off + 16 * 64 * 2);
__nv_bfloat16* qg = reinterpret_cast<__nv_bfloat16*>(smem + off);
off = align16(off + 16 * D * 2);
float* alphas = reinterpret_cast<float*>(smem + off);
const long long head = static_cast<long long>(bh) * S;
const __nv_bfloat16* kbase = k + head * D;
const __nv_bfloat16* vbase = v + head * D;
const int qn4 = tile_q * D / 8;
const uint4* qsrc = reinterpret_cast<const uint4*>(q + (head + q0) * D);
uint4* qdst = reinterpret_cast<uint4*>(qsm);
for (int i = threadIdx.x; i < qn4; i += THREADS) qdst[i] = qsrc[i];
for (int i = threadIdx.x; i < tile_q * D; i += THREADS) accs[i] = 0.f;
for (int i = threadIdx.x; i < tile_q; i += THREADS) {
mls[i * 2] = -INFINITY;
mls[i * 2 + 1] = 0.f;
}
__shared__ int nact;
__syncthreads();
for (int bi = 0; bi < n_blocks; ++bi) {
const int s0 = bi * BLOCK;
if (s0 > tile_max_t) break;
const int slen = min(BLOCK, S - s0);
if (threadIdx.x == 0) nact = 0;
__syncthreads();
for (int ql = threadIdx.x; ql < tile_q; ql += THREADS) {
const int t = q0 + ql;
int r0 = -1, r1 = -1;
if (s0 <= t) {
const short* tp = topk + (head + t) * 8;
const bool sel = tp[0] == bi || tp[1] == bi || tp[2] == bi || tp[3] == bi ||
tp[4] == bi || tp[5] == bi || tp[6] == bi || tp[7] == bi;
if (sel) {
const int end = min(s0 + slen, t + 1);
if (s0 < end) { r0 = 0; r1 = end - s0; }
} else {
int w0 = t + 1 - WINDOW;
if (w0 < 0) w0 = 0;
const int a = max(w0, s0);
const int b = min(t + 1, s0 + slen);
if (a < b) { r0 = a - s0; r1 = b - s0; }
}
}
if (r0 >= 0) {
const int slot = atomicAdd(&nact, 1);
act[slot] = ql;
ar0[slot] = r0;
ar1[slot] = r1;
}
}
const int n4 = slen * (D * 2 / 16);
const uint4* ksrc = reinterpret_cast<const uint4*>(kbase + static_cast<long long>(s0) * D);
const uint4* vsrc = reinterpret_cast<const uint4*>(vbase + static_cast<long long>(s0) * D);
for (int i = threadIdx.x; i < n4; i += THREADS) {
reinterpret_cast<uint4*>(ksm)[i] = ksrc[i];
reinterpret_cast<uint4*>(vsm)[i] = vsrc[i];
}
for (int i = threadIdx.x + slen * D; i < 64 * D; i += THREADS) {
ksm[i] = __nv_bfloat16{};
vsm[i] = __nv_bfloat16{};
}
__syncthreads();
const int nactive = nact;
for (int g = 0; g < nactive; g += 16) {
const int ngrp = min(16, nactive - g);
for (int i = threadIdx.x; i < ngrp * D; i += THREADS) {
const int rr = i / D;
const int dd = i - rr * D;
qg[rr * D + dd] = qsm[act[g + rr] * D + dd];
}
for (int i = threadIdx.x + ngrp * D; i < 16 * D; i += THREADS) qg[i] = __nv_bfloat16{};
__syncthreads();
if (warp == 0) {
for (int n0 = 0; n0 < 64; n0 += 16) {
fragment<accumulator, 16, 16, 16, float> c;
fill_fragment(c, 0.f);
for (int kk = 0; kk < D; kk += 16) {
fragment<matrix_a, 16, 16, 16, __nv_bfloat16, row_major> a;
fragment<matrix_b, 16, 16, 16, __nv_bfloat16, col_major> b;
load_matrix_sync(a, qg + kk, D);
load_matrix_sync(b, ksm + n0 * D + kk, D);
mma_sync(c, a, b, c);
}
store_matrix_sync(scores + n0, c, 64, mem_row_major);
}
__syncwarp();
const int row = lane >> 1;
const int sub = lane & 1;
const bool real = row < ngrp;
const int r0 = real ? ar0[g + row] : -1;
const int r1 = real ? ar1[g + row] : -1;
float tmax = -INFINITY;
#pragma unroll
for (int c = 0; c < 32; ++c) {
const int col = sub * 32 + c;
float s = scores[row * 64 + col] * scale;
if (!real || col < r0 || col >= r1) s = -INFINITY;
scores[row * 64 + col] = s;
tmax = fmaxf(tmax, s);
}
tmax = fmaxf(tmax, __shfl_xor_sync(0xffffffff, tmax, 1));
const bool active = real && tmax > -1.0e30f;
const int ql = real ? act[g + row] : 0;
float m_old = active ? mls[ql * 2] : -INFINITY;
float l_old = active ? mls[ql * 2 + 1] : 0.f;
float m_new = active ? fmaxf(m_old, tmax) : m_old;
float alpha = 1.f;
if (active) alpha = (m_old > -1.0e30f) ? __expf(m_old - m_new) : 0.f;
float l_add = 0.f;
#pragma unroll
for (int c = 0; c < 32; ++c) {
const int col = sub * 32 + c;
float s = scores[row * 64 + col];
float w = (active && s > -1.0e30f) ? __expf(s - m_new) : 0.f;
l_add += w;
ps[row * 64 + col] = __float2bfloat16(w);
}
l_add += __shfl_xor_sync(0xffffffff, l_add, 1);
if (sub == 0) {
alphas[row] = alpha;
if (active) {
mls[ql * 2] = m_new;
mls[ql * 2 + 1] = l_old * alpha + l_add;
}
}
}
__syncthreads();
// scale scattered acc rows, then PV, then add
for (int row = threadIdx.x; row < ngrp; row += THREADS) {
float* arow = accs + act[g + row] * D;
const float al = alphas[row];
for (int d = 0; d < D; ++d) arow[d] *= al;
}
__syncthreads();
if (warp == 0) {
for (int d0 = 0; d0 < D; d0 += 16) {
fragment<accumulator, 16, 16, 16, float> c;
fill_fragment(c, 0.f);
for (int n0 = 0; n0 < 64; n0 += 16) {
fragment<matrix_a, 16, 16, 16, __nv_bfloat16, row_major> a;
fragment<matrix_b, 16, 16, 16, __nv_bfloat16, row_major> b;
load_matrix_sync(a, ps + n0, 64);
load_matrix_sync(b, vsm + n0 * D + d0, D);
mma_sync(c, a, b, c);
}
store_matrix_sync(scores + d0, c, D, mem_row_major);
}
}
__syncthreads();
for (int i = threadIdx.x; i < ngrp * D; i += THREADS) {
const int rr = i / D;
const int dd = i - rr * D;
accs[act[g + rr] * D + dd] += scores[rr * D + dd];
}
__syncthreads();
}
__syncthreads();
}
for (int ql = threadIdx.x; ql < tile_q; ql += THREADS) {
const int t = q0 + ql;
const float l = mls[ql * 2 + 1];
const float inv = (l > 0.f) ? (1.0f / l) : 0.f;
__nv_bfloat16* orow = o + (head + t) * D;
const float* arow = accs + ql * D;
for (int d = 0; d < D; d += 2) {
__nv_bfloat162 r;
r.x = __float2bfloat16(arow[d] * inv);
r.y = __float2bfloat16(arow[d + 1] * inv);
*reinterpret_cast<__nv_bfloat162*>(orow + d) = r;
}
}
}
template <int D, int NW>
__global__ void __launch_bounds__(NW * 32, 1)
tc_attend_kernel(const __nv_bfloat16* __restrict__ q,
const __nv_bfloat16* __restrict__ k,
const __nv_bfloat16* __restrict__ v,
const short* __restrict__ topk,
__nv_bfloat16* __restrict__ o,
int S, int n_blocks, float scale) {
constexpr int TILE = NW * 16;
constexpr int THREADS = NW * 32;
constexpr int BLOCK = 64;
constexpr int WINDOW = 64;
const int lane = threadIdx.x & 31;
const int warp = threadIdx.x >> 5;
const int bh = blockIdx.y;
const int q0 = blockIdx.x * TILE;
if (q0 >= S) return;
const int tile_q = min(TILE, S - q0);
const int tile_max_t = q0 + tile_q - 1;
extern __shared__ __align__(16) char smem[];
const auto lay = tc_layout<D, NW>();
__nv_bfloat16* ksm = reinterpret_cast<__nv_bfloat16*>(smem + lay.ksm);
__nv_bfloat16* vsm = reinterpret_cast<__nv_bfloat16*>(smem + lay.vsm);
__nv_bfloat16* qsm = reinterpret_cast<__nv_bfloat16*>(smem + lay.qsm);
float* accs = reinterpret_cast<float*>(smem + lay.acc);
float* mls = reinterpret_cast<float*>(smem + lay.ml);
int scratch = 16 * D * 4;
if (16 * 64 * 4 > scratch) scratch = 16 * 64 * 4;
scratch = (scratch + 15) & ~15;
float* sc_base = reinterpret_cast<float*>(smem + lay.sc);
__nv_bfloat16* p_base = reinterpret_cast<__nv_bfloat16*>(smem + lay.p);
const long long head = static_cast<long long>(bh) * S;
const __nv_bfloat16* kbase = k + head * D;
const __nv_bfloat16* vbase = v + head * D;
// Q for the tile, zero softmax state
const int qn = tile_q * D;
const uint4* qsrc = reinterpret_cast<const uint4*>(q + (head + q0) * D);
uint4* qdst = reinterpret_cast<uint4*>(qsm);
const int qn4 = qn * (int)sizeof(__nv_bfloat16) / 16;
for (int i = threadIdx.x; i < qn4; i += THREADS) qdst[i] = qsrc[i];
for (int i = threadIdx.x; i < tile_q * D; i += THREADS) accs[i] = 0.f;
for (int i = threadIdx.x; i < tile_q; i += THREADS) {
mls[i * 2] = -INFINITY;
mls[i * 2 + 1] = 0.f;
}
__syncthreads();
for (int bi = 0; bi < n_blocks; ++bi) {
const int s0 = bi * BLOCK;
if (s0 > tile_max_t) break;
const int slen = min(BLOCK, S - s0);
const int n4 = slen * (D * 2 / 16);
const uint4* ksrc = reinterpret_cast<const uint4*>(kbase + static_cast<long long>(s0) * D);
const uint4* vsrc = reinterpret_cast<const uint4*>(vbase + static_cast<long long>(s0) * D);
uint4* kdst = reinterpret_cast<uint4*>(ksm);
uint4* vdst = reinterpret_cast<uint4*>(vsm);
for (int i = threadIdx.x; i < n4; i += THREADS) {
kdst[i] = ksrc[i];
vdst[i] = vsrc[i];
}
// Tail rows are read by WMMA even when the block is short. Zero them so
// a masked P=0 cannot pick up NaN from uninitialized smem.
for (int i = threadIdx.x + slen * D; i < 64 * D; i += THREADS) {
ksm[i] = __nv_bfloat16{};
vsm[i] = __nv_bfloat16{};
}
__syncthreads();
// Each warp owns 16 queries. Inactive warps (past tile) skip uniformly.
const int row = lane >> 1; // 0..15
const int sub = lane & 1; // column half
const int ql = warp * 16 + row;
const int t = q0 + ql;
const bool in_tile = ql < tile_q && t < S;
int r0 = -1, r1 = -1;
if (in_tile && s0 <= t) {
const short* tp = topk + (head + t) * 8;
const bool sel = tp[0] == bi || tp[1] == bi || tp[2] == bi || tp[3] == bi ||
tp[4] == bi || tp[5] == bi || tp[6] == bi || tp[7] == bi;
if (sel) {
const int end = min(s0 + slen, t + 1);
if (s0 < end) {
r0 = 0;
r1 = end - s0;
}
} else {
int w0 = t + 1 - WINDOW;
if (w0 < 0) w0 = 0;
const int a = max(w0, s0);
const int b = min(t + 1, s0 + slen);
if (a < b) {
r0 = a - s0;
r1 = b - s0;
}
}
}
const bool any = __any_sync(0xffffffff, r0 >= 0);
if (any) {
float* scores = reinterpret_cast<float*>(
reinterpret_cast<char*>(sc_base) + warp * scratch);
__nv_bfloat16* ps = p_base + warp * 16 * 64;
__nv_bfloat16* qwarp = qsm + warp * 16 * D;
// Q[16,D] @ K[64,D]^T -> scores[16,64]
for (int n0 = 0; n0 < 64; n0 += 16) {
fragment<accumulator, 16, 16, 16, float> c;
fill_fragment(c, 0.f);
for (int kk = 0; kk < D; kk += 16) {
fragment<matrix_a, 16, 16, 16, __nv_bfloat16, row_major> a;
fragment<matrix_b, 16, 16, 16, __nv_bfloat16, col_major> b;
load_matrix_sync(a, qwarp + kk, D);
load_matrix_sync(b, ksm + n0 * D + kk, D);
mma_sync(c, a, b, c);
}
store_matrix_sync(scores + n0, c, 64, mem_row_major);
}
__syncwarp();
// Mask, online-softmax stats, bf16 weights. Two lanes share a row.
float tmax = -INFINITY;
float l_add = 0.f;
#pragma unroll
for (int c = 0; c < 32; ++c) {
const int col = sub * 32 + c;
float s = scores[row * 64 + col] * scale;
if (r0 < 0 || col < r0 || col >= r1) s = -INFINITY;
scores[row * 64 + col] = s;
tmax = fmaxf(tmax, s);
}
tmax = fmaxf(tmax, __shfl_xor_sync(0xffffffff, tmax, 1));
const bool active = tmax > -1.0e30f;
const int qli = warp * 16 + row;
float m_old = in_tile ? mls[qli * 2] : -INFINITY;
float l_old = in_tile ? mls[qli * 2 + 1] : 0.f;
float m_new = active ? fmaxf(m_old, tmax) : m_old;
float alpha = 1.f;
if (active) alpha = (m_old > -1.0e30f) ? __expf(m_old - m_new) : 0.f;
#pragma unroll
for (int c = 0; c < 32; ++c) {
const int col = sub * 32 + c;
float s = scores[row * 64 + col];
float w = 0.f;
if (active && s > -1.0e30f) w = __expf(s - m_new);
l_add += w;
ps[row * 64 + col] = __float2bfloat16(w);
}
l_add += __shfl_xor_sync(0xffffffff, l_add, 1);
if (sub == 0 && in_tile) {
mls[qli * 2] = m_new;
mls[qli * 2 + 1] = l_old * alpha + l_add;
}
__syncwarp();
// Weights already live in ps. Park per-row alpha at scores[row, 0].
if (sub == 0) scores[row * 64] = alpha;
__syncwarp();
float* acc_w = accs + warp * 16 * D;
for (int i = lane; i < 16 * D; i += 32)
acc_w[i] *= scores[(i / D) * 64];
__syncwarp();
// P @ V -> pv tmp aliased over this warp's score buffer
for (int d0 = 0; d0 < D; d0 += 16) {
fragment<accumulator, 16, 16, 16, float> c;
fill_fragment(c, 0.f);
for (int n0 = 0; n0 < 64; n0 += 16) {
fragment<matrix_a, 16, 16, 16, __nv_bfloat16, row_major> a;
fragment<matrix_b, 16, 16, 16, __nv_bfloat16, row_major> b;
load_matrix_sync(a, ps + n0, 64);
load_matrix_sync(b, vsm + n0 * D + d0, D);
mma_sync(c, a, b, c);
}
store_matrix_sync(scores + d0, c, D, mem_row_major);
}
__syncwarp();
for (int i = lane; i < 16 * D; i += 32) acc_w[i] += scores[i];
__syncwarp();
}
__syncthreads();
}
for (int ql = threadIdx.x; ql < tile_q; ql += THREADS) {
const int t = q0 + ql;
float l = mls[ql * 2 + 1];
float inv = (l > 0.f) ? (1.0f / l) : 0.f;
__nv_bfloat16* orow = o + (head + t) * D;
const float* arow = accs + ql * D;
for (int d = 0; d < D; d += 2) {
__nv_bfloat162 r;
r.x = __float2bfloat16(arow[d] * inv);
r.y = __float2bfloat16(arow[d + 1] * inv);
*reinterpret_cast<__nv_bfloat162*>(orow + d) = r;
}
}
}
template <int D, int NW>
int tc_smem() {
return tc_layout<D, NW>().bytes;
}
// One warp per query. High occupancy, state in registers. Best when the selected
// set is a small fraction of S and L2 already holds K and V.
template <int D>
__global__ void __launch_bounds__(128, 4)
warp_attend_kernel(const __nv_bfloat16* __restrict__ q,
const __nv_bfloat16* __restrict__ k,
const __nv_bfloat16* __restrict__ v,
const float* __restrict__ block_sum,
__nv_bfloat16* __restrict__ o,
int S, int n_blocks, float scale) {
constexpr int V = D / 32;
constexpr int BLOCK = 64;
constexpr int WINDOW = 64;
const int lane = threadIdx.x & 31;
const int warp = threadIdx.x >> 5;
const int warps = blockDim.x >> 5;
const int bh = blockIdx.y;
const int t = blockIdx.x * warps + warp;
if (t >= S) return;
const long long head = static_cast<long long>(bh) * S;
const __nv_bfloat16* qrow = q + (head + t) * D;
const __nv_bfloat16* kbase = k + head * D;
const __nv_bfloat16* vbase = v + head * D;
const float* bs = block_sum + static_cast<long long>(bh) * n_blocks * D;
float qreg[V];
#pragma unroll
for (int i = 0; i < V; i += 2) {
__nv_bfloat162 qq = *reinterpret_cast<const __nv_bfloat162*>(qrow + lane * V + i);
float2 f = __bfloat1622float2(qq);
qreg[i] = f.x;
qreg[i + 1] = f.y;
}
float ts0 = -INFINITY, ts1 = -INFINITY, ts2 = -INFINITY, ts3 = -INFINITY;
float ts4 = -INFINITY, ts5 = -INFINITY, ts6 = -INFINITY, ts7 = -INFINITY;
int tb0 = 32767, tb1 = 32767, tb2 = 32767, tb3 = 32767;
int tb4 = 32767, tb5 = 32767, tb6 = 32767, tb7 = 32767;
const int cb = t >> 6;
const int block_start = cb << 6;
const float inv_block = 1.0f / 64.0f;
for (int bi = 0; bi < cb; ++bi) {
float dot = dot_f32<D>(bs + static_cast<long long>(bi) * D, qreg, lane);
CONSIDER(dot * scale * inv_block, bi);
}
float psum = 0.f;
const int count = t - block_start + 1;
for (int j = block_start; j <= t; ++j)
psum += dot_bf16<D>(kbase + static_cast<long long>(j) * D, qreg, lane);
CONSIDER(psum * scale / static_cast<float>(count), cb);
uint32_t sel0 = 0, sel1 = 0, sel2 = 0, sel3 = 0;
auto setbit = [&](int b) {
if (b < 0 || b >= 128) return;
if (b < 32) sel0 |= 1u << b;
else if (b < 64) sel1 |= 1u << (b - 32);
else if (b < 96) sel2 |= 1u << (b - 64);
else sel3 |= 1u << (b - 96);
};
setbit(tb0); setbit(tb1); setbit(tb2); setbit(tb3);
setbit(tb4); setbit(tb5); setbit(tb6); setbit(tb7);
auto selected = [&](int b) {
if (b < 32) return (sel0 >> b) & 1u;
if (b < 64) return (sel1 >> (b - 32)) & 1u;
if (b < 96) return (sel2 >> (b - 64)) & 1u;
return (sel3 >> (b - 96)) & 1u;
};
float acc[V];
#pragma unroll
for (int i = 0; i < V; ++i) acc[i] = 0.f;
float m = -INFINITY, l = 0.f;
auto attend = [&](int s0, int s1) {
for (int j = s0; j < s1; ++j) {
float score = dot_bf16<D>(kbase + static_cast<long long>(j) * D, qreg, lane) * scale;
float m_new = fmaxf(m, score);
float alpha = (m > -1.0e30f) ? __expf(m - m_new) : 0.f;
float p = __expf(score - m_new);
l = fmaf(l, alpha, p);
const __nv_bfloat16* vp = vbase + static_cast<long long>(j) * D + lane * V;
#pragma unroll
for (int i = 0; i < V; i += 2) {
__nv_bfloat162 vv = *reinterpret_cast<const __nv_bfloat162*>(vp + i);
float2 f = __bfloat1622float2(vv);
acc[i] = fmaf(p, f.x, acc[i] * alpha);
acc[i + 1] = fmaf(p, f.y, acc[i + 1] * alpha);
}
m = m_new;
}
};
for (int bi = 0; bi < cb; ++bi) {
if (!selected(bi)) continue;
attend(bi * BLOCK, bi * BLOCK + BLOCK);
}
if (selected(cb)) attend(block_start, t + 1);
int w0 = t + 1 - WINDOW;
if (w0 < 0) w0 = 0;
if (cb > 0 && !selected(cb - 1)) {
int prev = (cb - 1) * BLOCK;
int a = w0 > prev ? w0 : prev;
if (a < block_start) attend(a, block_start);
}
if (!selected(cb)) {
int a = w0 > block_start ? w0 : block_start;
if (a <= t) attend(a, t + 1);
}
float inv = (l > 0.f) ? (1.0f / l) : 0.f;
__nv_bfloat16* orow = o + (head + t) * D + lane * V;
#pragma unroll
for (int i = 0; i < V; i += 2) {
__nv_bfloat162 r;
r.x = __float2bfloat16(acc[i] * inv);
r.y = __float2bfloat16(acc[i + 1] * inv);
*reinterpret_cast<__nv_bfloat162*>(orow + i) = r;
}
}
void nsa_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v,
torch::Tensor block_sum, torch::Tensor topk, torch::Tensor out) {
TORCH_CHECK(q.is_cuda() && q.is_contiguous() && k.is_contiguous() && v.is_contiguous(), "nsa");
TORCH_CHECK(q.scalar_type() == torch::kBFloat16, "nsa bf16");
const int B = q.size(0);
const int H = q.size(1);
const int S = q.size(2);
const int D = q.size(3);
TORCH_CHECK(D == 64 || D == 128, "nsa D");
const int n_blocks = (S + 63) / 64;
TORCH_CHECK(n_blocks <= 128, "nsa n_blocks");
const float scale = 1.0f / sqrtf(static_cast<float>(D));
auto stream = at::cuda::getCurrentCUDAStream();
const auto* q_p = reinterpret_cast<const __nv_bfloat16*>(q.data_ptr<at::BFloat16>());
const auto* k_p = reinterpret_cast<const __nv_bfloat16*>(k.data_ptr<at::BFloat16>());
const auto* v_p = reinterpret_cast<const __nv_bfloat16*>(v.data_ptr<at::BFloat16>());
auto* o_p = reinterpret_cast<__nv_bfloat16*>(out.data_ptr<at::BFloat16>());
float* bs = block_sum.data_ptr<float>();
short* tk = topk.data_ptr<int16_t>();
dim3 bsg(n_blocks, B * H);
dim3 sg((S + 3) / 4, B * H);
dim3 wg((S + 3) / 4, B * H);
// Short sequences are dense enough that WMMA wins. Long sequences are
// L2-resident gathers; a warp per query beats padded tensor-core tiles.
if (D == 64 && S <= 3072) {
block_sum_kernel<64><<<bsg, 32, 0, stream>>>(k_p, bs, S, n_blocks);
score_kernel<64><<<sg, 128, 0, stream>>>(q_p, k_p, bs, tk, S, n_blocks, scale);
constexpr int NW = 4;
int smem = tc_smem<64, NW>();
auto kern = tc_attend_kernel<64, NW>;
cudaFuncSetAttribute(kern, cudaFuncAttributeMaxDynamicSharedMemorySize, smem);
dim3 ag((S + NW * 16 - 1) / (NW * 16), B * H);
kern<<<ag, NW * 32, smem, stream>>>(q_p, k_p, v_p, tk, o_p, S, n_blocks, scale);
} else if (D == 64) {
block_sum_kernel<64><<<bsg, 32, 0, stream>>>(k_p, bs, S, n_blocks);
warp_attend_kernel<64><<<wg, 128, 0, stream>>>(q_p, k_p, v_p, bs, o_p, S, n_blocks, scale);
} else {
block_sum_kernel<128><<<bsg, 32, 0, stream>>>(k_p, bs, S, n_blocks);
warp_attend_kernel<128><<<wg, 128, 0, stream>>>(q_p, k_p, v_p, bs, o_p, S, n_blocks, scale);
}
}
"""
def _build():
os.environ.setdefault("TORCH_CUDA_ARCH_LIST", "12.0")
return load_inline(
name="nsa_sm120_v8",
cpp_sources=[_CPP],
cuda_sources=[_CUDA],
functions=["nsa_forward"],
extra_cuda_cflags=[
"-O3",
"--use_fast_math",
"-std=c++20",
"--expt-relaxed-constexpr",
],
with_cuda=True,
verbose=False,
)
_EXT = None
def _ext():
global _EXT
if _EXT is None:
_EXT = _build()
return _EXT
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))
self._bsum = None
self._topk = None
self._out = None
def forward(self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor) -> torch.Tensor:
if not q.is_contiguous():
q = q.contiguous()
if not k.is_contiguous():
k = k.contiguous()
if not v.is_contiguous():
v = v.contiguous()
B, H, S, D = q.shape
n_blocks = (S + _BLOCK - 1) // _BLOCK
dev = q.device
if self._bsum is None or self._bsum.shape != (B, H, n_blocks, D) or self._bsum.device != dev:
self._bsum = torch.empty(B, H, n_blocks, D, dtype=torch.float32, device=dev)
if self._topk is None or self._topk.shape != (B, H, S, 8) or self._topk.device != dev:
self._topk = torch.empty(B, H, S, 8, dtype=torch.int16, device=dev)
if self._out is None or self._out.shape != q.shape or self._out.device != dev:
self._out = torch.empty_like(q)
_ext().nsa_forward(q, k, v, self._bsum, self._topk, self._out)
return self._out
20260917_005835_grok_grok-4.7_02_deepseek_nsa