kernelbench.com

KernelBench cuda · RTX PRO 6000

DeepSeek NSA Claude Opus 5

103.7%geomean peak fraction across shapes

manually audited: clean

1541-line nsa_kernel.cu (cpp_extension.load): kbar block-means in fp16 hi/lo pairs, mma-based block-importance selection with reference tie-break, work-unit flash-style attention producing (m,l,o/l) partials + combine kernel; 64-bit block-mask k_flash path for S<=4096. Real f16 mma.sync/ ldmatrix throughout. Only cache is _plan()/_CACHE holding shape-keyed scratch workspaces, never outputs. EMPIRICAL (mandatory for >1.0): output vs reference max-err 0.0075; same q/k/v buffers overwritten in place -> output changed by 5.31 max-abs and matched a fresh reference on the new values at 0.0081. Recomputes on live data. peak>1.0 is the documented dense-equivalent FLOPs roofline artifact; ms is the headline metric for this problem. template_mutated=false, numeric stress on, zero cross-run access. Sequential isolated re-grade on anvil GPU0 2026-07-26 (contended 1.0413, delta -0.4%).

harnessor-opusagent session13h 43mtotal wall13h 44mcheck2sbenchmark2soutput tokensgpu-lock wait0sgpu-lock held29sregimecompute

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

1×16×2048×640.067 ms51.6%258 TFLOPS · 52% of 500 TF bf16 peak · also 0.25 TB/s (14% of HBM)
1×16×4127×640.110 ms126.4%632 TFLOPS · 126% of 500 TF bf16 peak · also 0.31 TB/s (17% of HBM)
1×8×8192×640.118 ms233.4%1,167 TFLOPS · 233% of 500 TF bf16 peak · also 0.28 TB/s (16% of HBM)
1×8×8191×1280.221 ms248.9%1,245 TFLOPS · 249% of 500 TF bf16 peak · also 0.30 TB/s (17% of HBM)
4×8×1024×640.045 ms38.0%190 TFLOPS · 38% of 500 TF bf16 peak · also 0.37 TB/s (21% of HBM)
2×8×3000×640.086 ms86.2%431 TFLOPS · 86% of 500 TF bf16 peak · also 0.29 TB/s (16% of HBM)

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

geomean(51.6% · 126.4% · 233.4% · 248.9% · 38.0% · 86.2%) = 103.7%

Kernel source (redacted)
"""DeepSeek-NSA style sparse attention — fused CUDA kernels for SM120 (Blackwell).

Pipeline (see nsa_kernel.cu for the kernels themselves):
  1. k_kbar    - per key-block mean of K, stored as an fp16 hi/lo pair so that
                 q . kbar keeps fp32 accuracy through the tensor cores.
  2. k_select  - block importance (mma) + exact diagonal-block importance,
                 top-8 with the reference tie-break, scatter into per-key-block
                 query lists.
  3. k_attend  - work-unit driven mma attention (dense causal prefix, sparse
                 selected blocks, sliding window) producing one partial
                 (m, l, o/l) per (query, slot).
  4. k_combine - merge the <= 9 partials of each query into the bf16 output.

For S <= 4096 a query's selection fits in a 64-bit block mask and step 2 emits
masks instead of query lists: k_flash then sweeps every causal key block of a
query tile with per-row column masks (Q read once, ceil(tile/16) partials).
"""

import math
import os

import torch
import torch.nn as nn
from torch.utils.cpp_extension import load

BLOCK_SIZE = 64
TOP_N_BLOCKS = 8
SLIDING_WINDOW = 64
NSLOT = 9
# k_flash partial slots per query: the flat work split can straddle a tile boundary,
# so a query's row is spread over at most this many units (must match nsa_kernel.cu).
FLASH_MAXTILE = 40  # k_flash ntile bound (match .cu)

_HERE = os.path.dirname(os.path.abspath(__file__))
# The container preseeds a multi-arch list; we only ever run on SM120 and the
# mma/ldmatrix PTX below needs sm_80+, so pin the target (also cuts build time).
os.environ["TORCH_CUDA_ARCH_LIST"] = "12.0"

_EXT = None


def _ext():
    global _EXT
    if _EXT is None:
        _EXT = load(
            name="nsa_sparse_attn",
            sources=[os.path.join(_HERE, "nsa_kernel.cu")],
            extra_cuda_cflags=[
                "-O3",
                "--use_fast_math",
                "-lineinfo",
                "--expt-relaxed-constexpr",
            ],
            verbose=False,
        )
    return _EXT


_CACHE = {}


def _plan(B, H, S, D, device):
    key = (B, H, S, D, device)
    p = _CACHE.get(key)
    if p is not None:
        return p
    BH = B * H
    ntile = (S + BLOCK_SIZE - 1) // BLOCK_SIZE
    dev = torch.device(device)
    f16 = dict(dtype=torch.float16, device=dev)
    i32 = dict(dtype=torch.int32, device=dev)
    sparse = ntile > 8
    if ntile <= FLASH_MAXTILE:  # k_flash: no query lists, NSF slots, one range
        nslot = 4
        p = dict(
            RS=S,
            kb_hi=torch.empty(BH * ntile * D, **f16) if sparse else torch.empty(1, **f16),
            kb_lo=torch.empty(BH * ntile * D, **f16) if sparse else torch.empty(1, **f16),
            ent=torch.empty(1, **i32),
            cnt=torch.zeros(1, **i32),
            po=torch.empty(BH * S * nslot * D, **f16),
            pml=torch.empty(BH * S * nslot * 2, dtype=torch.float32, device=dev),
            msk=torch.empty(BH * S, dtype=torch.int64, device=dev),
            # one arrival ticket per (bh, query tile) for k_flash's fused merge; k_flash
            # resets each ticket as it consumes it, so this stays zero between calls
            ctr=torch.zeros(BH * ntile, **i32),
        )
        _CACHE[key] = p
        return p
    # Partial-output working set is the dominant traffic; keep it L2 resident by
    # splitting the query axis into ranges of RS queries.  L2 is 128 MB: 96 MB keeps
    # the partials resident with room for the K/V stream, and stays on the safe side
    # of the cliff where a single range spills L2 (measured: -12% at S=8191, D=128).
    budget = 96 << 20
    n = 1
    while True:
        RS = 512 * ((S + 512 * n - 1) // (512 * n))
        RS = max(RS, 512)
        if BH * RS * NSLOT * D * 2 <= budget or RS <= 512:
            break
        n += 1
    RS = min(RS, 512 * ((S + 511) // 512))
    nrange = (S + RS - 1) // RS
    p = dict(
        RS=RS,
        kb_hi=torch.empty(BH * ntile * D, **f16),
        kb_lo=torch.empty(BH * ntile * D, **f16),
        ent=torch.empty(BH * ntile * RS, **i32),
        cnt=torch.zeros(nrange * BH * ntile, **i32),
        po=torch.empty(BH * RS * NSLOT * D, **f16),
        pml=torch.empty(BH * RS * NSLOT * 2, dtype=torch.float32, device=dev),
        msk=torch.empty(1, dtype=torch.int64, device=dev),
        ctr=torch.zeros(1, **i32),
    )
    _CACHE[key] = p
    return p


def nsa_attend(q, k, v):
    B, H, S, D = q.shape
    q = q.contiguous()
    k = k.contiguous()
    v = v.contiguous()
    p = _plan(B, H, S, D, q.device.index if q.device.index is not None else 0)
    out = torch.empty_like(q)
    _ext().nsa_run(
        q, k, v, out, p["kb_hi"], p["kb_lo"], p["ent"], p["cnt"], p["po"], p["pml"], p["msk"],
        p["ctr"], p["RS"],
    )
    return out


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, k, v):
        return nsa_attend(q, k, v)


B, H, S, D = 1, 16, 1024, 64


def get_init_inputs():
    return [B, H, S, D]


def get_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]


# ==================================================================
# ===== sidecar: nsa_kernel.cu (71573 bytes, loaded by solution.py) =====
# ==================================================================

// DeepSeek-NSA style sparse attention, fused CUDA implementation (sm_120 / Blackwell).
//
// Bench semantics (per query t, causal):
//   1. keys split in blocks of 64
//   2. block importance = mean over causal keys j in block of (q_t . k_j)/sqrt(D)
//        == q_t . kbar_block / sqrt(D)         (full blocks: kbar = block mean of k)
//        == q_t . (prefix sum of k)/(n) /sqrt(D)  (the diagonal, partially-causal block)
//   3. top-8 blocks by (importance, block_idx) descending, union sliding window [t-63, t]
//   4. softmax attention over that union.
//
// Strategy: the selection is per-query and (for random data) essentially uncorrelated
// between neighbouring queries, so a query-tile kernel would degenerate to dense.
// Instead we *invert* the sparsity: for every key block we gather the list of queries
// that selected it, run dense 64x64 mma tiles over (query-chunk x key-block), emit one
// partial softmax (m, l, o/l) per (query, slot), and combine the <=9 partials per query.
//
//   k_kbar    : block means of K (fp16 hi/lo split -> exact fp32-quality dot products)
//   k_select  : importance mma + diagonal-block importance + top-8 + scatter to per-block lists
//   k_attend  : work-unit driven mma attention -> partials
//   k_combine : merge partials -> bf16 output
//
// Queries with t < 512 have <= 8 causal blocks, so *all* of them are selected: that
// region is exactly dense causal attention and is handled by "dense" work units.
//
// Inverting the sparsity costs one Q read and one partial write per (query, selected
// block) -- 9x each -- and at small S that L2 traffic, not the mma, is the wall.  So
// when the whole selection of a query fits in a 64-bit mask (ntile <= 64) we use the
// other structure instead: k_select emits per-query bitmasks and k_flash sweeps all
// causal key blocks of a query tile with per-row column masks, reading Q once per work
// unit and emitting at most one partial per query.  Above that ntile the wasted mma
// wins and the inverted path takes over.

#include <ATen/cuda/CUDAContext.h>
#include <cuda_bf16.h>
#include <cuda_fp16.h>
#include <cuda_runtime.h>
#include <torch/extension.h>

#include <algorithm>
#include <cmath>

#define DEVI __device__ __forceinline__
#define BS 64        // block_size
#define TQ 64        // query tile
#define NSLOT 9      // 8 selected blocks + 1 merged sliding-window partial
#define NWARP 4
#define NTHREAD 128
#define NEG 1e30f
#define LOG2E 1.4426950408889634f
#define MAXCAND 8    // candidate blocks per lane in the k_select_big scan (S <= 32*64*MAXCAND)
#define SELNCH 2     // 64-key-block chunks held in registers by k_select (S <= 4096*SELNCH)
#define SELMB1 6     // resident k_select blocks per SM, one candidate chunk (free: 70 regs)
#define SELMB2 5     // ... two chunks, but few tiles reach the second: a small spill buys
#define SELMB3 4     // ... two chunks in earnest: 128 registers, no spill
#define QPW 1  // rows per warp in k_combine
#define NSF 4  // k_flash partial slots per query; W is picked so a row never needs more
// A key block selected by many queries is one k_attend unit with many query chunks, and
// the fattest such unit can outweigh a whole block's share of the launch (its selection
// count runs ~3x the mean).  Split a unit into up to GSPLIT slices, but only while each
// slice keeps >= SMINCH chunks, so thin units stay whole and pay no extra K/V staging.
#define GSPLIT 2
#define SMINCH 3
#define FLASH_MAXTILE 40  // masked-dense sweeps ntile*(ntile+1)/2 tile pairs per bh
// while the inverted path costs ~10*ntile plus gather/partial/merge overhead, so the
// two cross near ntile ~40 (measured: ntile 32 favours dense by 5%, ntile 47 favours
// the inverted path by 24%)  // k_flash needs the selection of a query to fit in one u64

// ---------------------------------------------------------------- small helpers

DEVI unsigned sm_u32(const void* p) { return (unsigned)__cvta_generic_to_shared(p); }

DEVI float bf2f(const __nv_bfloat16 x) {
  return __int_as_float(((unsigned)(*(const unsigned short*)&x)) << 16);
}

// 8 bf16 (16B) -> 8 fp16 (16B)
DEVI void cvt8(const uint4 src, uint4& dst) {
  const unsigned short* s = (const unsigned short*)&src;
  __half h[8];
#pragma unroll
  for (int i = 0; i < 8; ++i) h[i] = __float2half_rn(__int_as_float(((unsigned)s[i]) << 16));
  dst = *(const uint4*)h;
}

DEVI void ldm_x4(unsigned a, unsigned& r0, unsigned& r1, unsigned& r2, unsigned& r3) {
  asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0,%1,%2,%3}, [%4];\n"
               : "=r"(r0), "=r"(r1), "=r"(r2), "=r"(r3) : "r"(a));
}
DEVI void ldm_x2(unsigned a, unsigned& r0, unsigned& r1) {
  asm volatile("ldmatrix.sync.aligned.m8n8.x2.shared.b16 {%0,%1}, [%2];\n"
               : "=r"(r0), "=r"(r1) : "r"(a));
}
DEVI void ldm_x2t(unsigned a, unsigned& r0, unsigned& r1) {
  asm volatile("ldmatrix.sync.aligned.m8n8.x2.trans.shared.b16 {%0,%1}, [%2];\n"
               : "=r"(r0), "=r"(r1) : "r"(a));
}
DEVI void ldm_x4t(unsigned a, unsigned& r0, unsigned& r1, unsigned& r2, unsigned& r3) {
  asm volatile("ldmatrix.sync.aligned.m8n8.x4.trans.shared.b16 {%0,%1,%2,%3}, [%4];\n"
               : "=r"(r0), "=r"(r1), "=r"(r2), "=r"(r3) : "r"(a));
}
// D += A(16x16) * B(16x8), fp16 in / fp32 acc.  q,k,v are bf16 in memory, hence exactly
// representable in fp16, so these products carry no extra rounding vs fp32 inputs.
// Scalar refs (not float*/unsigned* arrays): taking the address of an accumulator
// tile is enough to make ptxas keep the whole array in local memory.
DEVI void mma_f16(float& d0, float& d1, float& d2, float& d3, unsigned a0, unsigned a1,
                  unsigned a2, unsigned a3, unsigned b0, unsigned b1) {
  asm volatile("mma.sync.aligned.m16n8k16.row.col.f32.f16.f16.f32 "
               "{%0,%1,%2,%3}, {%4,%5,%6,%7}, {%8,%9}, {%0,%1,%2,%3};\n"
               : "+f"(d0), "+f"(d1), "+f"(d2), "+f"(d3)
               : "r"(a0), "r"(a1), "r"(a2), "r"(a3), "r"(b0), "r"(b1));
}
// 16B global -> shared async copy; !pred zero-fills (out-of-range rows)
DEVI void cp16(unsigned dst, const void* src, bool pred) {
  asm volatile("cp.async.ca.shared.global [%0], [%1], 16, %2;\n" ::"r"(dst), "l"(src),
               "r"(pred ? 16 : 0));
}
DEVI void cp_commit() { asm volatile("cp.async.commit_group;\n" ::); }
template <int N>
DEVI void cp_wait() {
  asm volatile("cp.async.wait_group %0;\n" ::"n"(N));
}
// bf16 variant of mma_f16: q/k/v are bf16 in memory so this needs no conversion at
// all, and P in bf16 costs far less accuracy than the bf16 output rounding.
DEVI void mma_bf16(float& d0, float& d1, float& d2, float& d3, unsigned a0, unsigned a1,
                   unsigned a2, unsigned a3, unsigned b0, unsigned b1) {
  asm volatile("mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32 "
               "{%0,%1,%2,%3}, {%4,%5,%6,%7}, {%8,%9}, {%0,%1,%2,%3};\n"
               : "+f"(d0), "+f"(d1), "+f"(d2), "+f"(d3)
               : "r"(a0), "r"(a1), "r"(a2), "r"(a3), "r"(b0), "r"(b1));
}
DEVI unsigned pack2bf(float lo, float hi) {
  __nv_bfloat162 h = __floats2bfloat162_rn(lo, hi);
  return *(unsigned*)&h;
}
DEVI unsigned pack2(float lo, float hi) {
  __half2 h = __floats2half2_rn(lo, hi);
  return *(unsigned*)&h;
}
// order preserving float -> uint32 (for 64-bit (imp, idx) max-reduction keys)
DEVI unsigned f2u(float f) {
  unsigned b = __float_as_uint(f);
  return (b & 0x80000000u) ? ~b : (b | 0x80000000u);
}

