kernelbench.com

KernelBench cuda · RTX PRO 6000

DeepSeek NSA Claude Opus 4.8

17.8%geomean peak fraction across shapes

manually audited: clean

Genuine raw-CUDA DeepSeek-style NSA sparse attention. A first CUDA kernel mean-pools live K by 64-token block; a second query-tiled kernel computes the exact causal diagonal prefix score and full-block importance, selects the per-query top eight with the reference tie-break, unions the sliding window, builds a block-to-query CSR schedule, streams selected live K/V blocks, and performs online-softmax attention into a fresh output. The mean-pooling identity reduces block scoring without changing semantics. There is no output cache, fixed selection table, dense-attention substitute on a scored shape, forbidden op, Triton/DSL path, cross-run artifact reuse, grader mutation, or numeric-stress bypass. Language sidecars say cuda_raw with real CUDA evidence and no cheat flags. For this ms-headline problem the six official timings are the ground truth; the archived 0.1784 scalar is internally reproduced but its dense-equivalent FLOP ceiling is not a physically meaningful performance headline.

harnessclaudeagent session4h 56mtotal wall5h 17mcheck14mbenchmark7moutput tokens315cost$33.61gpu-lock wait3h 12mgpu-lock held21mregimecompute

Per-shape vs governing ceilingeach shape graded against whichever binds — bf16 compute or HBM bandwidth

1×16×2048×640.316 ms10.9%54 TFLOPS · 11% of 500 TF bf16 peak · also 0.05 TB/s (3% of HBM)
1×16×4127×640.631 ms22.1%110 TFLOPS · 22% of 500 TF bf16 peak · also 0.05 TB/s (3% of HBM)
1×8×8192×640.747 ms36.8%184 TFLOPS · 37% of 500 TF bf16 peak · also 0.04 TB/s (2% of HBM)
1×8×8191×1281.635 ms33.6%168 TFLOPS · 34% of 500 TF bf16 peak · also 0.04 TB/s (2% of HBM)
4×8×1024×640.271 ms6.3%32 TFLOPS · 6% of 500 TF bf16 peak · also 0.06 TB/s (3% of HBM)
2×8×3000×640.431 ms17.1%86 TFLOPS · 17% of 500 TF bf16 peak · also 0.06 TB/s (3% of HBM)

compute-bound memory-bound · bar + right column = official fraction of the ceiling (the geomean input)

geomean(10.9% · 22.1% · 36.8% · 33.6% · 6.3% · 17.1%) = 17.8%

Kernel source (redacted)
"""DeepSeek NSA-inspired sparse attention — fused CUDA kernel for SM120 (RTX PRO 6000).

Implements the bench semantics of reference.nsa_attend: per query, score key blocks
by the mean of q.k over the block's causal keys, take the top-8, union with the
64-token sliding window, and softmax-attend over that key set only.

The kernel fuses block scoring, top-n selection, gather and online sparse attention
into one launch per (batch, head, query block).
"""
from __future__ import annotations

import math
import os

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 <cuda_runtime.h>
#include <cuda_bf16.h>
#include <torch/extension.h>
#include <c10/cuda/CUDAStream.h>

#define BLK 64          // block_size
#define TOPN 8          // top_n_blocks
#define WIN 64          // sliding_window
// 16 warps/CTA keeps registers at 54 (vs 72 at 8 warps) so two 512-thread CTAs
// fit per SM: 1024 threads = 32 warps = 8 per scheduler. This kernel is latency-
// bound (a barrier + a dependent global->smem load per key block), so warps in
// flight matter far more than per-warp efficiency.
#define NWARPS 16
#define NTHREADS (NWARPS * 32)
#define MAX_ENTRIES_PER_Q (TOPN + 2)

// ---------------------------------------------------------------- helpers

__device__ __forceinline__ unsigned int order_f32(float x) {
    // Monotonic float->uint map so unsigned max == float max (finite values).
    unsigned int u = __float_as_uint(x);
    return (u & 0x80000000u) ? ~u : (u | 0x80000000u);
}

__device__ __forceinline__ float bf2f(const __nv_bfloat16 h) { return __bfloat162float(h); }

__device__ __forceinline__ float warp_max(float v) {
#pragma unroll
    for (int off = 16; off > 0; off >>= 1) v = fmaxf(v, __shfl_xor_sync(0xffffffffu, v, off));
    return v;
}

