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