// fragment row/col ownership of an m16n8k16 accumulator
#define FRAG_R0(lane) ((lane) >> 2)
#define FRAG_C(lane) (((lane) & 3) << 1)

// XOR swizzle for an unpadded (row stride == D halves) shared tile: the 16B unit index
// inside a row is xored with the low 3 bits of the row, so the 8 rows one ldmatrix
// group touches land in 8 distinct 16B chunks == all 32 banks, exactly once.  Costs
// nothing and, unlike the +8 pad, keeps the tile a power of two -- which is what buys
// k_attend a second (D=128: 1 -> 2) resident block per SM.
#define SWZ(row, col) ((row) * D + (((((col) >> 3) ^ ((row) & 7))) << 3) + ((col) & 7))

// ---------------------------------------------------------------- k_kbar

// grid (ntile, BH), NTHREAD threads.  Also zeroes the per-range counters.
template <int D>
__global__ void k_kbar(const __nv_bfloat16* __restrict__ Kg, __half* __restrict__ kb_hi,
                       __half* __restrict__ kb_lo, int S, int ntile, int* __restrict__ cnt,
                       int nrange) {
  const int bi = blockIdx.x, bh = blockIdx.y, BH = gridDim.y;
  for (int r = threadIdx.x; r < nrange; r += NTHREAD) cnt[(r * BH + bh) * ntile + bi] = 0;

  const int s0 = bi * BS;
  __half* oh = kb_hi + (size_t)(bh * ntile + bi) * D;
  __half* ol = kb_lo + (size_t)(bh * ntile + bi) * D;
  if (s0 + BS > S) {  // last, partial block: never used as a "full" block
    for (int d = threadIdx.x; d < D; d += NTHREAD) { oh[d] = __float2half(0.f); ol[d] = __float2half(0.f); }
    return;
  }
  // 16B per thread per load.  One element per thread turns the 8 KB tile into 4096
  // two-byte requests, which left this kernel latency-bound at 0.8 TB/s on a cold L2
  // (it is pure streaming: nothing else here can hide a load).
  constexpr int CPR = D / 8;            // 16B chunks per row
  constexpr int RPP = NTHREAD / CPR;    // rows covered per pass
  const int tid = threadIdx.x, c = tid % CPR, r0 = tid / CPR;
  const __nv_bfloat16* kp = Kg + (size_t)(bh * S + s0) * D + c * 8;
  float a[8];
#pragma unroll
  for (int e = 0; e < 8; ++e) a[e] = 0.f;
#pragma unroll
  for (int i = 0; i < BS / RPP; ++i) {
    const uint4 raw = *(const uint4*)(kp + (size_t)(r0 + i * RPP) * D);
    const unsigned short* s = (const unsigned short*)&raw;
#pragma unroll
    for (int e = 0; e < 8; ++e) a[e] += __int_as_float(((unsigned)s[e]) << 16);
  }
  __shared__ float red[NTHREAD * 9];  // 9, not 8: the store below is bank-conflict free
#pragma unroll
  for (int e = 0; e < 8; ++e) red[tid * 9 + e] = a[e];
  __syncthreads();
  for (int d = tid; d < D; d += NTHREAD) {
    float acc = 0.f;
#pragma unroll
    for (int g = 0; g < RPP; ++g) acc += red[(g * CPR + (d >> 3)) * 9 + (d & 7)];
    const float mean = acc * (1.f / (float)BS);
    const __half hi = __float2half_rn(mean);
    oh[d] = hi;
    ol[d] = __float2half_rn(mean - __half2float(hi));
  }
}

// ---------------------------------------------------------------- k_select
//
// One block per (query tile >= 8, bh).  Computes all block importances for the 64
// queries of the tile, selects the top-8 with the reference tie-break (importance
// desc, block index desc) and scatters (query, slot) into the per-key-block lists.
//
// The importances never leave the registers.  In an m16n8k16 accumulator lane l
// owns rows (l>>2, (l>>2)+8) and columns (l&3)*2 + {0,1} of every n-tile, so the
// 4 lanes of a row group hold all 64 candidates of a 64-block chunk and the whole
// 64-query tile is selected in 8 warp-parallel rounds -- versus 64x8 reductions
// over a shared importance matrix, which measured 60% of this kernel's time.
// NCH chunks are kept live, i.e. S <= NCH*4096; longer goes to k_select_big.

// 64 queries x 64 keys, one warp = 16 query rows.  A macro, not a function: passing
// the accumulator by pointer is enough to make ptxas spill it to local memory.
#define MMA_TILE_A(SQ, SB, ACC, ADR)                                                            \
  do {                                                                                          \
    _Pragma("unroll 1")                                                                         \
    for (int kt = 0; kt < D / 16; ++kt) {                                                       \
      unsigned a0, a1, a2, a3;                                                                  \
      ldm_x4(sm_u32(ADR(SQ, warp * 16 + 8 * ((lane >> 3) & 1) + (lane & 7),                     \
                        kt * 16 + 8 * (lane >> 4))),                                            \
             a0, a1, a2, a3);                                                                   \
      _Pragma("unroll")                                                                         \
      for (int nt = 0; nt < 8; ++nt) {                                                          \
        unsigned b0, b1;                                                                        \
        ldm_x2(sm_u32(ADR(SB, nt * 8 + (lane & 7), kt * 16 + 8 * ((lane >> 3) & 1))), b0, b1);   \
        mma_f16(ACC[nt][0], ACC[nt][1], ACC[nt][2], ACC[nt][3], a0, a1, a2, a3, b0, b1);        \
      }                                                                                         \
    }                                                                                           \
  } while (0)
#define ADR_S(B, r, c) ((B) + SWZ(r, c))
#define MMA_TILE_W(SQ, SB, ACC) MMA_TILE_A(SQ, SB, ACC, ADR_S)