__device__ __forceinline__ float warp_sum(float v) {
#pragma unroll
    for (int off = 16; off > 0; off >>= 1) v += __shfl_xor_sync(0xffffffffu, v, off);
    return v;
}

__device__ __forceinline__ unsigned long long warp_maxu64(unsigned long long v) {
#pragma unroll
    for (int off = 16; off > 0; off >>= 1) {
        unsigned long long o = __shfl_xor_sync(0xffffffffu, v, off);
        v = (o > v) ? o : v;
    }
    return v;
}

// Insert into a descending-sorted top-TOPN list (single thread).
__device__ __forceinline__ void top_insert(unsigned long long* arr, unsigned long long v) {
    if (v <= arr[TOPN - 1]) return;
    int p = TOPN - 1;
    while (p > 0 && arr[p - 1] < v) {
        arr[p] = arr[p - 1];
        --p;
    }
    arr[p] = v;
}

// ---------------------------------------------------------------- kbar kernel
// Mean-pooled key per block ("compress" branch of NSA). Because the block
// importance is a mean of dot products and the mean is linear,
//     mean_j (q . k_j) == q . (mean_j k_j)
// so a fully-causal block needs only this one pooled key, turning the O(S^2)
// scoring pass into O(S * NB).
//
// Layout is (B, H, D, NB) -- transposed so a warp reading consecutive block ids
// at a fixed d is fully coalesced during scoring.
template <int D>
__global__ void kbar_kernel(const __nv_bfloat16* __restrict__ K, float* __restrict__ KbarT,
                            int S, int NB) {
    const int bi = blockIdx.x;
    const int h = blockIdx.y;
    const int b = blockIdx.z;
    const int H = gridDim.y;
    const long bh = (long)b * H + h;
    const __nv_bfloat16* Kb = K + bh * (long)S * D;

    const int j0 = bi * BLK;
    const int j1 = min(j0 + BLK, S);
    const int d = threadIdx.x;  // blockDim.x == D

    float sum = 0.f;
#pragma unroll 8
    for (int j = j0; j < j1; ++j) sum += bf2f(Kb[(long)j * D + d]);
    const int cnt = j1 - j0;
    KbarT[(bh * D + d) * (long)NB + bi] = (cnt > 0) ? sum / (float)cnt : 0.f;
}

// ---------------------------------------------------------------- smem plan

struct SmemPlan {
    int q, k, v, acc, top, p, m, l, diag, cnt, off, list, total;
};

__host__ __device__ __forceinline__ int smem_push(int* o, int bytes) {
    const int at = *o;
    *o += (bytes + 15) & ~15;
    return at;
}

__host__ __device__ __forceinline__ SmemPlan plan_smem(int D, int NB) {
    const int KSTRIDE = D + 8;
    SmemPlan s;
    int o = 0;
    s.q = smem_push(&o, BLK * D * 2);
    s.k = smem_push(&o, BLK * KSTRIDE * 2);
    s.v = smem_push(&o, BLK * KSTRIDE * 2);
    s.acc = smem_push(&o, BLK * D * 4);
    s.top = smem_push(&o, BLK * TOPN * 8);
    // sTop is dead once the CSR lists exist, so the softmax-probability scratch
    // reuses its space. Those 2KB are what buys a 2nd resident CTA per SM
    // (2 x 50176 <= 102400), doubling warps/SM from 8 to 16.
    s.p = s.top;
    s.m = smem_push(&o, BLK * 4);
    s.l = smem_push(&o, BLK * 4);
    s.diag = smem_push(&o, BLK * 4);
    s.cnt = smem_push(&o, NB * 4);
    s.off = smem_push(&o, NB * 4);
    s.list = smem_push(&o, BLK * MAX_ENTRIES_PER_Q * 2);
    s.total = o;
    return s;
}

// ---------------------------------------------------------------- main kernel
//
// One CTA per (batch, head, query block of 64).
//
// Each query picks its own ~10 key blocks, so the union over the 64 queries in a
// CTA covers nearly every causal block: gathering per query from global would
// re-read K/V ~10x per query. We invert the loop instead -- stream each causal key
// block into shared memory ONCE and process only the (query, block) pairs that
// actually selected it, kept as a compact CSR list. Compute stays sparse (~10
// blocks per query) while traffic stays at ~one pass of K/V per query block.