// With mskg != nullptr the selection is emitted as one 64-bit block mask per query
// (for k_flash) and the whole scatter machinery below is skipped.
// MINB is the resident-block target: the top-8 rounds below are a long dependency chain
// with little to overlap, so the kernel wants warps, and paying a small spill for one more
// block per SM is worth it (measured -4% on k_select at NCH=2).
template <int D, int NCH, int MINB = 4>
__global__ __launch_bounds__(NTHREAD, MINB) void k_select(
    const __nv_bfloat16* __restrict__ Qg, const __nv_bfloat16* __restrict__ Kg,
    const __half* __restrict__ kb_hi, const __half* __restrict__ kb_lo,
    int* __restrict__ ent, int* __restrict__ cnt, float2* __restrict__ pml,
    unsigned long long* __restrict__ mskg, int S, int ntile, int tile0, int T0, int RS, int cap,
    float scale) {
  const int LD = D;  // unpadded: rows are XOR-swizzled (SWZ) instead
  extern __shared__ char smem[];
  __half* sq = (__half*)smem;  // TQ x D
  __half* sb = sq + TQ * LD;   // 64 x D : diagonal K, then each kbar chunk

  const int tile = tile0 + blockIdx.x, bh = blockIdx.y;
  const int tid = threadIdx.x, lane = tid & 31, warp = tid >> 5;
  const int t_base = tile * BS;
  const int CPR = D / 8;  // 16B chunks per row

  // ---- load q tile (bf16 -> fp16)
  for (int i0 = 0; i0 < TQ * CPR; i0 += NTHREAD * 4) {
    uint4 raw[4];  // all four in flight: one round trip, not four
#pragma unroll
    for (int j = 0; j < 4; ++j) {
      const int idx = i0 + tid + j * NTHREAD, rr = idx / CPR, cc = (idx - rr * CPR) * 8;
      raw[j] = make_uint4(0, 0, 0, 0);
      if (idx < TQ * CPR && t_base + rr < S)
        raw[j] = *(const uint4*)(Qg + (size_t)(bh * S + t_base + rr) * D + cc);
    }
#pragma unroll
    for (int j = 0; j < 4; ++j) {
      const int idx = i0 + tid + j * NTHREAD, rr = idx / CPR, cc = (idx - rr * CPR) * 8;
      if (idx < TQ * CPR) {
        uint4 o;
        cvt8(raw[j], o);
        *(uint4*)(sq + SWZ(rr, cc)) = o;
      }
    }
  }

  const int r0 = warp * 16 + FRAG_R0(lane), r1 = r0 + 8, cb = FRAG_C(lane);
  const int gl = lane & 3;  // column position inside the row group

  // ---- diagonal block: importance = causal row-mean of its own 64x64 score tile.
  // First, so that its accumulator is dead before the chunk accumulators go live.
  float dv0, dv1;
  {
    for (int i0 = 0; i0 < BS * CPR; i0 += NTHREAD * 4) {
      uint4 raw[4];
#pragma unroll
      for (int j = 0; j < 4; ++j) {
        const int idx = i0 + tid + j * NTHREAD, rr = idx / CPR, cc = (idx - rr * CPR) * 8;
        raw[j] = make_uint4(0, 0, 0, 0);
        if (idx < BS * CPR && t_base + rr < S)
          raw[j] = *(const uint4*)(Kg + (size_t)(bh * S + t_base + rr) * D + cc);
      }
#pragma unroll
      for (int j = 0; j < 4; ++j) {
        const int idx = i0 + tid + j * NTHREAD, rr = idx / CPR, cc = (idx - rr * CPR) * 8;
        if (idx < BS * CPR) {
          uint4 o;
          cvt8(raw[j], o);
          *(uint4*)(sb + SWZ(rr, cc)) = o;
        }
      }
    }
    __syncthreads();
    float d[8][4];
#pragma unroll
    for (int nt = 0; nt < 8; ++nt)
#pragma unroll
      for (int i = 0; i < 4; ++i) d[nt][i] = 0.f;
    MMA_TILE_W(sq, sb, d);
    float s0 = 0.f, s1 = 0.f;
#pragma unroll
    for (int nt = 0; nt < 8; ++nt) {
      const int c = nt * 8 + cb;
      if (c <= r0) s0 += d[nt][0];
      if (c + 1 <= r0) s0 += d[nt][1];
      if (c <= r1) s1 += d[nt][2];
      if (c + 1 <= r1) s1 += d[nt][3];
    }
#pragma unroll
    for (int off = 1; off < 4; off <<= 1) {  // all 4 lanes end up with the row sum
      s0 += __shfl_xor_sync(0xffffffffu, s0, off);
      s1 += __shfl_xor_sync(0xffffffffu, s1, off);
    }
    dv0 = s0 * scale / (float)(r0 + 1);
    dv1 = s1 * scale / (float)(r1 + 1);
  }

  // ---- importance of the full blocks 0..tile-1  (kbar = hi + lo, two mma passes)
  float acc[NCH][8][4];
#pragma unroll
  for (int ch = 0; ch < NCH; ++ch) {
#pragma unroll
    for (int nt = 0; nt < 8; ++nt)
#pragma unroll
      for (int i = 0; i < 4; ++i) acc[ch][nt][i] = 0.f;
    if (ch * BS >= tile) continue;  // uniform over the block
    const int nbk = min(BS, tile - ch * BS);
#pragma unroll 1
    for (int pass = 0; pass < 2; ++pass) {
      const __half* src = pass ? kb_lo : kb_hi;
      __syncthreads();
      for (int i0 = 0; i0 < BS * CPR; i0 += NTHREAD * 4) {
        uint4 raw[4];
#pragma unroll
        for (int j = 0; j < 4; ++j) {
          const int idx = i0 + tid + j * NTHREAD, rr = idx / CPR, cc = (idx - rr * CPR) * 8;
          raw[j] = make_uint4(0, 0, 0, 0);
          if (idx < BS * CPR && rr < nbk)
            raw[j] = *(const uint4*)(src + (size_t)(bh * ntile + ch * BS + rr) * D + cc);
        }
#pragma unroll
        for (int j = 0; j < 4; ++j) {
          const int idx = i0 + tid + j * NTHREAD, rr = idx / CPR, cc = (idx - rr * CPR) * 8;
          if (idx < BS * CPR) *(uint4*)(sb + SWZ(rr, cc)) = raw[j];
        }
      }
      __syncthreads();
      MMA_TILE_W(sq, sb, acc[ch]);
    }
  }
#pragma unroll
  for (int ch = 0; ch < NCH; ++ch)
#pragma unroll
    for (int nt = 0; nt < 8; ++nt) {
      const int c = ch * BS + nt * 8 + cb;
      acc[ch][nt][0] = (c < tile) ? acc[ch][nt][0] * scale : -NEG;
      acc[ch][nt][2] = (c < tile) ? acc[ch][nt][2] * scale : -NEG;
      acc[ch][nt][1] = (c + 1 < tile) ? acc[ch][nt][1] * scale : -NEG;
      acc[ch][nt][3] = (c + 1 < tile) ? acc[ch][nt][3] * scale : -NEG;
    }

  // ---- scatter staging (the q/kbar tiles are dead from here on)
  int* lcnt = (int*)smem;     // ntile
  int* lbase = lcnt + ntile;  // ntile
  int* sA = lbase + ntile;    // TQ*8 : bi<<10 | rank-within-this-block
  int* sB = sA + TQ * 8;      // TQ*8 : (t<<4) | slot
  if (!mskg) {
    __syncthreads();
    for (int i = tid; i < ntile; i += NTHREAD) lcnt[i] = 0;
    for (int i = tid; i < TQ * 8; i += NTHREAD) sA[i] = -1;
    __syncthreads();
  }

  // ---- top-8: 8 rounds of (max over the row's candidates, drop the winner).
  // Every row group runs its own 2 rows, so all 64 queries advance together.
  int sl_a0 = 0, sl_a1 = 0, sl_b0 = 0, sl_b1 = 0;  // slots gl and gl+4 of rows r0/r1
#pragma unroll 1
  for (int s = 0; s < 8; ++s) {
    float bv0 = dv0, bv1 = dv1;
#pragma unroll
    for (int ch = 0; ch < NCH; ++ch) {
      if (ch * BS >= tile) continue;
#pragma unroll
      for (int nt = 0; nt < 8; ++nt) {
        bv0 = fmaxf(bv0, fmaxf(acc[ch][nt][0], acc[ch][nt][1]));
        bv1 = fmaxf(bv1, fmaxf(acc[ch][nt][2], acc[ch][nt][3]));
      }
    }
#pragma unroll
    for (int off = 1; off < 4; off <<= 1) {
      bv0 = fmaxf(bv0, __shfl_xor_sync(0xffffffffu, bv0, off));
      bv1 = fmaxf(bv1, __shfl_xor_sync(0xffffffffu, bv1, off));
    }
    // largest block index wins ties, as in the reference sort key
    int wi0 = (dv0 == bv0) ? tile : -1, wi1 = (dv1 == bv1) ? tile : -1;
#pragma unroll
    for (int ch = 0; ch < NCH; ++ch) {
      if (ch * BS >= tile) continue;
#pragma unroll
      for (int nt = 0; nt < 8; ++nt) {
        const int c = ch * BS + nt * 8 + cb;
        if (acc[ch][nt][0] == bv0) wi0 = max(wi0, c);
        if (acc[ch][nt][1] == bv0) wi0 = max(wi0, c + 1);
        if (acc[ch][nt][2] == bv1) wi1 = max(wi1, c);
        if (acc[ch][nt][3] == bv1) wi1 = max(wi1, c + 1);
      }
    }
#pragma unroll
    for (int off = 1; off < 4; off <<= 1) {
      wi0 = max(wi0, __shfl_xor_sync(0xffffffffu, wi0, off));
      wi1 = max(wi1, __shfl_xor_sync(0xffffffffu, wi1, off));
    }
    if (wi0 == tile) dv0 = -NEG;
    if (wi1 == tile) dv1 = -NEG;
#pragma unroll
    for (int ch = 0; ch < NCH; ++ch) {
      if (ch * BS >= tile) continue;
#pragma unroll
      for (int nt = 0; nt < 8; ++nt) {
        const int c = ch * BS + nt * 8 + cb;
        if (c == wi0) acc[ch][nt][0] = -NEG;
        if (c + 1 == wi0) acc[ch][nt][1] = -NEG;
        if (c == wi1) acc[ch][nt][2] = -NEG;
        if (c + 1 == wi1) acc[ch][nt][3] = -NEG;
      }
    }
    if (gl == (s & 3)) {
      if (s < 4) {
        sl_a0 = wi0;
        sl_a1 = wi1;
      } else {
        sl_b0 = wi0;
        sl_b1 = wi1;
      }
    }
  }

  // ---- k_flash form: OR the 4 lanes of a row group into the query's block mask
  if (mskg) {
    unsigned long long b0 = (1ull << sl_a0) | (1ull << sl_b0);
    unsigned long long b1 = (1ull << sl_a1) | (1ull << sl_b1);
#pragma unroll
    for (int off = 1; off < 4; off <<= 1) {
      b0 |= __shfl_xor_sync(0xffffffffu, b0, off);
      b1 |= __shfl_xor_sync(0xffffffffu, b1, off);
    }
    if (gl == 0) {
      if (t_base + r0 < S) mskg[(size_t)bh * S + t_base + r0] = b0;
      if (t_base + r1 < S) mskg[(size_t)bh * S + t_base + r1] = b1;
    }
    return;
  }

  // ---- stage one entry per (row, slot); 4 per lane covers the warp's 16 rows
#define EMIT(RR, SLOT, BI)                                                    \
  do {                                                                        \
    const int t_ = t_base + (RR);                                             \
    if (t_ < S) {                                                             \
      const int bi_ = (BI), sl_ = (SLOT);                                     \
      if (bi_ * BS >= t_ - 63) { /* fully inside the sliding window */        \
        pml[((size_t)bh * RS + (t_ - T0)) * NSLOT + sl_] = make_float2(0.f, 0.f); \
      } else {                                                                \
        const int rk = atomicAdd(&lcnt[bi_], 1); /* shared, not global */      \
        sA[(RR) * 8 + sl_] = (bi_ << 10) | rk;                                \
        sB[(RR) * 8 + sl_] = (t_ << 4) | sl_;                                 \
      }                                                                       \
    }                                                                         \
  } while (0)

  EMIT(r0, gl, sl_a0);
  EMIT(r0, gl + 4, sl_b0);
  EMIT(r1, gl, sl_a1);
  EMIT(r1, gl + 4, sl_b1);
#undef EMIT

  // ---- one global atomic per (block, key block) instead of one per entry
  __syncthreads();
  for (int i = tid; i < ntile; i += NTHREAD)
    if (lcnt[i]) lbase[i] = atomicAdd(&cnt[bh * ntile + i], lcnt[i]);
  __syncthreads();
  for (int e = tid; e < TQ * 8; e += NTHREAD) {
    const int a = sA[e];
    if (a < 0) continue;
    const int bi = a >> 10, p = lbase[bi] + (a & 1023);
    if (p < cap) ent[(size_t)(bh * ntile + bi) * cap + p] = sB[e];
  }
}

// ------------------------------------------------------------ k_select_big
//
// Same contract as k_select, for ntile > 64*NCH: importances go through a shared
// (64 x ntile) matrix and each warp scans its 16 queries one at a time.  Only
// reachable for S > 8192, so it is kept simple rather than fast.

template <int D>
__global__ __launch_bounds__(NTHREAD) void k_select_big(
    const __nv_bfloat16* __restrict__ Qg, const __nv_bfloat16* __restrict__ Kg,
    const __half* __restrict__ kb_hi, const __half* __restrict__ kb_lo,
    int* __restrict__ ent, int* __restrict__ cnt, float2* __restrict__ pml, int S, int ntile,
    int tile0, int T0, int RS, int cap, float scale) {
  const int LD = D + 8;
  extern __shared__ char smem[];
  __half* sq = (__half*)smem;            // TQ x LD
  __half* sb0 = sq + TQ * LD;            // 64 x LD
  __half* sb1 = sb0 + BS * LD;           // 64 x LD
  float* imp = (float*)(sb1 + BS * LD);  // TQ x ntile

  const int tile = tile0 + blockIdx.x, bh = blockIdx.y;
  const int tid = threadIdx.x, lane = tid & 31, warp = tid >> 5;
  const int t_base = tile * BS;
  const int CPR = D / 8;  // 16B chunks per row

  // ---- load q tile (bf16 -> fp16)
  for (int idx = tid; idx < TQ * CPR; idx += NTHREAD) {
    int rr = idx / CPR, cc = (idx - rr * CPR) * 8;
    int t = t_base + rr;
    uint4 out = make_uint4(0, 0, 0, 0);
    if (t < S) cvt8(*(const uint4*)(Qg + (size_t)(bh * S + t) * D + cc), out);
    *(uint4*)(sq + rr * LD + cc) = out;
  }

  float acc[8][4];
  const int r0 = warp * 16 + FRAG_R0(lane), r1 = r0 + 8, cb = FRAG_C(lane);

  // ---- importance of the full blocks 0..tile-1  (kbar = hi + lo, two mma passes)
  for (int bc = 0; bc * BS < tile; ++bc) {
    __syncthreads();
    const int nbk = min(BS, tile - bc * BS);
    for (int idx = tid; idx < BS * CPR; idx += NTHREAD) {
      int rr = idx / CPR, cc = (idx - rr * CPR) * 8;
      uint4 z = make_uint4(0, 0, 0, 0);
      const bool ok = rr < nbk;
      *(uint4*)(sb0 + rr * LD + cc) =
          ok ? *(const uint4*)(kb_hi + (size_t)(bh * ntile + bc * BS + rr) * D + cc) : z;
      *(uint4*)(sb1 + rr * LD + cc) =
          ok ? *(const uint4*)(kb_lo + (size_t)(bh * ntile + bc * BS + rr) * D + cc) : z;
    }
    __syncthreads();
#pragma unroll
    for (int nt = 0; nt < 8; ++nt)
#pragma unroll
      for (int i = 0; i < 4; ++i) acc[nt][i] = 0.f;

    for (int pass = 0; pass < 2; ++pass) {
      const __half* sb = pass ? sb1 : sb0;
#pragma unroll 1
      for (int kt = 0; kt < D / 16; ++kt) {
        unsigned a0, a1, a2, a3;
        ldm_x4(sm_u32(sq + (warp * 16 + 8 * ((lane >> 3) & 1) + (lane & 7)) * LD + kt * 16 +
                      8 * (lane >> 4)),
               a0, a1, a2, a3);
#pragma unroll
        for (int nt = 0; nt < 8; ++nt) {
          unsigned b0, b1;
          ldm_x2(sm_u32(sb + (nt * 8 + (lane & 7)) * LD + kt * 16 + 8 * ((lane >> 3) & 1)), b0, b1);
          mma_f16(acc[nt][0], acc[nt][1], acc[nt][2], acc[nt][3], a0, a1, a2, a3, b0, b1);
        }
      }
    }
#pragma unroll
    for (int nt = 0; nt < 8; ++nt) {
      int c = bc * BS + nt * 8 + cb;
      float* p0 = imp + (size_t)r0 * ntile;
      float* p1 = imp + (size_t)r1 * ntile;
      // only entries < tile are ever read; skipping the rest also keeps the
      // stores inside the (TQ x ntile) tile when ntile is not a multiple of 64.
      if (c < tile) {
        p0[c] = acc[nt][0] * scale;
        p1[c] = acc[nt][2] * scale;
      }
      if (c + 1 < tile) {
        p0[c + 1] = acc[nt][1] * scale;
        p1[c + 1] = acc[nt][3] * scale;
      }
    }
  }

  // ---- diagonal block: importance = causal row-mean of its 64x64 score tile
  __syncthreads();
  for (int idx = tid; idx < BS * CPR; idx += NTHREAD) {
    int rr = idx / CPR, cc = (idx - rr * CPR) * 8;
    uint4 out = make_uint4(0, 0, 0, 0);
    if (t_base + rr < S) cvt8(*(const uint4*)(Kg + (size_t)(bh * S + t_base + rr) * D + cc), out);
    *(uint4*)(sb0 + rr * LD + cc) = out;
  }
  __syncthreads();
#pragma unroll
  for (int nt = 0; nt < 8; ++nt)
#pragma unroll
    for (int i = 0; i < 4; ++i) acc[nt][i] = 0.f;
#pragma unroll 1
  for (int kt = 0; kt < D / 16; ++kt) {
    unsigned a0, a1, a2, a3;
    ldm_x4(sm_u32(sq + (warp * 16 + 8 * ((lane >> 3) & 1) + (lane & 7)) * LD + kt * 16 +
                  8 * (lane >> 4)),
           a0, a1, a2, a3);
#pragma unroll
    for (int nt = 0; nt < 8; ++nt) {
      unsigned b0, b1;
      ldm_x2(sm_u32(sb0 + (nt * 8 + (lane & 7)) * LD + kt * 16 + 8 * ((lane >> 3) & 1)), b0, b1);
      mma_f16(acc[nt][0], acc[nt][1], acc[nt][2], acc[nt][3], a0, a1, a2, a3, b0, b1);
    }
  }
  {
    float s0 = 0.f, s1 = 0.f;
#pragma unroll
    for (int nt = 0; nt < 8; ++nt) {
      int c = nt * 8 + cb;
      if (c <= r0) s0 += acc[nt][0];
      if (c + 1 <= r0) s0 += acc[nt][1];
      if (c <= r1) s1 += acc[nt][2];
      if (c + 1 <= r1) s1 += acc[nt][3];
    }
#pragma unroll
    for (int off = 1; off < 4; off <<= 1) {
      s0 += __shfl_xor_sync(0xffffffffu, s0, off);
      s1 += __shfl_xor_sync(0xffffffffu, s1, off);
    }
    if ((lane & 3) == 0) {
      imp[(size_t)r0 * ntile + tile] = s0 * scale / (float)(r0 + 1);
      imp[(size_t)r1 * ntile + tile] = s1 * scale / (float)(r1 + 1);
    }
  }
  __syncthreads();  // diagonal mma reads of sb0 must finish before we reuse it

  // ---- scatter staging (overlaid on the kbar tiles, which are dead from here on)
  int* lcnt = (int*)sb0;          // ntile
  int* lbase = lcnt + ntile;      // ntile
  int* sA = lbase + ntile;        // TQ*8 : bi<<10 | rank-within-this-block
  int* sB = sA + TQ * 8;          // TQ*8 : (t<<4) | slot
  for (int i = tid; i < ntile; i += NTHREAD) lcnt[i] = 0;
  for (int i = tid; i < TQ * 8; i += NTHREAD) sA[i] = -1;
  __syncthreads();

  // ---- top-8 per query, one warp per 16 queries
  const int MAXC = MAXCAND;
  for (int q = warp * 16; q < warp * 16 + 16; ++q) {
    const int t = t_base + q;
    if (t >= S) continue;
    unsigned long long ky[MAXC];
    const float* ip = imp + (size_t)q * ntile;
#pragma unroll
    for (int c = 0; c < MAXC; ++c) {
      int bi = lane + 32 * c;
      ky[c] = (bi <= tile) ? ((((unsigned long long)f2u(ip[bi])) << 32) | (unsigned)bi) : 0ull;
    }
    // 8 rounds of 64-bit warp max-reduce; all array indices stay compile-time
    // constant (a dynamic index would push ky[] into local memory).
    int mysel = 0;
#pragma unroll 1
    for (int s = 0; s < 8; ++s) {
      unsigned long long best = 0ull;
#pragma unroll
      for (int c = 0; c < MAXC; ++c) best = ky[c] > best ? ky[c] : best;
#pragma unroll
      for (int off = 16; off; off >>= 1) {
        unsigned long long o = __shfl_xor_sync(0xffffffffu, best, off);
        best = o > best ? o : best;
      }
#pragma unroll
      for (int c = 0; c < MAXC; ++c)
        if (ky[c] == best) ky[c] = 0ull;  // keys are unique: at most one hit
      if (s == lane) mysel = (int)(unsigned)(best & 0xffffffffull);
    }
    const int w0 = t - 63;  // t >= 512 here, so w0 > 0
    if (lane < 8) {
      const int bi = mysel;
      if (bi * BS >= w0) {  // fully inside the sliding window -> no separate partial
        pml[((size_t)bh * RS + (t - T0)) * NSLOT + lane] = make_float2(0.f, 0.f);
      } else {
        int rank = atomicAdd(&lcnt[bi], 1);  // shared, not global
        sA[q * 8 + lane] = (bi << 10) | rank;
        sB[q * 8 + lane] = (t << 4) | lane;
      }
    }
  }

  // ---- one global atomic per (block, key block) instead of one per entry
  __syncthreads();
  for (int i = tid; i < ntile; i += NTHREAD)
    if (lcnt[i]) lbase[i] = atomicAdd(&cnt[bh * ntile + i], lcnt[i]);
  __syncthreads();
  for (int e = tid; e < TQ * 8; e += NTHREAD) {
    const int a = sA[e];
    if (a < 0) continue;
    const int bi = a >> 10, p = lbase[bi] + (a & 1023);
    if (p < cap) ent[(size_t)(bh * ntile + bi) * cap + p] = sB[e];
  }
}

// ---------------------------------------------------------------- k_attend