template <int D, int MAXB>
__global__ __launch_bounds__(NTHREADS, 2) void nsa_kernel(
    const __nv_bfloat16* __restrict__ Q, const __nv_bfloat16* __restrict__ K,
    const __nv_bfloat16* __restrict__ V, const float* __restrict__ KbarT,
    __nv_bfloat16* __restrict__ O, int S, int NB, int n_qtiles) {
    constexpr int KSTRIDE = D + 8;     // pad rows so smem reads are conflict-free
    constexpr int NQW = BLK / NWARPS;  // queries per warp during scoring
    constexpr int DV = D / 8;          // int4 chunks per row

    extern __shared__ char smem_raw[];
    const SmemPlan sp = plan_smem(D, NB);
    __nv_bfloat16* sQ = (__nv_bfloat16*)(smem_raw + sp.q);
    __nv_bfloat16* sK = (__nv_bfloat16*)(smem_raw + sp.k);
    __nv_bfloat16* sV = (__nv_bfloat16*)(smem_raw + sp.v);
    float* sAcc = (float*)(smem_raw + sp.acc);
    unsigned long long* sTop = (unsigned long long*)(smem_raw + sp.top);
    float* sP = (float*)(smem_raw + sp.p);
    float* sM = (float*)(smem_raw + sp.m);
    float* sL = (float*)(smem_raw + sp.l);
    float* sDiag = (float*)(smem_raw + sp.diag);
    int* sCount = (int*)(smem_raw + sp.cnt);
    int* sOff = (int*)(smem_raw + sp.off);
    unsigned short* sList = (unsigned short*)(smem_raw + sp.list);

    const int tid = threadIdx.x;
    const int lane = tid & 31;
    const int warp = tid >> 5;

    // Descending work order: causal cost grows with the tile index, so the most
    // expensive tiles are scheduled first.
    const int bq = n_qtiles - 1 - blockIdx.x;
    const int h = blockIdx.y;
    const int b = blockIdx.z;
    const int H = gridDim.y;
    const long bh = (long)b * H + h;
    const int qbase = bq * BLK;

    const __nv_bfloat16* Qb = Q + bh * (long)S * D;
    const __nv_bfloat16* Kb = K + bh * (long)S * D;
    const __nv_bfloat16* Vb = V + bh * (long)S * D;
    __nv_bfloat16* Ob = O + bh * (long)S * D;
    const float* KbarTb = KbarT + bh * (long)D * NB;

    const float scale = rsqrtf((float)D);

    // ---- phase 0: load Q tile ----
    for (int idx = tid; idx < BLK * DV; idx += NTHREADS) {
        const int row = idx / DV, c = idx - row * DV;
        const int t = qbase + row;
        int4 val = make_int4(0, 0, 0, 0);
        if (t < S) val = reinterpret_cast<const int4*>(Qb + (long)t * D)[c];
        reinterpret_cast<int4*>(&sQ[row * D])[c] = val;
    }

    // ---- phase 1: diagonal block -> per-query importance of block bq ----
    // Only keys [64*bq, t] are causal there, so the importance is a prefix mean:
    //   imp = q_t . (sum_{j<=t} k_j) / ((t - 64*bq + 1) * sqrt(D))
    for (int idx = tid; idx < BLK * DV; idx += NTHREADS) {
        const int row = idx / DV, c = idx - row * DV;
        const int j = qbase + row;
        int4 val = make_int4(0, 0, 0, 0);
        if (j < S) val = reinterpret_cast<const int4*>(Kb + (long)j * D)[c];
        reinterpret_cast<int4*>(&sK[row * KSTRIDE])[c] = val;
    }
    __syncthreads();

    if (tid < D) {  // running in-block cumulative sum of K (sAcc is scratch here)
        float run = 0.f;
        for (int i = 0; i < BLK; ++i) {
            const int t = qbase + i;
            run += (t < S) ? bf2f(sK[i * KSTRIDE + tid]) : 0.f;
            sAcc[i * D + tid] = run;
        }
    }
    __syncthreads();

    for (int ii = 0; ii < NQW; ++ii) {
        const int i = warp * NQW + ii;
        float part = 0.f;
        for (int d = lane; d < D; d += 32) part += bf2f(sQ[i * D + d]) * sAcc[i * D + d];
        part = warp_sum(part);
        if (lane == 0) sDiag[i] = part * scale / (float)(i + 1);
    }

    for (int i = tid; i < BLK * TOPN; i += NTHREADS) sTop[i] = 0ULL;
    for (int i = tid; i < NB; i += NTHREADS) sCount[i] = 0;
    __syncthreads();

    // ---- phase 2: block scoring + per-query top-8 ----
    // Keys pack (importance, block id) into one u64 so an unsigned max reproduces
    // the reference's sort order exactly, including its tie-break (larger block id
    // wins ties, matching `block_imp.sort(reverse=True)` on (imp, bi) tuples).
    // Pooled block keys are staged once per CTA into the sK/sV/sAcc scratch
    // (all dead here): every warp scores against the same table, so reading it
    // straight from global costs NWARPS redundant passes over it.
    constexpr int CHUNK = 32 * MAXB;
    float* sKbar = (float*)(smem_raw + sp.k);

    const int nb_cand = bq + 1;  // causal blocks are 0..bq
    for (int cb = 0; cb < nb_cand; cb += CHUNK) {
        const int nchunk = min(CHUNK, nb_cand - cb);
        __syncthreads();
        for (int i = tid; i < D * nchunk; i += NTHREADS) {
            const int d = i / nchunk, bi = i - d * nchunk;
            sKbar[d * CHUNK + bi] = KbarTb[(long)d * NB + cb + bi];
        }
        __syncthreads();

        float dotv[NQW][MAXB];
#pragma unroll
        for (int ii = 0; ii < NQW; ++ii)
#pragma unroll
            for (int c = 0; c < MAXB; ++c) dotv[ii][c] = 0.f;

        for (int d = 0; d < D; ++d) {
            float kb[MAXB];
#pragma unroll
            for (int c = 0; c < MAXB; ++c) {
                const int j = c * 32 + lane;
                kb[c] = (j < nchunk) ? sKbar[d * CHUNK + j] : 0.f;
            }
#pragma unroll
            for (int ii = 0; ii < NQW; ++ii) {
                const float qv = bf2f(sQ[(warp * NQW + ii) * D + d]);
#pragma unroll
                for (int c = 0; c < MAXB; ++c) dotv[ii][c] = fmaf(qv, kb[c], dotv[ii][c]);
            }
        }

#pragma unroll
        for (int ii = 0; ii < NQW; ++ii) {
            const int i = warp * NQW + ii;
            unsigned long long keys[MAXB];
#pragma unroll
            for (int c = 0; c < MAXB; ++c) {
                const int bi = cb + c * 32 + lane;
                if (bi < nb_cand) {
                    const float imp = (bi == bq) ? sDiag[i] : dotv[ii][c] * scale;
                    keys[c] = ((unsigned long long)order_f32(imp) << 32) | (unsigned)bi;
                } else {
                    keys[c] = 0ULL;
                }
            }
            // Successive extractions are already descending, so on the first (and,
            // for every benchmarked length, only) pass they can be stored straight
            // down: no sorted insert, no shared read-back, no __syncwarp per round.
            // Later passes must merge into the running list.
            const bool first_pass = (cb == 0);
            for (int r = 0; r < TOPN; ++r) {
                unsigned long long best = 0ULL;
#pragma unroll
                for (int c = 0; c < MAXB; ++c) best = (keys[c] > best) ? keys[c] : best;
                best = warp_maxu64(best);
                if (best == 0ULL) break;  // warp-uniform
                if (first_pass) {
                    if (lane == 0) sTop[i * TOPN + r] = best;
                } else {
                    if (best <= sTop[i * TOPN + TOPN - 1]) break;  // warp-uniform
                    if (lane == 0) top_insert(&sTop[i * TOPN], best);
                    __syncwarp();
                }
#pragma unroll
                for (int c = 0; c < MAXB; ++c)
                    if (keys[c] == best) keys[c] = 0ULL;
            }
        }
    }
    __syncthreads();


    // ---- phase 3: CSR build (key block -> queries that selected it) ----
    // Entry = (query index << 1) | in_top8, where in_top8 distinguishes "whole
    // block selected" from "sliding-window tail only".
    const bool q_active = (tid < BLK) && (qbase + tid < S);
    int myBlk[MAX_ENTRIES_PER_Q];
    int myTop[MAX_ENTRIES_PER_Q];
    int myN = 0;
    if (q_active) {
        const int t = qbase + tid;
        bool top_has_prev = false;
        for (int r = 0; r < TOPN; ++r) {
            const unsigned long long v = sTop[tid * TOPN + r];
            if (v == 0ULL) break;
            const int bi = (int)(v & 0xffffffffULL);
            if (bi == bq) continue;  // added unconditionally below
            if (bi == bq - 1) {      // may also be a window block; add once
                top_has_prev = true;
                continue;
            }
            myBlk[myN] = bi;
            myTop[myN] = 1;
            myN++;
        }
        // Diagonal block is always present via the window, and its causal range
        // [64*bq, t] is identical whether or not it also made the top-8.
        myBlk[myN] = bq;
        myTop[myN] = 1;
        myN++;
        if (bq >= 1) {
            const bool win_overlap = (t - WIN + 1) < bq * BLK;  // i.e. t % 64 < 63
            if (top_has_prev || win_overlap) {
                myBlk[myN] = bq - 1;
                myTop[myN] = top_has_prev ? 1 : 0;
                myN++;
            }
        }
        for (int e = 0; e < myN; ++e) atomicAdd(&sCount[myBlk[e]], 1);
    }
    __syncthreads();


    if (warp == 0) {  // exclusive prefix sum over sCount -> sOff
        int running = 0;
        for (int base = 0; base < NB; base += 32) {
            const int i = base + lane;
            const int c = (i < NB) ? sCount[i] : 0;
            int scan = c;
#pragma unroll
            for (int off = 1; off < 32; off <<= 1) {
                const int n = __shfl_up_sync(0xffffffffu, scan, off);
                if (lane >= off) scan += n;
            }
            const int total = __shfl_sync(0xffffffffu, scan, 31);
            scan -= c;  // inclusive -> exclusive
            if (i < NB) {
                sOff[i] = running + scan;
                sCount[i] = 0;  // reused as a fill cursor, then as the final count
            }
            running += total;
        }
    }
    __syncthreads();

    if (q_active) {
        for (int e = 0; e < myN; ++e) {
            const int bi = myBlk[e];
            const int slot = atomicAdd(&sCount[bi], 1);
            sList[sOff[bi] + slot] = (unsigned short)((tid << 1) | myTop[e]);
        }
    }

    // ---- phase 4: init online-softmax accumulators ----
    for (int i = tid; i < BLK * D; i += NTHREADS) sAcc[i] = 0.f;
    for (int i = tid; i < BLK; i += NTHREADS) {
        sM[i] = -INFINITY;
        sL[i] = 0.f;
    }
    __syncthreads();

    // ---- phase 5: sparse attention over the selected blocks ----
    const int jg = lane >> 3;  // key group (4 keys in flight)
    const int dc = lane & 7;   // dim chunk (8 dims each)

    for (int bk = 0; bk <= bq; ++bk) {
        const int na = sCount[bk];  // CTA-uniform
        if (na == 0) continue;
        const int base = sOff[bk];
        const int kbase = bk * BLK;

        __syncthreads();  // previous iteration has finished reading sK/sV
        for (int idx = tid; idx < BLK * DV; idx += NTHREADS) {
            const int row = idx / DV, c = idx - row * DV;
            const int j = kbase + row;
            int4 kval = make_int4(0, 0, 0, 0), vval = make_int4(0, 0, 0, 0);
            if (j < S) {
                kval = reinterpret_cast<const int4*>(Kb + (long)j * D)[c];
                vval = reinterpret_cast<const int4*>(Vb + (long)j * D)[c];
            }
            reinterpret_cast<int4*>(&sK[row * KSTRIDE])[c] = kval;
            reinterpret_cast<int4*>(&sV[row * KSTRIDE])[c] = vval;
        }
        __syncthreads();

        for (int idx = warp; idx < na; idx += NWARPS) {
            const unsigned short e = sList[base + idx];
            const int qi = e >> 1;
            const int is_top = e & 1;
            const int t = qbase + qi;

            // Whole block if selected; otherwise only the sliding-window tail.
            const int lo = is_top ? kbase : max(kbase, t - WIN + 1);
            const int hi = min(kbase + BLK, t + 1);

            int4 qreg[DV];
#pragma unroll
            for (int c = 0; c < DV; ++c) qreg[c] = reinterpret_cast<const int4*>(&sQ[qi * D])[c];

            // Both of this lane's keys are walked inside one pass over D so each q
            // element is converted once and reused for both (a key-outer loop makes
            // the compiler redo the q conversions: 3 -> 2.5 instructions per MAC).
            // Four accumulators keep four independent FMA chains in flight.
            float a0 = 0.f, a1 = 0.f, b0 = 0.f, b1 = 0.f;
            {
                const int4* krow0 = reinterpret_cast<const int4*>(&sK[lane * KSTRIDE]);
                const int4* krow1 = reinterpret_cast<const int4*>(&sK[(lane + 32) * KSTRIDE]);
#pragma unroll
                for (int c = 0; c < DV; ++c) {
                    const int4 kk0 = krow0[c];
                    const int4 kk1 = krow1[c];
                    const int4 qv = qreg[c];
#pragma unroll
                    for (int u = 0; u < 4; ++u) {
                        const float2 qa =
                            __bfloat1622float2(reinterpret_cast<const __nv_bfloat162*>(&qv)[u]);
                        const float2 k0 =
                            __bfloat1622float2(reinterpret_cast<const __nv_bfloat162*>(&kk0)[u]);
                        const float2 k1 =
                            __bfloat1622float2(reinterpret_cast<const __nv_bfloat162*>(&kk1)[u]);
                        a0 = fmaf(qa.x, k0.x, a0);
                        a1 = fmaf(qa.y, k0.y, a1);
                        b0 = fmaf(qa.x, k1.x, b0);
                        b1 = fmaf(qa.y, k1.y, b1);
                    }
                }
            }
            float s[2];
            {
                const int j0 = kbase + lane, j1 = kbase + lane + 32;
                s[0] = (j0 >= lo && j0 < hi) ? (a0 + a1) * scale : -INFINITY;
                s[1] = (j1 >= lo && j1 < hi) ? (b0 + b1) * scale : -INFINITY;
            }

            const float mb = warp_max(fmaxf(s[0], s[1]));
            const float m_prev = sM[qi];
            const float l_prev = sL[qi];
            const float m_new = fmaxf(m_prev, mb);
            const float p0 = (s[0] == -INFINITY) ? 0.f : __expf(s[0] - m_new);
            const float p1 = (s[1] == -INFINITY) ? 0.f : __expf(s[1] - m_new);
            const float lb = warp_sum(p0 + p1);
            const float alpha = (m_prev == -INFINITY) ? 0.f : __expf(m_prev - m_new);
            if (lane == 0) {
                sM[qi] = m_new;
                sL[qi] = l_prev * alpha + lb;
            }
            sP[warp * BLK + lane] = p0;
            sP[warp * BLK + lane + 32] = p1;
            __syncwarp();

            // acc[d] = acc[d] * alpha + sum_j p_j * V[j][d]
#pragma unroll
            for (int dblk = 0; dblk < D; dblk += 64) {
                float pacc[8];
#pragma unroll
                for (int u = 0; u < 8; ++u) pacc[u] = 0.f;
                for (int jb = 0; jb < BLK; jb += 4) {
                    const int j = jb + jg;
                    const float pj = sP[warp * BLK + j];
                    const int4 vv = *reinterpret_cast<const int4*>(&sV[j * KSTRIDE + dblk + dc * 8]);
#pragma unroll
                    for (int u = 0; u < 4; ++u) {
                        const float2 f =
                            __bfloat1622float2(reinterpret_cast<const __nv_bfloat162*>(&vv)[u]);
                        pacc[2 * u] = fmaf(pj, f.x, pacc[2 * u]);
                        pacc[2 * u + 1] = fmaf(pj, f.y, pacc[2 * u + 1]);
                    }
                }
#pragma unroll
                for (int u = 0; u < 8; ++u) {
                    float x = pacc[u];
                    x += __shfl_xor_sync(0xffffffffu, x, 8);
                    x += __shfl_xor_sync(0xffffffffu, x, 16);
                    pacc[u] = x;
                }
                if (jg == 0) {
                    float* ap = &sAcc[qi * D + dblk + dc * 8];
#pragma unroll
                    for (int u = 0; u < 8; ++u) ap[u] = fmaf(ap[u], alpha, pacc[u]);
                }
            }
            __syncwarp();
        }
    }
    __syncthreads();

    // ---- phase 6: normalize + write ----
    for (int idx = tid; idx < BLK * DV; idx += NTHREADS) {
        const int row = idx / DV, c = idx - row * DV;
        const int t = qbase + row;
        if (t >= S) continue;
        const float inv = 1.f / sL[row];
        const float* ap = &sAcc[row * D + c * 8];
        __nv_bfloat16 out[8];
#pragma unroll
        for (int u = 0; u < 8; ++u) out[u] = __float2bfloat16(ap[u] * inv);
        reinterpret_cast<int4*>(Ob + (long)t * D)[c] = *reinterpret_cast<const int4*>(out);
    }
}