// role: 0 = dense causal tile (t < 512: every causal block is selected, so the tile is
//           plain causal attention -- flashed over all its key blocks and written
//           straight to O, no partials),
//       1 = sparse selected block (keys <= t-64), 2 = sliding window (t-63 <= keys <= t)
template <int D>
__global__ __launch_bounds__(NTHREAD) void k_attend(
    const __nv_bfloat16* __restrict__ Qg, const __nv_bfloat16* __restrict__ Kg,
    const __nv_bfloat16* __restrict__ Vg, __nv_bfloat16* __restrict__ O,
    const int* __restrict__ ent, const int* __restrict__ cnt, __half* __restrict__ po,
    float2* __restrict__ pml, int S, int ntile, int T0, int RS, int cap, int n_sp, int n_win,
    int win_t0, int n_den, int den_t0, int per_bh, int total_units, int* __restrict__ counter,
    float scale) {
  const int LD = D;  // unpadded: rows are XOR-swizzled (SWZ) instead
  extern __shared__ char smem[];
  __half* sq = (__half*)smem;   // TQ x D  (also output staging)
  __half* sk = sq + TQ * LD;    // 64 x D
  __half* sv = sk + BS * LD;    // 64 x D
  int* sqi = (int*)(sv + BS * LD);
  int* ssl = sqi + TQ;
  __shared__ int s_unit;

  const int tid = threadIdx.x, lane = tid & 31, warp = tid >> 5;
  const int CPR = D / 8;
  const int rr0 = warp * 16 + FRAG_R0(lane), rr1 = rr0 + 8, cb = FRAG_C(lane);
  const int NO = D / 8;
  // scores are carried in the exp2 domain: the mask folds log2(e) into the scale so
  // every softmax is a bare subtract.  m in the partials is log2 too -- merge_tile
  // and k_combine weight with exp2f(m - M), no LOG2E of their own.
  const float scaleL = scale * LOG2E;

  for (;;) {
    __syncthreads();
    if (tid == 0) s_unit = atomicAdd(counter, 1);
    __syncthreads();
    const int unit = s_unit;
    if (unit >= total_units) break;

    // Units are claimed in increasing order, so the claim order should be decreasing
    // in cost or the tail runs alone.  Cost falls off hard with the key-block index
    // (the first quartile of blocks takes ~13x the selections of the last), so keep r
    // ascending but interleave bh into the low bits: without this, the fat r=0 unit of
    // the last bh is claimed after every other bh is nearly done.
    const int nbh = total_units / per_bh;
    int r = unit / nbh, bh = unit - r * nbh;
    int role, kvb0 = 0, kvb1 = 0, nkv = 1, nchunk = 1, listc = 0, bi = 0, tile = 0, nqf = TQ;
    int ch0 = 0;
    if (r < n_sp) {  // ---- sparse: a slice of the query chunks of key block r / GSPLIT
      role = 1;
      bi = r / GSPLIT;
      const int g = r - bi * GSPLIT;
      kvb0 = bi;
      listc = cnt[bh * ntile + bi];
      if (listc > cap) listc = cap;
      if (listc == 0) continue;
      const int nch = (listc + TQ - 1) / TQ;
      const int gs = min(GSPLIT, max(1, nch / SMINCH));  // slices this unit is worth
      if (g >= gs) continue;                             // slice beyond the split: nothing to do
      const int per = (nch + gs - 1) / gs;
      ch0 = g * per;
      nchunk = min(nch, ch0 + per) - ch0;
      if (nchunk <= 0) continue;
    } else if (r - n_sp < n_den) {  // ---- dense causal tile: blocks 0..tile, one softmax
      role = 0;                       // longest tile first: nkv = tile + 1 key blocks
      tile = den_t0 + n_den - 1 - (r - n_sp);
      nkv = tile + 1;
      nqf = min(TQ, S - tile * BS);
    } else {  // ---- sliding window, two key blocks
      role = 2;
      tile = win_t0 + (r - n_sp - n_den);
      nkv = 2;
      kvb0 = tile - 1;
      kvb1 = tile;
      nqf = min(TQ, S - tile * BS);
    }

    // the sparse role keeps one key block across all its query chunks: load it once
    if (role == 1) {
      __syncthreads();
      const int kb = kvb0 * BS;
      for (int idx = tid; idx < BS * CPR; idx += NTHREAD) {
        int j = idx / CPR, cc = (idx - j * CPR) * 8;
        const int jj = min(kb + j, S - 1);
        cp16(sm_u32(sk + SWZ(j, cc)), Kg + (size_t)(bh * S + jj) * D + cc, kb + j < S);
        cp16(sm_u32(sv + SWZ(j, cc)), Vg + (size_t)(bh * S + jj) * D + cc, kb + j < S);
      }
      cp_commit();
    }

    // ent -> Q is a dependent global chain (indices, then the rows they point at) and
    // TQ <= NTHREAD, so one entry per thread lives in a register: fetch the next chunk's
    // indices right after issuing this chunk's Q gather, so the latency hides behind the
    // chunk's own mma work instead of stalling in front of it.
    const int* epb = ent + (size_t)(bh * ntile + bi) * cap + (size_t)ch0 * TQ;
    int e_pf = -1;
    if (role == 1 && tid < min(TQ, listc - ch0 * TQ)) e_pf = epb[tid];

    for (int ch = 0; ch < nchunk; ++ch) {
      __syncthreads();
      int nq;
      if (role == 1) {
        nq = min(TQ, listc - (ch0 + ch) * TQ);
        if (tid < TQ) {
          const int e = e_pf;
          sqi[tid] = (e >= 0) ? (e >> 4) : -1;
          ssl[tid] = (e >= 0) ? (e & 15) : 0;
        }
      } else {
        nq = nqf;
        for (int i = tid; i < TQ; i += NTHREAD) {
          sqi[i] = (i < nq) ? (tile * BS + i) : -1;
          ssl[i] = (role == 2) ? 8 : bi;
        }
      }
      __syncthreads();
      for (int idx = tid; idx < TQ * CPR; idx += NTHREAD) {
        int i = idx / CPR, cc = (idx - i * CPR) * 8;
        const int t = sqi[i];
        cp16(sm_u32(sq + SWZ(i, cc)), Qg + (size_t)(bh * S + max(t, 0)) * D + cc, t >= 0);
      }
      if (role == 1 && ch + 1 < nchunk)  // in flight across the whole chunk below
        e_pf = (tid < min(TQ, listc - (ch0 + ch + 1) * TQ)) ? epb[(ch + 1) * TQ + tid] : -1;
      cp_commit();
      cp_wait<0>();
      __syncthreads();

      const int t0 = sqi[rr0], t1 = sqi[rr1];
      unsigned qa[D / 16][4];  // A fragments live across the whole chunk now
#pragma unroll
      for (int kt = 0; kt < D / 16; ++kt)
        ldm_x4(sm_u32(sq + SWZ(warp * 16 + 8 * ((lane >> 3) & 1) + (lane & 7),
                               kt * 16 + 8 * (lane >> 4))),
               qa[kt][0], qa[kt][1], qa[kt][2], qa[kt][3]);
      float m0 = -NEG, m1 = -NEG, l0 = 0.f, l1 = 0.f;
      float o[16][4];
#pragma unroll
      for (int nt = 0; nt < NO; ++nt)
#pragma unroll
        for (int i = 0; i < 4; ++i) o[nt][i] = 0.f;

      for (int kvi = 0; kvi < nkv; ++kvi) {
        const int kbi = (role == 1) ? kvb0 : (role == 2 ? (kvi ? kvb1 : kvb0) : kvi);
        if (role != 1) {
          __syncthreads();
          const int kb = kbi * BS;
          for (int idx = tid; idx < BS * CPR; idx += NTHREAD) {
            int j = idx / CPR, cc = (idx - j * CPR) * 8;
            const int jj = min(kb + j, S - 1);
            cp16(sm_u32(sk + SWZ(j, cc)), Kg + (size_t)(bh * S + jj) * D + cc, kb + j < S);
            cp16(sm_u32(sv + SWZ(j, cc)), Vg + (size_t)(bh * S + jj) * D + cc, kb + j < S);
          }
          cp_commit();
          cp_wait<0>();
          __syncthreads();
        }
        const int kbase = kbi * BS;
        // per-row valid column window
        int lo0, hi0, lo1, hi1;
        if (role == 0) {
          lo0 = 0; hi0 = min(63, t0 - kbase);
          lo1 = 0; hi1 = min(63, t1 - kbase);
        } else if (role == 1) {
          lo0 = 0; hi0 = min(63, t0 - 64 - kbase);
          lo1 = 0; hi1 = min(63, t1 - 64 - kbase);
        } else {
          lo0 = max(0, t0 - 63 - kbase); hi0 = min(63, t0 - kbase);
          lo1 = max(0, t1 - 63 - kbase); hi1 = min(63, t1 - kbase);
        }
        if (t0 < 0) { lo0 = 1; hi0 = 0; }
        if (t1 < 0) { lo1 = 1; hi1 = 0; }

        // n-tile at a time: acc[nt] is final after its own D/16 mma, so the mask/max
        // chain of one tile issues into the mma shadow of the next
        float acc[8][4];
        float mt0 = -NEG, mt1 = -NEG;
#define QK_NT(MASK)                                                                          \
  _Pragma("unroll") for (int nt = 0; nt < 8; ++nt) {                                         \
    float a0 = 0.f, a1 = 0.f, a2 = 0.f, a3 = 0.f;                                            \
    _Pragma("unroll") for (int kt = 0; kt < D / 16; ++kt) {                                  \
      unsigned b0, b1;                                                                       \
      ldm_x2(sm_u32(sk + SWZ(nt * 8 + (lane & 7), kt * 16 + 8 * ((lane >> 3) & 1))), b0, b1); \
      mma_bf16(a0, a1, a2, a3, qa[kt][0], qa[kt][1], qa[kt][2], qa[kt][3], b0, b1);           \
    }                                                                                        \
    const int c = nt * 8 + cb;                                                               \
    MASK                                                                                     \
    acc[nt][0] = a0; acc[nt][1] = a1; acc[nt][2] = a2; acc[nt][3] = a3;                      \
    mt0 = fmaxf(mt0, fmaxf(a0, a1));                                                         \
    mt1 = fmaxf(mt1, fmaxf(a2, a3));                                                         \
  }
        // a selected block is fully below the window for nearly every query that picked
        // it, and a warp's 16 list entries are neighbours in t: one ballot, no compares
        if (__all_sync(0xffffffffu, lo0 == 0 && hi0 == 63 && lo1 == 0 && hi1 == 63)) {
          QK_NT(a0 *= scaleL; a1 *= scaleL; a2 *= scaleL; a3 *= scaleL;)
        } else {
          QK_NT(a0 = (c >= lo0 && c <= hi0) ? a0 * scaleL : -NEG;
                a1 = (c + 1 >= lo0 && c + 1 <= hi0) ? a1 * scaleL : -NEG;
                a2 = (c >= lo1 && c <= hi1) ? a2 * scaleL : -NEG;
                a3 = (c + 1 >= lo1 && c + 1 <= hi1) ? a3 * scaleL : -NEG;)
        }
#undef QK_NT
#pragma unroll
        for (int off = 1; off < 4; off <<= 1) {
          mt0 = fmaxf(mt0, __shfl_xor_sync(0xffffffffu, mt0, off));
          mt1 = fmaxf(mt1, __shfl_xor_sync(0xffffffffu, mt1, off));
        }
        const float nm0 = fmaxf(m0, mt0), nm1 = fmaxf(m1, mt1);
        const float rs0 = exp2f(m0 - nm0), rs1 = exp2f(m1 - nm1);
        float sl0 = 0.f, sl1 = 0.f;
#pragma unroll
        for (int nt = 0; nt < 8; ++nt) {
          acc[nt][0] = exp2f(acc[nt][0] - nm0);
          acc[nt][1] = exp2f(acc[nt][1] - nm0);
          acc[nt][2] = exp2f(acc[nt][2] - nm1);
          acc[nt][3] = exp2f(acc[nt][3] - nm1);
          sl0 += acc[nt][0] + acc[nt][1];
          sl1 += acc[nt][2] + acc[nt][3];
        }
#pragma unroll
        for (int off = 1; off < 4; off <<= 1) {
          sl0 += __shfl_xor_sync(0xffffffffu, sl0, off);
          sl1 += __shfl_xor_sync(0xffffffffu, sl1, off);
        }
        l0 = l0 * rs0 + sl0;
        l1 = l1 * rs1 + sl1;
        m0 = nm0;
        m1 = nm1;
#pragma unroll
        for (int nt = 0; nt < NO; ++nt) {
          o[nt][0] *= rs0; o[nt][1] *= rs0;
          o[nt][2] *= rs1; o[nt][3] *= rs1;
        }
        // P (fp16) @ V  (must stay unrolled: acc[] indices have to be compile-time)
#pragma unroll
        for (int kt = 0; kt < 4; ++kt) {
          const unsigned a0 = pack2bf(acc[2 * kt][0], acc[2 * kt][1]),
                         a1 = pack2bf(acc[2 * kt][2], acc[2 * kt][3]),
                         a2 = pack2bf(acc[2 * kt + 1][0], acc[2 * kt + 1][1]),
                         a3 = pack2bf(acc[2 * kt + 1][2], acc[2 * kt + 1][3]);
#pragma unroll
          for (int nt = 0; nt < NO; ++nt) {
            unsigned b0, b1;
            ldm_x2t(sm_u32(sv + SWZ(kt * 16 + 8 * ((lane >> 3) & 1) + (lane & 7), nt * 8)), b0,
                    b1);
            mma_bf16(o[nt][0], o[nt][1], o[nt][2], o[nt][3], a0, a1, a2, a3, b0, b1);
          }
        }
      }

      // ---- normalise, stage in shared, coalesced store
      const float i0 = (l0 > 0.f) ? 1.f / l0 : 0.f, i1 = (l1 > 0.f) ? 1.f / l1 : 0.f;
      __syncthreads();
      if (role == 0) {  // complete rows: bf16 straight to O, nothing for k_combine
#pragma unroll
        for (int nt = 0; nt < NO; ++nt) {
          int c = nt * 8 + cb;
          *(__nv_bfloat162*)(sq + SWZ(rr0, c)) =
              __floats2bfloat162_rn(o[nt][0] * i0, o[nt][1] * i0);
          *(__nv_bfloat162*)(sq + SWZ(rr1, c)) =
              __floats2bfloat162_rn(o[nt][2] * i1, o[nt][3] * i1);
        }
        __syncthreads();
        for (int idx = tid; idx < TQ * CPR; idx += NTHREAD) {
          int i = idx / CPR, cc = (idx - i * CPR) * 8;
          int t = sqi[i];
          if (t < 0) continue;
          *(uint4*)(O + (size_t)(bh * S + t) * D + cc) = *(const uint4*)(sq + SWZ(i, cc));
        }
        continue;
      }
#pragma unroll
      for (int nt = 0; nt < NO; ++nt) {
        int c = nt * 8 + cb;
        *(__half2*)(sq + SWZ(rr0, c)) = __floats2half2_rn(o[nt][0] * i0, o[nt][1] * i0);
        *(__half2*)(sq + SWZ(rr1, c)) = __floats2half2_rn(o[nt][2] * i1, o[nt][3] * i1);
      }
      if ((lane & 3) == 0) {
        if (t0 >= 0) pml[((size_t)bh * RS + (t0 - T0)) * NSLOT + ssl[rr0]] = make_float2(m0, l0);
        if (t1 >= 0) pml[((size_t)bh * RS + (t1 - T0)) * NSLOT + ssl[rr1]] = make_float2(m1, l1);
      }
      __syncthreads();
      for (int idx = tid; idx < TQ * CPR; idx += NTHREAD) {
        int i = idx / CPR, cc = (idx - i * CPR) * 8;
        int t = sqi[i];
        if (t < 0) continue;
        *(uint4*)(po + (((size_t)bh * RS + (t - T0)) * NSLOT + ssl[i]) * D + cc) =
            *(const uint4*)(sq + SWZ(i, cc));
      }
    }
  }
}

// Merge the NS partials of one query tile into O.  Run by the last unit of the tile to
// arrive (see the ticket in k_flash), which is why NS is exact here and the slots are
// all live: no separate combine pass, and the partials are still L2 hot.  po/pml are
// read through .cg so the merging block cannot hit a stale L1 line of a neighbour's
// partial.
template <int D, int NS>
DEVI void merge_tile(const __half* __restrict__ po, const float2* __restrict__ pml,
                     __nv_bfloat16* __restrict__ O, size_t r_base, int nq, int nsf, int tid) {
  const int CPR = D / 8;
  for (int idx = tid; idx < TQ * CPR; idx += NTHREAD) {
    const int i = idx / CPR, cc = (idx - i * CPR) * 8;
    if (i >= nq) continue;
    const size_t row = r_base + i;
    const float2* ml = pml + row * nsf;
    float mv[NS], lv[NS], M = -NEG;
#pragma unroll
    for (int s = 0; s < NS; ++s) {
      const float2 x = __ldcg(ml + s);
      mv[s] = x.x;
      lv[s] = x.y;
      if (x.y > 0.f && x.x > M) M = x.x;
    }
    float den = 0.f, w[NS];
#pragma unroll
    for (int s = 0; s < NS; ++s) {
      w[s] = (lv[s] > 0.f) ? lv[s] * exp2f(mv[s] - M) : 0.f;
      den += w[s];
    }
    const float inv = (den > 0.f) ? 1.f / den : 0.f;
    float a[8];
#pragma unroll
    for (int e = 0; e < 8; ++e) a[e] = 0.f;
#pragma unroll
    for (int s = 0; s < NS; ++s) {
      const uint4 pk = __ldcg((const uint4*)(po + (row * nsf + s) * D + cc));
      const __half2* hp = (const __half2*)&pk;
      const float ws = w[s];
#pragma unroll
      for (int e = 0; e < 4; ++e) {
        const float2 x = __half22float2(hp[e]);
        a[2 * e] += ws * x.x;
        a[2 * e + 1] += ws * x.y;
      }
    }
    __nv_bfloat162* op = (__nv_bfloat162*)(O + row * D + cc);
#pragma unroll
    for (int e = 0; e < 4; ++e) op[e] = __floats2bfloat162_rn(a[2 * e] * inv, a[2 * e + 1] * inv);
  }
}

// ---------------------------------------------------------------- k_flash
//
// Masked-dense path, used when ntile <= 64 so a query's selection fits in one 64-bit
// mask.  A unit is (bh, a flat range of (query tile, key block) pairs); it walks every
// causal key block of its range and keeps, per row, only the columns that row's mask
// (union its sliding window) allows.  Compared to the inverted path this reads Q once
// per unit instead of once per selected block and emits at most one partial per query
// instead of 9 -- both were measured to be the dominant cost at small S.  The mma
// spent on rows that masked a block out is cheap in comparison while ntile is small.
//
// Work is split by an equal-size flat range: the (query tile, key block) pairs of one
// (b, h) form a triangle of L = ntile(ntile+1)/2 tiles, the whole grid covers BH * L of
// them, and unit u takes flat range [u*W, u*W+W).  W is picked so that exactly one unit
// fits per resident block, which makes every block do the same number of mma tiles --
// with the natural (tile, chunk) split the longest block ran 1.8-2.3x the mean.  A unit
// walks the 1-3 query tiles its range touches; a tile it covers completely is written
// straight to O, a tile it shares with a neighbour unit becomes a partial in slot
// (u - first unit of that tile) for k_combine.
//
// K/V come in by cp.async as raw bf16 (no staging registers, no fp16 conversion) and
// the mma runs on bf16 directly: K is fetched one key block ahead into a second
// buffer and V overlaps the QK mma of its own block, so the only exposed global
// latency per unit is its first key block.  Q stays in registers as A-fragments.