// ---------------------------------------------------------------- launcher

template <int D>
static void launch_for_D(const torch::Tensor& q, const torch::Tensor& k, const torch::Tensor& v,
                         torch::Tensor& kbarT, torch::Tensor& out, int B, int H, int S, int NB) {
    auto stream = at::cuda::getCurrentCUDAStream();
    const __nv_bfloat16* qp = (const __nv_bfloat16*)q.data_ptr();
    const __nv_bfloat16* kp = (const __nv_bfloat16*)k.data_ptr();
    const __nv_bfloat16* vp = (const __nv_bfloat16*)v.data_ptr();
    float* kbp = kbarT.data_ptr<float>();
    __nv_bfloat16* op = (__nv_bfloat16*)out.data_ptr();

    kbar_kernel<D><<<dim3(NB, H, B), D, 0, stream>>>(kp, kbp, S, NB);

    const int n_qtiles = NB;
    const int smem = plan_smem(D, NB).total;
    dim3 grid(n_qtiles, H, B);

// cudaFuncSetAttribute is a driver call; doing it per forward() adds a fixed
// host-side cost to every launch. Set it only when the request changes.
#define NSA_LAUNCH(MAXB)                                                                       \
    do {                                                                                       \
        static int last_smem_##MAXB = -1;                                                      \
        if (last_smem_##MAXB != smem) {                                                        \
            cudaFuncSetAttribute(nsa_kernel<D, MAXB>,                                          \
                                 cudaFuncAttributeMaxDynamicSharedMemorySize, smem);           \
            last_smem_##MAXB = smem;                                                           \
        }                                                                                      \
        nsa_kernel<D, MAXB><<<grid, NTHREADS, smem, stream>>>(qp, kp, vp, kbp, op, S, NB,      \
                                                              n_qtiles);                       \
    } while (0)

    // MAXB caps at 4 (128 blocks scored per pass); longer sequences loop passes
    // and merge into the running top-8, which also keeps the staged pooled-key
    // chunk inside the sK/sV/sAcc scratch.
    if (NB <= 32) {
        NSA_LAUNCH(1);
    } else if (NB <= 64) {
        NSA_LAUNCH(2);
    } else {
        NSA_LAUNCH(4);
    }
#undef NSA_LAUNCH
}