template <int D>
// shared memory, not registers, caps this at 3 blocks/SM (1 for D=128): tell ptxas so
// it stops squeezing to 128 registers and spilling.
__global__ __launch_bounds__(NTHREAD, D == 64 ? 3 : 1) void k_flash(
    const __nv_bfloat16* __restrict__ Qg, const __nv_bfloat16* __restrict__ Kg,
    const __nv_bfloat16* __restrict__ Vg, __nv_bfloat16* __restrict__ O,
    const unsigned long long* __restrict__ mskg, __half* __restrict__ po,
    float2* __restrict__ pml, int* __restrict__ tick, int S, int ntile, int nsf, int W, int L,
    int total_tiles, float scale) {
  const int LD = D + 8;
  extern __shared__ char smem[];
  __nv_bfloat16* sK = (__nv_bfloat16*)smem;  // 2 x BS x LD (double buffered)
  __nv_bfloat16* sV = sK + 2 * BS * LD;      // BS x LD, also Q input / O output staging
  unsigned long long* smsk = (unsigned long long*)(sV + BS * LD);  // TQ
  __shared__ int s_last;  // this unit closed a shared tile -> it merges the partials

  const int tid = threadIdx.x, lane = tid & 31, warp = tid >> 5;
  const int CPR = D / 8;  // 16B chunks per row
  const int lrow = tid / CPR, lstep = NTHREAD / CPR, lcol = (tid % CPR) * 8;
  const int rr0 = warp * 16 + FRAG_R0(lane), rr1 = rr0 + 8, cb = FRAG_C(lane);
  const int NO = D / 8;
  // scores are carried in the exp2 domain: the mask folds log2(e) into the scale so
  // every softmax is a bare subtract.  m in the partials is log2 too -- merge_tile
  // and k_combine weight with exp2f(m - M), no LOG2E of their own.
  const float scaleL = scale * LOG2E;
  const int arow = warp * 16 + 8 * ((lane >> 3) & 1) + (lane & 7), acol = 8 * (lane >> 4);
  const int brow = (lane & 7), bcol = 8 * ((lane >> 3) & 1);
  const int bhi = (lane >> 4) & 1, brow4 = bhi * 8 + brow;  // x4: lanes 16-31 -> next n-tile

  const int u = blockIdx.x, Fend = min(total_tiles, u * W + W);
  for (int F = u * W; F < Fend;) {
    const int bh = F / L, f = F - bh * L;
    int tq = (int)((sqrtf(8.f * f + 1.f) - 1.f) * 0.5f);  // largest tq with T(tq) <= f
    while ((tq + 1) * (tq + 2) / 2 <= f) ++tq;
    while (tq > 0 && tq * (tq + 1) / 2 > f) --tq;
    const int T = tq * (tq + 1) / 2;
    const int b_lo = f - T, b_hi = min(tq, Fend - bh * L - T - 1);
    const int Fb = bh * L + T;                   // flat index of the tile's first key block
    const int kc = u - Fb / W;                   // this unit's partial slot for the tile
    const bool whole = (b_lo == 0) && (b_hi == tq);
    const int t_base = tq * BS, nq = min(TQ, S - t_base);
    F = bh * L + T + b_hi + 1;

    // ---- Q tile -> A fragments (staged through sV, which V overwrites afterwards)
    for (int r = lrow; r < TQ; r += lstep)
      cp16(sm_u32(sV + r * LD + lcol), Qg + (size_t)(bh * S + t_base + r) * D + lcol, r < nq);
    for (int i = tid; i < TQ; i += NTHREAD)  // tiles < 8 have every causal block selected
      smsk[i] = (i >= nq) ? 0ull : (tq < 8 ? ~0ull : mskg[(size_t)bh * S + t_base + i]);
    cp_commit();
    cp_wait<0>();
    __syncthreads();
    unsigned qa[D / 16][4];
#pragma unroll
    for (int kt = 0; kt < D / 16; ++kt)
      ldm_x4(sm_u32(sV + arow * LD + kt * 16 + acol), qa[kt][0], qa[kt][1], qa[kt][2], qa[kt][3]);

    const int t0 = t_base + rr0, t1 = t_base + rr1;
    const unsigned long long mk0 = smsk[rr0], mk1 = smsk[rr1];
    float m0 = -NEG, m1 = -NEG, l0 = 0.f, l1 = 0.f;
    float o[16][4];
#pragma unroll
    for (int nt = 0; nt < NO; ++nt)
#pragma unroll
      for (int i = 0; i < 4; ++i) o[nt][i] = 0.f;

    __syncthreads();  // sV free
    for (int r = lrow; r < BS; r += lstep)
      cp16(sm_u32(sK + r * LD + lcol), Kg + (size_t)(bh * S + b_lo * BS + r) * D + lcol,
           b_lo * BS + r < S);
    cp_commit();

    for (int b = b_lo; b <= b_hi; ++b) {
      const int kb = b * BS;
      __nv_bfloat16* sKp = sK + ((b - b_lo) & 1) * (BS * LD);
      for (int r = lrow; r < BS; r += lstep)
        cp16(sm_u32(sV + r * LD + lcol), Vg + (size_t)(bh * S + kb + r) * D + lcol, kb + r < S);
      cp_commit();
      if (b < b_hi) {
        __nv_bfloat16* sKn = sK + ((b - b_lo + 1) & 1) * (BS * LD);
        for (int r = lrow; r < BS; r += lstep)
          cp16(sm_u32(sKn + r * LD + lcol), Kg + (size_t)(bh * S + kb + BS + r) * D + lcol,
               kb + BS + r < S);
      }
      cp_commit();
      cp_wait<2>();  // K of this block has landed
      __syncthreads();

      // n-tile pairs, not the whole 64x64 tile, are the unit here: acc[nt] is final
      // after its own D/16 mma, so the scale/mask/max chain of one pair issues into
      // the mma shadow of the next instead of after all 32 of them.  Two independent
      // accumulators per pair keep the tensor pipe fed despite the chaining.
      // selected -> the whole causal part of the block; otherwise only the sliding
      // window part, which is empty by construction for blocks below tile-1.
      float acc[8][4];
      float mt0 = -NEG, mt1 = -NEG;
      const bool cut = (b + 1 >= tq);  // block-uniform: can causality/window split it?
      const float bs0 = ((mk0 >> b) & 1) ? 0.f : -NEG;
      const float bs1 = ((mk1 >> b) & 1) ? 0.f : -NEG;
      const int lo0 = ((mk0 >> b) & 1) ? 0 : max(0, t0 - 63 - kb), hi0 = min(63, t0 - kb);
      const int lo1 = ((mk1 >> b) & 1) ? 0 : max(0, t1 - 63 - kb), hi1 = min(63, t1 - kb);
#define QK_PAIR(MASK)                                                                       \
  _Pragma("unroll") for (int nt = 0; nt < 8; nt += 2) {                                     \
    float a0 = 0.f, a1 = 0.f, a2 = 0.f, a3 = 0.f, a4 = 0.f, a5 = 0.f, a6 = 0.f, a7 = 0.f;   \
    _Pragma("unroll") for (int kt = 0; kt < D / 16; ++kt) { /* qa[] must stay in regs */    \
      unsigned c0, c1, c2, c3;                                                              \
      ldm_x4(sm_u32(sKp + (nt * 8 + brow4) * LD + kt * 16 + bcol), c0, c1, c2, c3);          \
      mma_bf16(a0, a1, a2, a3, qa[kt][0], qa[kt][1], qa[kt][2], qa[kt][3], c0, c1);          \
      mma_bf16(a4, a5, a6, a7, qa[kt][0], qa[kt][1], qa[kt][2], qa[kt][3], c2, c3);          \
    }                                                                                       \
    const int c = nt * 8 + cb;                                                              \
    MASK                                                                                    \
    acc[nt][0] = a0; acc[nt][1] = a1; acc[nt][2] = a2; acc[nt][3] = a3;                      \
    acc[nt + 1][0] = a4; acc[nt + 1][1] = a5; acc[nt + 1][2] = a6; acc[nt + 1][3] = a7;      \
    mt0 = fmaxf(mt0, fmaxf(fmaxf(a0, a1), fmaxf(a4, a5)));                                   \
    mt1 = fmaxf(mt1, fmaxf(fmaxf(a2, a3), fmaxf(a6, a7)));                                   \
  }
      if (!cut) {
        // one fma per element: every column is in iff the row selected the block
        QK_PAIR(a0 = fmaf(a0, scaleL, bs0); a1 = fmaf(a1, scaleL, bs0);
                a2 = fmaf(a2, scaleL, bs1); a3 = fmaf(a3, scaleL, bs1);
                a4 = fmaf(a4, scaleL, bs0); a5 = fmaf(a5, scaleL, bs0);
                a6 = fmaf(a6, scaleL, bs1); a7 = fmaf(a7, scaleL, bs1);)
      } else {
        QK_PAIR(a0 = (c >= lo0 && c <= hi0) ? a0 * scaleL : -NEG;
                a1 = (c + 1 >= lo0 && c + 1 <= hi0) ? a1 * scaleL : -NEG;
                a2 = (c >= lo1 && c <= hi1) ? a2 * scaleL : -NEG;
                a3 = (c + 1 >= lo1 && c + 1 <= hi1) ? a3 * scaleL : -NEG;
                a4 = (c + 8 >= lo0 && c + 8 <= hi0) ? a4 * scaleL : -NEG;
                a5 = (c + 9 >= lo0 && c + 9 <= hi0) ? a5 * scaleL : -NEG;
                a6 = (c + 8 >= lo1 && c + 8 <= hi1) ? a6 * scaleL : -NEG;
                a7 = (c + 9 >= lo1 && c + 9 <= hi1) ? a7 * scaleL : -NEG;)
      }
#undef QK_PAIR
#pragma unroll
      for (int off = 1; off < 4; off <<= 1) {
        mt0 = fmaxf(mt0, __shfl_xor_sync(0xffffffffu, mt0, off));
        mt1 = fmaxf(mt1, __shfl_xor_sync(0xffffffffu, mt1, off));
      }
      const float nm0 = fmaxf(m0, mt0), nm1 = fmaxf(m1, mt1);
      const float rs0 = exp2f(m0 - nm0), rs1 = exp2f(m1 - nm1);
      float sl0 = 0.f, sl1 = 0.f;
#pragma unroll
      for (int nt = 0; nt < 8; ++nt) {
        acc[nt][0] = exp2f(acc[nt][0] - nm0);
        acc[nt][1] = exp2f(acc[nt][1] - nm0);
        acc[nt][2] = exp2f(acc[nt][2] - nm1);
        acc[nt][3] = exp2f(acc[nt][3] - nm1);
        sl0 += acc[nt][0] + acc[nt][1];
        sl1 += acc[nt][2] + acc[nt][3];
      }
#pragma unroll
      for (int off = 1; off < 4; off <<= 1) {
        sl0 += __shfl_xor_sync(0xffffffffu, sl0, off);
        sl1 += __shfl_xor_sync(0xffffffffu, sl1, off);
      }
      l0 = l0 * rs0 + sl0;
      l1 = l1 * rs1 + sl1;
      m0 = nm0;
      m1 = nm1;
#pragma unroll
      for (int nt = 0; nt < NO; ++nt) {
        o[nt][0] *= rs0;
        o[nt][1] *= rs0;
        o[nt][2] *= rs1;
        o[nt][3] *= rs1;
      }
      cp_wait<1>();  // V of this block has landed
      __syncthreads();
#pragma unroll
      for (int kt = 0; kt < 4; ++kt) {
        const unsigned a0 = pack2bf(acc[2 * kt][0], acc[2 * kt][1]),
                       a1 = pack2bf(acc[2 * kt][2], acc[2 * kt][3]),
                       a2 = pack2bf(acc[2 * kt + 1][0], acc[2 * kt + 1][1]),
                       a3 = pack2bf(acc[2 * kt + 1][2], acc[2 * kt + 1][3]);
#pragma unroll
        for (int nt = 0; nt < NO; nt += 2) {
          unsigned b0, b1, b2, b3;
          ldm_x4t(sm_u32(sV + (kt * 16 + bcol + brow) * LD + (nt + bhi) * 8), b0, b1, b2, b3);
          mma_bf16(o[nt][0], o[nt][1], o[nt][2], o[nt][3], a0, a1, a2, a3, b0, b1);
          mma_bf16(o[nt + 1][0], o[nt + 1][1], o[nt + 1][2], o[nt + 1][3], a0, a1, a2, a3, b2, b3);
        }
      }
      __syncthreads();  // sV reusable by the next block
    }

    const float i0 = (l0 > 0.f) ? 1.f / l0 : 0.f, i1 = (l1 > 0.f) ? 1.f / l1 : 0.f;
    __nv_bfloat16* so = sK;  // K is dead, reuse it to coalesce the stores
#pragma unroll
    for (int nt = 0; nt < NO; ++nt) {
      const int c = nt * 8 + cb;
      if (whole) {
        *(__nv_bfloat162*)(so + rr0 * LD + c) = __floats2bfloat162_rn(o[nt][0] * i0, o[nt][1] * i0);
        *(__nv_bfloat162*)(so + rr1 * LD + c) = __floats2bfloat162_rn(o[nt][2] * i1, o[nt][3] * i1);
      } else {
        *(__half2*)(so + rr0 * LD + c) = __floats2half2_rn(o[nt][0] * i0, o[nt][1] * i0);
        *(__half2*)(so + rr1 * LD + c) = __floats2half2_rn(o[nt][2] * i1, o[nt][3] * i1);
      }
    }
    if (!whole && (lane & 3) == 0) {
      if (rr0 < nq) pml[((size_t)(bh * S + t0)) * nsf + kc] = make_float2(m0, l0);
      if (rr1 < nq) pml[((size_t)(bh * S + t1)) * nsf + kc] = make_float2(m1, l1);
    }
    __syncthreads();
    for (int idx = tid; idx < TQ * CPR; idx += NTHREAD) {
      const int i = idx / CPR, cc = (idx - i * CPR) * 8;
      if (i >= nq) continue;
      const uint4 val = *(const uint4*)(so + i * LD + cc);
      if (whole)
        *(uint4*)(O + (size_t)(bh * S + t_base + i) * D + cc) = val;
      else
        *(uint4*)(po + ((size_t)(bh * S + t_base + i) * nsf + kc) * D + cc) = val;
    }

    // ---- shared tile: the unit that arrives last merges the whole tile in place.  The
    // ticket is reset by that same unit, so the buffer is self-cleaning across launches
    // and no block ever spins: every unit either merges or leaves.
    if (!whole) {
      const int ns = (Fb + tq) / W - Fb / W + 1;
      __threadfence();
      __syncthreads();
      if (tid == 0) {
        int* tp = tick + bh * ntile + tq;
        const int old = atomicAdd(tp, 1);
        s_last = (old == ns - 1);
        if (s_last) *tp = 0;
      }
      __syncthreads();
      if (s_last) {
        const size_t r_base = (size_t)bh * S + t_base;
        if (ns == 2) merge_tile<D, 2>(po, pml, O, r_base, nq, nsf, tid);
        else if (ns == 3) merge_tile<D, 3>(po, pml, O, r_base, nq, nsf, tid);
        else merge_tile<D, NSF>(po, pml, O, r_base, nq, nsf, tid);
      }
    }
  }
}

// ---------------------------------------------------------------- k_combine

template <int D, int NSM>
__global__ void k_combine(const __half* __restrict__ po, const float2* __restrict__ pml,
                          __nv_bfloat16* __restrict__ O, int S, int T0, int nq_range, int RS,
                          int nq_total, int* __restrict__ counter, int nslot) {
  // One row per LPR-lane group, LPR = D/8, so every lane's slice of a partial is exactly
  // 16 B: a warp covers 32/LPR rows and keeps 8x the bytes in flight of a half2-per-lane
  // split.  This kernel is nothing but L2 bandwidth, so that is the whole game.
  constexpr int LPR = D / 8, EPL = 4;  // lanes per row; half2 per lane
  if (blockIdx.x == 0 && threadIdx.x == 0) *counter = 0;
  const int g = blockIdx.x * (NTHREAD / LPR) + (int)(threadIdx.x / LPR);
  if (g >= nq_total) return;
  const int gl = threadIdx.x & (LPR - 1);
  const int bh = g / nq_range, rq = g - bh * nq_range;
  const int t = T0 + rq;
  if (t < 8 * BS) return;  // dense tiles are written straight to O by k_attend
  // NSM is a compile-time bound so that mv/lv/w stay in registers: with a runtime bound
  // ptxas spills them to local memory and the whole kernel becomes latency bound.
  const float2* ml = pml + ((size_t)bh * RS + rq) * nslot;  // broadcast within the group
  float mv[NSM], lv[NSM], M = -NEG;
#pragma unroll
  for (int s = 0; s < NSM; ++s) {
    float2 x = ml[s];
    mv[s] = x.x;
    lv[s] = x.y;
    if (x.y > 0.f && x.x > M) M = x.x;
  }
  float den = 0.f, w[NSM];
#pragma unroll
  for (int s = 0; s < NSM; ++s) {
    w[s] = (lv[s] > 0.f) ? lv[s] * exp2f(mv[s] - M) : 0.f;
    den += w[s];
  }
  const float inv = (den > 0.f) ? 1.f / den : 0.f;
  float2 acc[EPL];
#pragma unroll
  for (int e = 0; e < EPL; ++e) acc[e] = make_float2(0.f, 0.f);
  const uint4* p = (const uint4*)(po + ((size_t)bh * RS + rq) * nslot * D) + gl;
  // every slot is read unconditionally (an unwritten slot has w == 0, and the select keeps
  // its garbage out of acc) so that all NSM loads are in flight at once
  uint4 u[NSM];
#pragma unroll
  for (int s = 0; s < NSM; ++s) u[s] = p[s * LPR];
#pragma unroll
  for (int s = 0; s < NSM; ++s) {
    const float ws = w[s];
    const __half2* y = (const __half2*)&u[s];
#pragma unroll
    for (int e = 0; e < EPL; ++e) {
      const float2 x = __half22float2(y[e]);
      acc[e].x += (ws > 0.f) ? ws * x.x : 0.f;
      acc[e].y += (ws > 0.f) ? ws * x.y : 0.f;
    }
  }
  __nv_bfloat162 ob[EPL];
#pragma unroll
  for (int e = 0; e < EPL; ++e) ob[e] = __floats2bfloat162_rn(acc[e].x * inv, acc[e].y * inv);
  *((uint4*)(O + (size_t)(bh * S + t) * D) + gl) = *(const uint4*)ob;
}

// ---------------------------------------------------------------- host launcher

template <int D>
static void run_d(const at::Tensor& q, const at::Tensor& k, const at::Tensor& v, at::Tensor& out,
                  at::Tensor& kb_hi, at::Tensor& kb_lo, at::Tensor& ent, at::Tensor& cnt,
                  at::Tensor& po, at::Tensor& pml, at::Tensor& msk, at::Tensor& ctr, int RS) {
  const int B = q.size(0), H = q.size(1), S = q.size(2);
  const int BH = B * H, ntile = (S + BS - 1) / BS;
  const int nrange = (S + RS - 1) / RS;
  const float scale = 1.0f / std::sqrt((float)D);
  auto stream = at::cuda::getCurrentCUDAStream();

  const int LD = D + 8;
  const bool selbig = ntile > 64 * SELNCH;
  // one candidate chunk covers 64 key blocks: half the accumulator registers, two more
  // resident blocks per SM and half the per-round scan of the top-8 loop.  Above that,
  // the second chunk's registers are only live for tiles >= 64, so while those are a
  // small minority of the launch (< ~1/4) it is still worth spilling for the occupancy.
  const bool sel1 = ntile <= 64, sel2 = ntile <= 82;
  // k_attend is swizzled, so its tile rows are unpadded (D, not LD): 2 blocks/SM at D=128
  const int sh_att = (int)((size_t)(TQ + 2 * BS) * D * sizeof(__half) + 2 * TQ * sizeof(int));
  const int sh_sel =
      selbig ? (int)((size_t)(TQ + 2 * BS) * LD * sizeof(__half) + (size_t)TQ * ntile * 4)
             : (int)std::max((size_t)(TQ + BS) * D * sizeof(__half),
                             (size_t)(2 * ntile + 2 * TQ * 8) * sizeof(int));
  static int set_att = 0, set_sel = 0, set_sel1 = 0, set_sel3 = 0, set_selbig = 0;
  if (ntile > 8 && sel1 && sh_sel > set_sel1) {
    C10_CUDA_CHECK(cudaFuncSetAttribute(k_select<D, 1, SELMB1>,
                                        cudaFuncAttributeMaxDynamicSharedMemorySize, sh_sel));
    set_sel1 = sh_sel;
  }
  if (sh_att > set_att) {
    C10_CUDA_CHECK(
        cudaFuncSetAttribute(k_attend<D>, cudaFuncAttributeMaxDynamicSharedMemorySize, sh_att));
    set_att = sh_att;
  }
  if (ntile > 8 && !selbig && !sel1 && sel2 && sh_sel > set_sel) {
    C10_CUDA_CHECK(cudaFuncSetAttribute(k_select<D, SELNCH, SELMB2>,
                                        cudaFuncAttributeMaxDynamicSharedMemorySize, sh_sel));
    set_sel = sh_sel;
  }
  if (ntile > 8 && !selbig && !sel2 && sh_sel > set_sel3) {
    C10_CUDA_CHECK(cudaFuncSetAttribute(k_select<D, SELNCH, SELMB3>,
                                        cudaFuncAttributeMaxDynamicSharedMemorySize, sh_sel));
    set_sel3 = sh_sel;
  }
  if (ntile > 8 && selbig && sh_sel > set_selbig) {
    C10_CUDA_CHECK(
        cudaFuncSetAttribute(k_select_big<D>, cudaFuncAttributeMaxDynamicSharedMemorySize, sh_sel));
    set_selbig = sh_sel;
  }
  TORCH_CHECK(ntile <= 32 * MAXCAND, "nsa: sequence too long for the selection kernel");

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

  // ---- masked-dense path: selection fits in a u64, so sweep whole causal rows
  if (ntile <= FLASH_MAXTILE) {
    const int nsf = NSF;                              // partial slots per query
    const int L = ntile * (ntile + 1) / 2, total = L * BH;  // mma tiles per bh / in total
    const int sh_fl = (int)((size_t)(3 * BS) * LD * sizeof(__nv_bfloat16) +
                            (size_t)TQ * sizeof(unsigned long long));
    static int set_fl = 0;
    if (sh_fl > set_fl) {
      C10_CUDA_CHECK(
          cudaFuncSetAttribute(k_flash<D>, cudaFuncAttributeMaxDynamicSharedMemorySize, sh_fl));
      set_fl = sh_fl;
    }
    // One unit per resident block, all units the same size: perfect load balance.
    static int resident = 0;
    if (!resident) {
      int nb = 0;
      C10_CUDA_CHECK(cudaOccupancyMaxActiveBlocksPerMultiprocessor(&nb, k_flash<D>, NTHREAD,
                                                                   (size_t)sh_fl));
      resident = std::max(1, nb) * at::cuda::getCurrentDeviceProperties()->multiProcessorCount;
    }
    const int W = std::max((total + resident - 1) / resident, (ntile + nsf - 2) / (nsf - 1));
    const int nunits = (total + W - 1) / W;
    if (ntile > 8) {
      // ntile <= 64 here, so one chunk of candidates (the attribute is set above).
      k_kbar<D><<<dim3(ntile, BH), NTHREAD, 0, stream>>>(kp, (__half*)kb_hi.data_ptr(),
                                                        (__half*)kb_lo.data_ptr(), S, ntile,
                                                        (int*)cnt.data_ptr(), 0);
      k_select<D, 1, SELMB1><<<dim3(ntile - 8, BH), NTHREAD, sh_sel, stream>>>(
          qp, kp, (const __half*)kb_hi.data_ptr(), (const __half*)kb_lo.data_ptr(), nullptr,
          nullptr, nullptr, (unsigned long long*)msk.data_ptr(), S, ntile, 8, 0, S, 0, scale);
    }
    k_flash<D><<<nunits, NTHREAD, sh_fl, stream>>>(
        qp, kp, vp, (__nv_bfloat16*)out.data_ptr(), (const unsigned long long*)msk.data_ptr(),
        (__half*)po.data_ptr(), (float2*)pml.data_ptr(), (int*)ctr.data_ptr(), S, ntile, nsf, W, L,
        total, scale);
    return;
  }

  if (ntile > 8) {
    k_kbar<D><<<dim3(ntile, BH), NTHREAD, 0, stream>>>(
        kp, (__half*)kb_hi.data_ptr(), (__half*)kb_lo.data_ptr(), S, ntile,
        (int*)cnt.data_ptr(), nrange);
  }

  const int cap = RS;
  for (int rg = 0; rg < nrange; ++rg) {
    const int T0 = rg * RS, T1 = std::min(S, T0 + RS);
    const int tl0 = T0 / BS, tl1 = (T1 + BS - 1) / BS;  // tiles [tl0, tl1)
    int* cnt_r = (int*)cnt.data_ptr() + (size_t)rg * BH * ntile;

    const int sel_t0 = std::max(tl0, 8), sel_n = std::max(0, tl1 - sel_t0);
    if (sel_n > 0) {
      if (selbig) {
        k_select_big<D><<<dim3(sel_n, BH), NTHREAD, sh_sel, stream>>>(
            qp, kp, (const __half*)kb_hi.data_ptr(), (const __half*)kb_lo.data_ptr(),
            (int*)ent.data_ptr(), cnt_r, (float2*)pml.data_ptr(), S, ntile, sel_t0, T0, RS, cap,
            scale);
      } else if (sel1) {
        k_select<D, 1, SELMB1><<<dim3(sel_n, BH), NTHREAD, sh_sel, stream>>>(
            qp, kp, (const __half*)kb_hi.data_ptr(), (const __half*)kb_lo.data_ptr(),
            (int*)ent.data_ptr(), cnt_r, (float2*)pml.data_ptr(), nullptr, S, ntile, sel_t0, T0, RS,
            cap, scale);
      } else if (sel2) {
        k_select<D, SELNCH, SELMB2><<<dim3(sel_n, BH), NTHREAD, sh_sel, stream>>>(
            qp, kp, (const __half*)kb_hi.data_ptr(), (const __half*)kb_lo.data_ptr(),
            (int*)ent.data_ptr(), cnt_r, (float2*)pml.data_ptr(), nullptr, S, ntile, sel_t0, T0, RS,
            cap, scale);
      } else {
        k_select<D, SELNCH, SELMB3><<<dim3(sel_n, BH), NTHREAD, sh_sel, stream>>>(
            qp, kp, (const __half*)kb_hi.data_ptr(), (const __half*)kb_lo.data_ptr(),
            (int*)ent.data_ptr(), cnt_r, (float2*)pml.data_ptr(), nullptr, S, ntile, sel_t0, T0, RS,
            cap, scale);
      }
    }
    const int n_sp = (sel_n > 0) ? ntile * GSPLIT : 0;
    const int win_t0 = sel_t0, n_win = sel_n;
    const int den_t0 = tl0, n_den = std::max(0, std::min(tl1, 8) - tl0);
    const int per_bh = n_sp + n_win + n_den;
    const int total = per_bh * BH;
    if (total > 0) {
      // Persistent blocks: one per resident slot.  Anything beyond that cannot start
      // until a slot frees, and slots only free once the queue is drained.
      static int res_att = 0;
      if (!res_att) {
        int nb = 0;
        C10_CUDA_CHECK(cudaOccupancyMaxActiveBlocksPerMultiprocessor(&nb, k_attend<D>, NTHREAD,
                                                                     (size_t)sh_att));
        res_att = std::max(1, nb) * at::cuda::getCurrentDeviceProperties()->multiProcessorCount;
      }
      int grid = std::min(total, res_att);
      k_attend<D><<<grid, NTHREAD, sh_att, stream>>>(
          qp, kp, vp, (__nv_bfloat16*)out.data_ptr(), (const int*)ent.data_ptr(), cnt_r,
          (__half*)po.data_ptr(), (float2*)pml.data_ptr(), S, ntile, T0, RS, cap, n_sp, n_win,
          win_t0, n_den, den_t0, per_bh, total, (int*)ctr.data_ptr(), scale);
    }
    const int nq_range = T1 - T0, nq_total = BH * nq_range;
    const int rpb = NTHREAD / (D / 8);  // rows per block: one per D/8-lane group
    const int cb = (nq_total + rpb - 1) / rpb;
    k_combine<D, NSLOT><<<cb, NTHREAD, 0, stream>>>(
        (const __half*)po.data_ptr(), (const float2*)pml.data_ptr(),
        (__nv_bfloat16*)out.data_ptr(), S, T0, nq_range, RS, nq_total, (int*)ctr.data_ptr(),
        NSLOT);
  }
}

void nsa_run(at::Tensor q, at::Tensor k, at::Tensor v, at::Tensor out, at::Tensor kb_hi,
             at::Tensor kb_lo, at::Tensor ent, at::Tensor cnt, at::Tensor po, at::Tensor pml,
             at::Tensor msk, at::Tensor ctr, int64_t RS) {
  const int D = q.size(3);
  if (D == 64) run_d<64>(q, k, v, out, kb_hi, kb_lo, ent, cnt, po, pml, msk, ctr, (int)RS);
  else if (D == 128) run_d<128>(q, k, v, out, kb_hi, kb_lo, ent, cnt, po, pml, msk, ctr, (int)RS);
  else TORCH_CHECK(false, "nsa: unsupported head dim ", D);
}

PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { m.def("nsa_run", &nsa_run, "fused NSA attention"); }

20260725_023423_or-opus_anthropic_claude-opus-5_02_deepseek_nsa