torch::Tensor nsa_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v) {
    TORCH_CHECK(q.is_cuda() && k.is_cuda() && v.is_cuda(), "inputs must be CUDA tensors");
    TORCH_CHECK(q.scalar_type() == torch::kBFloat16, "q must be bf16");
    q = q.contiguous();
    k = k.contiguous();
    v = v.contiguous();
    const int B = q.size(0), H = q.size(1), S = q.size(2), D = q.size(3);
    const int NB = (S + BLK - 1) / BLK;
    TORCH_CHECK(NB <= 256, "sequence too long for the CUDA path");
    {   // the staged pooled-key chunk must fit the sK/sV/sAcc scratch
        const auto sp = plan_smem(D, NB);
        const int scratch = sp.acc + BLK * D * 4 - sp.k;
        const int chunk = (NB <= 32) ? 32 : (NB <= 64 ? 64 : 128);
        TORCH_CHECK(D * chunk * 4 <= scratch, "kbar scratch too small: need ", D * chunk * 4,
                    " have ", scratch);
    }

    auto out = torch::empty_like(q);
    auto kbarT = torch::empty({B, H, D, NB}, q.options().dtype(torch::kFloat32));

    if (D == 64) {
        launch_for_D<64>(q, k, v, kbarT, out, B, H, S, NB);
    } else if (D == 128) {
        launch_for_D<128>(q, k, v, kbarT, out, B, H, S, NB);
    } else {
        TORCH_CHECK(false, "unsupported head dim ", D);
    }
    return out;
}
"""

_CPP_SRC = "torch::Tensor nsa_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v);"

_EXT = None


def _ext():
    global _EXT
    if _EXT is None:
        os.environ.setdefault("TORCH_CUDA_ARCH_LIST", "12.0")
        _EXT = load_inline(
            name="nsa_sparse_attn_sm120",
            cpp_sources=[_CPP_SRC],
            cuda_sources=[_CUDA_SRC],
            functions=["nsa_forward"],
            extra_cuda_cflags=[
                "-O3",
                "-lineinfo",
                "--expt-relaxed-constexpr",
                "-gencode=arch=compute_120,code=sm_120",
            ],
            extra_cflags=["-O3"],
            verbose=False,
        )
    return _EXT


# Supported by the CUDA path: head dims 64/128 and up to 256 key blocks (S <= 16384).
def _cuda_supported(q: torch.Tensor) -> bool:
    S, D = q.shape[2], q.shape[3]
    return q.is_cuda and D in (64, 128) and (S + BLOCK_SIZE - 1) // BLOCK_SIZE <= 256


def _nsa_fallback(
    q: torch.Tensor,
    k: torch.Tensor,
    v: torch.Tensor,
    block_size: int = BLOCK_SIZE,
    top_n_blocks: int = TOP_N_BLOCKS,
    sliding_window: int = SLIDING_WINDOW,
) -> torch.Tensor:
    """Vectorised torch implementation of the same semantics.

    Only used for shapes outside the CUDA path's envelope (head dim not 64/128,
    or S > 16384). Correct but dense in memory; never hit by the benchmark deck.
    """
    B, H, S, D = q.shape
    scale = 1.0 / math.sqrt(D)
    nb = (S + block_size - 1) // block_size
    pad = nb * block_size - S
    idx = torch.arange(S, device=q.device)
    causal = idx[:, None] >= idx[None, :]
    cnt = torch.nn.functional.pad(causal.float(), (0, pad)).view(S, nb, block_size).sum(-1)
    window = (idx[:, None] - idx[None, :]) < sliding_window
    out = torch.empty_like(q)
    for b in range(B):
        for h in range(H):
            sc = (q[b, h].float() @ k[b, h].float().T) * scale
            blk = (
                torch.nn.functional.pad(torch.where(causal, sc, torch.zeros((), device=q.device)), (0, pad))
                .view(S, nb, block_size)
                .sum(-1)
            )
            imp = torch.where(cnt > 0, blk / cnt.clamp(min=1), torch.full_like(blk, -1e9))
            top = imp.topk(min(top_n_blocks, nb), dim=-1).indices
            selb = torch.zeros(S, nb, dtype=torch.bool, device=q.device)
            selb.scatter_(-1, top, True)
            keyblk = torch.arange(nb * block_size, device=q.device) // block_size
            mask = selb[:, keyblk][:, :S]
            keep = causal & (mask | window)
            sc = sc.masked_fill(~keep, float("-inf"))
            out[b, h] = (torch.softmax(sc, dim=-1) @ v[b, h].float()).to(q.dtype)
    return out


def nsa_attend(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor) -> torch.Tensor:
    if _cuda_supported(q):
        return _ext().nsa_forward(q, k, v)
    return _nsa_fallback(q, k, v)


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))

    def forward(self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor) -> torch.Tensor:
        return nsa_attend(q, k, v)


def get_init_inputs():
    return [1, 16, 1024, 64]


def get_inputs():
    B, H, S, D = get_init_inputs()
    q = torch.randn(B, H, S, D, dtype=torch.bfloat16)
    k = torch.randn(B, H, S, D, dtype=torch.bfloat16)
    v = torch.randn(B, H, S, D, dtype=torch.bfloat16)
    return [q, k, v]

20260716_140658_claude_claude-opus-4-8_02_deepseek_nsa