From-scratch SM80 (A100/A800) CUDA/CUTE kernels for the GDN (Gated DeltaNet) and QSA (query-key sparse attention) attention operators.
Public implementations of this architecture family are mostly Triton kernels (e.g. fla), while first-class CUDA kernel libraries (FlashMLA, FlashKDA, …) target SM90+ — leaving the large installed base of A100/A800 (SM80) GPUs without hand-written CUDA kernels for GDN/QSA. This repo fills that gap.
The kernels implement the compute core of modern open-weight GDN + QSA models
(chunked gated delta-rule linear attention, block-level sparse indexer, sparse
core attention, and the gated output projection) — written from scratch for
A800 (compute capability 8.0), with tensor-core mma.sync + cp.async
kernels and reproducible benchmarks against public baselines.
Status: all four operators shipped —
gdn_chunk(M1),qsa_indexer+output_gate(M2),qsa_core(M3).Validation: clean-A800 build / test / benchmark record in
docs/VALIDATION_LOG.md(37/37 tests PASS).v0.2.4: see
RELEASE_NOTES.md—qsa_indexer3.0× vs eager at S=8192,gdn_chunk1.27× vs fla at S=8192 with peak memory halved;qsa_corereuse path unchanged at S=8192 9.23ms (2.81× TC pass2, 9.89× scalar).
- 2026.09.08 · v0.2.4 — perf release:
qsa_indexerwarp-shuffle hybrid bitonic final sort (45 barriers → 6; topk S=8192 0.79→0.49ms) + a 2×2 register-blocked long-sequence score kernel (score 1.13→0.85ms) → 3.0× vs vectorized eager at S=8192 (27.6× at S=512, 6.3× at S=2048);gdn_chunkworkspace-free fused reset (peak memory ~halved: S=8192 582→356 MB) and 8-warp mma register-persistent replay → 1.27× vs fla 0.5.2 at S=8192. - 2026.09.07 · v0.2.3 — perf release:
qsa_corereuse kernel pair-packed P (one u32 per token pair, half the smem slots/stores), in-warp rowsum (5__shflreductions,plsumarray dropped), and online-softmax state fused into the P phase via ansm_m/sm_lping-pong double buffer — S=8192 11.50→9.23ms, now 2.81× over the per-query TC pass-2 and 9.89× over scalar. - 2026.09.06 · v0.2.2 — perf release:
qsa_corereuse kernel inner-loop fusion — row-max merged into QK (register scores) and rescale fused into P (drops thesm_lbarrier) — S=8192 12.88→11.50ms, now 2.26× over the per-query TC pass-2 and 7.93× over scalar. - 2026.09.06 · v0.2.1 — perf release:
qsa_corereuse kernel round-2 optimizations (loop-invariant fragment/token hoists, mma-fragment smem swizzle, 16-col uint32 gathers) — S=8192 20.85→12.88ms, now 2.01× over the per-query TC pass-2 and 7.07× over scalar; addedRELEASE_NOTES.md. - 2026.09.05 · v0.2.0 — perf release:
qsa_indexerradix-select TopK, beats vectorized eager at all lengths (1.71× at S=8192, 19× at S=512);qsa_coreTC pass-2 v3 (3.52× vs scalar) + query-tile K/V reuse kernel (20.85 ms, 4.37× vs scalar at S=8192);gdn_chunkdynamic GC dispatch; tests 33→37. - 2026.09.04 · v0.1.0 — initial public release: all four operators shipped, 33/33 tests PASS, clean-A800 validation log.
- M0: repo skeleton
- M1:
gdn_chunk— Gated DeltaNet chunked linear attention - M2:
qsa_indexer+output_gate - M3:
qsa_core— sparse attention core (scalar + TC pass2)
Deliberately out of scope for this repo:
fused_linear_ce→ experimental / training-side, not in the main APIflashmla-sm80→ a separate SM80 MLA decode project, to be released independently
- gdn_chunk up to 1.44× vs fla 0.5.2 (1.27× at S=8192, peak memory ~halved); qsa_indexer 3.0× vs vectorized eager at S=8192 (27.6× at S=512, 6.3× at S=2048); qsa_core TC pass-2 3.52× vs scalar, reuse path 9.89× — full reproduce commands in Benchmarks.
- From-scratch SM80 CUDA/CUTE kernels (not Triton wrappers).
- Tensor-core kernels via
mma.sync+cp.async, tuned for A800. - Fair, reproducible benchmarks vs public baselines (fla) — see
docs/BENCHMARK_METHODOLOGY.md. - Numeric correctness policy documented in
docs/CORRECTNESS_POLICY.md. qsa_coreships a scalar path (all dtypes) and a tensor-core pass-2 (v3) path.
| Operator | Component | Target | Dtypes | Notes |
|---|---|---|---|---|
gdn_chunk |
Gated DeltaNet (linear attention) | SM80 (A100/A800) | bf16 | serial / reset-fast-path / two-level scan, auto-dispatch by S |
qsa_indexer |
QSA indexer (MQA 4Q/1K) | SM80 | fp32 | two-stage path: 3.0× vs vectorized eager at S=8192, 27.6× at S=512 (see Benchmarks) |
output_gate |
Gated residual output gate | SM80 | bf16 | RMSNormGated + CUTLASS GEMM |
qsa_core |
QSA sparse-block attention | SM80 | scalar: all / TC: bf16 | TC pass-2 requires D=256 |
Default benchmark shapes follow the public GDN+QSA architecture
(GDN Hq=16/Hv=32/D=128, QSA 24Q/2KV/D=256, bf16). Kernels are specialized
for the published head dimensions and SM80, not locked to any single model.
csrc/ CUDA/C++ sources, one dir per operator
gdn_chunk/ Gated DeltaNet: prepare / stage1 / stage2 / stage3 + host wrapper
qsa_indexer/ block-level MQA indexer (kernel + pybind)
output_gate/ RMSNormGated + out_proj (self-written + CUTLASS GEMM)
qsa_core/ sparse core attention (scalar) + TC pass-2 (v3)
gdn_qsa_sm80/ Python package: per-op functional API + torch references
reference/ pure-torch correctness anchors
tests/ pytest, one file per operator (37 tests)
benchmarks/ per-op benchmark scripts (bench_<op>.py)
docs/ per-op notes, methodology, correctness policy, validation log
tools/ analysis / verification helpers (reuse analysis, proto checks)
third_party/ vendored CUTLASS/CuTe headers (BSD-3-Clause)
scripts/ build.sh / verify_all.sh / bench_all.sh
Source build only (no wheels provided):
export TORCH_CUDA_ARCH_LIST=8.0
pip install -e . --no-build-isolationRequires: CUDA ≥ 12.0, sm_80 target (A800), PyTorch with CUDA, and a C++17
compiler. The extension is compiled in-place; gdn_qsa_sm80.*_cuda .so files
appear under the package after a successful build.
To build only a subset of operators:
GDN_QSA_BUILD_OPS=gdn_chunk,qsa_core bash scripts/build.sh
# or legacy alias: GDN_QSA_BUILD_GDN_ONLY=1 bash scripts/build.shRun on a clean A800 (CUDA_VISIBLE_DEVICES=0):
bash scripts/verify_all.sh # correctness gate: pytest tests/ -x
bash scripts/bench_all.sh # gate + all four benchmark tablesfrom gdn_qsa_sm80 import gdn_chunk
B, S, Hk, Hv, D = 1, 8192, 16, 32, 128
q = torch.randn(B, S, Hk, D, dtype=torch.bfloat16, device="cuda")
k = torch.randn(B, S, Hk, D, dtype=torch.bfloat16, device="cuda")
v = torch.randn(B, S, Hv, D, dtype=torch.bfloat16, device="cuda")
g = -torch.rand(B, S, Hv, dtype=torch.bfloat16, device="cuda") * 2.0
beta = torch.rand(B, S, Hv, dtype=torch.bfloat16, device="cuda").sigmoid()
out, final_state = gdn_chunk(q, k, v, g, beta, output_final_state=True)
# out: [B, S, Hv, D] final_state: [B, Hv, D, D]gdn_chunk auto-dispatches serial / reset-fast-path / two-level scan by decay
strength and S; the reset fast path is workspace-free (peak memory ~halved).
See docs/GDN_CHUNK.md.
from gdn_qsa_sm80 import qsa_indexer, qsa_indexer_topk_only
# full API: also returns the dense [B,S,NB] block_scores matrix (debug/compat)
block_scores, block_indices, selected_scores = qsa_indexer(
q, raw_keys, cos_q, sin_q, cos_k, sin_k, r=4, block_topk=512)
# q/raw_keys fp32, cos/sin [B,S,R] -> block_indices [B,S,KB] int32 (-1 pad)
# memory-light API: fused score+TopK, never materializes block_scores
block_indices, selected_scores = qsa_indexer_topk_only(
q, raw_keys, cos_q, sin_q, cos_k, sin_k, r=4, block_topk=512)See docs/QSA_INDEXER.md.
from gdn_qsa_sm80 import qsa_sparse_core_attention, qsa_expand, qsa_pass2_tc
# scalar (all dtypes)
out = qsa_sparse_core_attention(q, k, v, block_idx, r)
# TC-accelerated pass2 (bf16, D=256)
sel_idx, sel_cnt = qsa_expand(block_idx, r)
out_tc = qsa_pass2_tc(q, k, v, sel_idx, sel_cnt, r)
# query-tile local K/V reuse pass2 — ~2.81x faster than qsa_pass2_tc at
# S=8192 (9.23ms vs 25.89ms): shares each gathered 64-token tile across 4
# adjacent queries, packs P as token pairs, computes the rowsum in-warp, and
# fuses the online-softmax state into the P phase via a ping-pong double
# buffer; auto-falls back to v3 for S>8192.
out_re = qsa_pass2_tc_reuse(q, k, v, sel_idx, sel_cnt, r)See docs/QSA_CORE.md.
from gdn_qsa_sm80 import rmsnorm_gated, out_proj_gemm_cutlass
gated = rmsnorm_gated(y, z, weight) # [N,128] bf16
out = out_proj_gemm_cutlass(gated.reshape(T, 4096), W) # W [N,K] bf16See docs/OUTPUT_GATE.md.
All numbers from a clean A800, same input / same dtype / same GPU / warmup +
median. Baselines: fla for gdn_chunk; vectorized eager with identical math
for qsa_indexer. Full methodology:
docs/BENCHMARK_METHODOLOGY.md.
| S | ours (ms) | fla (ms) | speedup | peak mem (ours, MB) |
|---|---|---|---|---|
| 2048 | 0.423 | 0.500 | 1.18x | 102.2 |
| 4096 | 0.561 | 0.640 | 1.14x | 186.6 |
| 8192 | 0.950 | 1.205 | 1.27x | 355.5 (582 before) |
| 32768 | 3.187 | 4.189 | 1.31x (1.31–1.55 across sessions) | 1385.2 (2291 before) |
Reproduce: CUDA_VISIBLE_DEVICES=0 python benchmarks/bench_gdn_chunk.py
gdn_chunk auto-dispatches serial / reset-fast-path / two-level scan by decay
strength and sequence length. The reset fast path is workspace-free: the
per-group gt metric and last-chunk B_g are recomputed in-CTA from raw
k/v/g/beta (fused stage-1), and each replay CTA recomputes its chunks'
kd/qd/kr/INV/Mqk in-CTA from raw q/k/g/beta (fused stage-3, bit-identical
math to the prepare kernel), so the ~216MB prepare workspace is only
allocated by the exact-scan fallback — peak memory ~halves on the reset
path (S=8192 582→356 MB, S=32768 2291→1385 MB). The replay state is now
register-persistent: each warp owns its 16-col state blocks across the
whole group (per-chunk serial mma latency scales as 1/kWarps) and all 8 warps
join the MMA phases. The superchunk group size is swept per S (the serial
per-group replay chain shortens with smaller groups while the cross-group scan
amortizes over larger ones): GC=8 in the S<=2048 reset band, GC=16 for
S<=4096, GC=32 for the 8192..16384 band, GC=64 beyond (A800 sweep,
g=-rand*2.0).
Same inputs, same dtype (bf16), same GPU, warmup + median. rel-L2 is
||out − fp32_ref|| / ||fp32_ref||, computed for every backend against the
same pure-torch fp32 reference (gdn_chunk_reference), so accuracy and speed
are read off one table. vs X = X_ms / ours_ms (>1 means ours is faster).
| S | ours (ms) | fla (ms) | sglang (ms) | vs fla | vs sglang | rel-L2 ours | rel-L2 fla | rel-L2 sglang |
|---|---|---|---|---|---|---|---|---|
| 2048 | 0.423 | 0.527 | 0.370 | 1.25x | 0.87x | 5.04e-03 | 3.26e-03 | 3.26e-03 |
| 4096 | 0.560 | 0.616 | 0.666 | 1.10x | 1.19x | 5.04e-03 | 3.25e-03 | 3.25e-03 |
| 8192 | 0.950 | 1.174 | 1.251 | 1.24x | 1.32x | 5.05e-03 | 3.26e-03 | 3.26e-03 |
| 32768 | 2.650 | 4.111 | 4.421 | 1.55x | 1.67x | 5.05e-03 | 3.25e-03 | 3.25e-03 |
Reproduce: PYTHONPATH=<repo> python benchmarks/bench_gdn_backends.py
(fla / sglang are optional; each is labelled and skipped if not importable.)
Read this table honestly — it is not a clean sweep:
- vs
fla: ours leads at every length (1.10x → 1.55x, widening with S). - vs the sglang backend: ours is slower at S=2048 (0.87x) and only takes the lead from S=4096 onward (1.19x → 1.67x). That kernel is the serving-path kernel, tuned for short-context decode; at S=2048 its fixed overhead is lower than ours. If your workload is short-sequence, the sglang kernel is the better choice — this table says so.
- Accuracy: ours is ~1.5x the rel-L2 of the fla lineage (5.0e-03 vs 3.3e-03
at every length).
flaand the sglang backend are numerically identical (3.25–3.26e-03), as expected — the sglang kernel is derived from it. Both are ordinary bf16 error levels, but the gap is real and reproducible, and a caller with a tight accuracy budget should know it before switching.
Baseline versions: fla 0.6.0 (from source), sglang kernels/ops/attention/fla/chunk.py.
The
flacolumn in the table further up was measured against a different fla release. The PyPI artifact forfla 0.5.2no longer shipsfla/ops/at all (onlylayers/ models/ modules/ utils/), so that version is not reproducible from PyPI; the numbers here use 0.6.0 from source. Do not compare the two tables'flacolumns directly.
| S | ours (ms) | eager (ms) | speedup |
|---|---|---|---|
| 512 | 0.034 | ~0.93 | 27.6x |
| 2048 | 0.148 | ~0.93 | 6.3x |
| 8192 | 1.126 | 3.373 | 3.0x |
Reproduce: CUDA_VISIBLE_DEVICES=0 python benchmarks/bench_qsa_indexer.py
The two-stage path (pool → encode → tiled score → per-query TopK) beats the
vectorized eager baseline at every length, including S=8192 (3.0x; the
absolute ours ms drift ~±20% across sessions on shared A800 infra, so quote
the ratios — see docs/BENCHMARK_METHODOLOGY.md). The wins come from
SM80-scalar work-splitting and sorting:
- Coalesced preprocess kernels. The pool-keys and encode kernels each run
one warp per (batch·block / batch·query·head) row, so every lane owns one
dims-contiguous
float4; the RMSNorm sum is a 32-lane warp reduction and the partial RoPE pair is exchanged with a single__shfl_xor. The old thread-per-row versions read rows stridedDfloats apart (~2–3% coalescing) and under-filled the GPU at short lengths — pool 117→9 µs, encode 349→27 µs atS=8192, andS=512dropped ~6× overall. - Score kernel: invisible-tile early exit + cp.async staging + float2 keys.
Tiles the causal mask fully hides now write
-infand return before staging (skips the global loads + barrier); the query tile is staged withcp.asyncso its copies overlap the block-key staging; aSCORE_CQ=8tile keeps 3 blocks resident; key rows padded toD+2load 4 dims as 2 float2 reads (half the load-issue slots) with independent per-head accumulators. - TopK round-2 slab narrowing. The
==pivotslab of the radix select is narrowed one byte at a time (CUBblock_topk_airstyle) instead of a fullO(eq_n log² eq_n)bitonic sort; once it fits one warp it is sorted in registers with warp-shuffle bitonic. The exactKwinners are then ordered by a warp-shuffle hybrid bitonic network — consecutive pairs live in registers, everyj<=32exchange is a__shfl_xor, and only the64/128/256distances of the k-merges round-trip through shared memory (6 barriers instead of the previous 45). Candidates are packed to fixed positions (no single-counteratomicAddcontention), andP/K_effcome from a thread-0 pivot walk. - Score kernel: 2×2 register-blocked variant for long sequences. 256
threads but four output cells each (2 query rows × 2 block columns); the
staged block keys stay in shared as unpadded
float4rows and both q and k stage viacp.async(SCORE2_CQ=16/SCORE2_CB=64, 64 B/cell). Dispatched whenNB>=256 && S*NB>=512·1024so every k tile stays full; short problems keep theSCORE_CQ=8kernel.
A fused topK-only variant (qsa_indexer_topk_only) skips materializing the
dense [S,NB] score matrix entirely (lower peak memory, no 64MB write/read at
S=8192) but is per-query CTA-based, so it loses the cross-query block-key
reuse of the tiled score kernel and is not the fast path at long sequences —
the two-stage qsa_indexer is recommended for S=8192.
| T | gate (ms) | proj self (ms) | proj cutlass (ms) | cutlass speedup |
|---|---|---|---|---|
| 512 | 0.011 | 0.279 | 0.087 | 3.19x |
| 2048 | 0.036 | 0.730 | 0.215 | 3.40x |
| 8192 | 0.128 | 2.505 | 0.678 | 3.69x |
Reproduce: CUDA_VISIBLE_DEVICES=0 python benchmarks/bench_output_gate.py
qsa_pass2_tc is the TC pass-2 (v3) kernel for any S; qsa_pass2_tc_reuse
adds query-tile local K/V reuse (auto-dispatched for 1024 ≤ S ≤ 8192).
| S | scalar (ms) | TC pass2 v3 (ms) | TC pass2 reuse (ms) | v3 vs scalar | reuse vs v3 |
|---|---|---|---|---|---|
| 512 | 1.37 | 0.35 | 0.35* | 3.91x | 1.00x |
| 2048 | 14.63 | 3.85 | 1.52 | 3.80x | 2.54x |
| 8192 | 91.27 | 25.89 | 9.23 | 3.52x | 2.81x |
* S=512 still routes to v3 (union-build overhead; dispatch gate unchanged).
Reuse is within noise of v3 at S=512 (its union-build overhead doesn't pay at
short sequences), so the auto-dispatch falls back to v3 below S=1024.
The S=8192 reuse path is 9.89x over scalar.
Reproduce: CUDA_VISIBLE_DEVICES=0 python benchmarks/bench_qsa_core.py
All tables are from the clean-A800 bash scripts/bench_all.sh run logged in
docs/VALIDATION_LOG.md.
gdn_chunkhas noinitial_stateinput — it always starts from a zero recurrent state.grep -rn "initial_state" gdn_qsa_sm80/returns nothing: no parameter, no buffer, no write-back path. A serving stack that carries a persistent state pool across decode steps therefore cannot use this kernel as a drop-in replacement. sglang's GDN backend is exactly that stack — it callschunk_gated_delta_rule(..., initial_state=ssm_states, initial_state_indices=cache_indices, inplace_update=...)on every extend, including the continuation path where the state is non-zero.- So there is no serve-level comparison, and that is a deliberate choice, not
a gap in effort. The zero-state path is verified equivalent to sglang's
GDN math (
benchmarks/check_sglang_equivalence.py): 5.29e-03 rel-L2 at S=2048 and S=8192, at the bf16 noise floor. The non-zero-state case returns 9.29e-02 — that number is the capability gap, not an accuracy figure, because the two sides are not computing the same function. Standing a server on top of it would produce a service that starts, emits tokens, and computes something else; its tok/s and speculative-accept numbers would be meaningless. Publishing them would be worse than publishing nothing. - What is already shown to be viable: variable-length batching. Running the
kernel per segment and stitching the outputs reproduces sglang's
cu_seqlenspath to 5.28e-03 rel-L2 (ragged 512+1024+2048). An integration needs per-segment dispatch, not a new varlen kernel. - In-kernel q/k L2 normalisation is not implemented. sglang normalises
inside the kernel (
use_qk_l2norm_in_kernel=True). The equivalence gate applies the same normalisation in torch on both sides so the comparison isolates the math; an integration must either move it into the kernel or do it in Python ahead of the call. - Accuracy is ~1.5x the rel-L2 of the fla lineage (5.0e-03 vs 3.3e-03 — see the three-way table above). Ordinary bf16 levels, but real and reproducible.
- vs the sglang backend, this kernel is slower at S=2048 (0.87x) and only takes the lead from S=4096. For short-sequence work the sglang kernel is the better choice.
qsa_indexernow beats the vectorized eager baseline at all lengths, includingS=8192(3.0x), via coalesced preprocess kernels, a narrowing radix-select TopK with a warp-shuffle hybrid bitonic final sort, cross-query block-key reuse, and a 2×2 register-blocked long-sequence score kernel. Further gains would come from a fused score+radix kernel (single pass, no dense[S,NB]materialization) and tensor-core score with fp32-emulation precision.qsa_coreTC pass-2 is a 3.52x win over scalar (S=8192 25.89ms). A query-tile local K/V reuse kernel (qsa_pass2_tc_reuse) is shipped for 1024≤S≤8192: it groups four adjacent queries per CTA, builds the union of their selected token sets in smem, and shares each gathered 64-token K/V tile — S=8192 9.23ms (2.81x over the per-query v3 path, 9.89x over scalar), S=2048 1.52ms (2.54x over v3). On top of the union-sharing it (a) fuses the softmax P/plsum into one barrier, (b) swizzles the Q/K/P smem layouts so each mma A/B fragment register (columns j and j+8) is one LDS.32 instead of two scattered LDS.32, (c) hoists the PV/QK fragment loads and the softmax token/ownership reads out of the per-query loops (the mma B operands and union_tok/qmap are qq-invariant), (d) gathers 16 cols/thread so swz pairs store as uint32, (e) packs P as token pairs — one u32 per(row, token-PAIR)halves the P smem footprint and P-store instructions, (f) computes the rowsum entirely in-warp via 5__shflreductions (theplsumsmem array is dropped), and (g) fuses the online-softmax state update into the P phase via ansm_m/sm_lping-pong double buffer (no separate phase-5 barrier). Round-2 fusion merges the softmax row-max into QK (register scores) and the rescale into P (mnewphase, nosm_lbarrier). Block-reuse analysis (tools/analyze_qsa_core_reuse.py) showed ~90% of the gather L2 traffic is shared across adjacent queries.fused_linear_ceandflashmla-sm80intentionally live outside this repo (see Scope).
- fla — public baseline for the GDN benchmark; kernel implementations here are independent.
flashmla-sm80— an SM80 FlashMLA decode optimization, maintained separately.- NVIDIA/CUTLASS — CuTe/CUTLASS headers
vendored under
third_party/cutlass(BSD-3-Clause).
Built on lessons from hand-optimizing GDN/QSA-family kernels for SM80 (two-level scan, tensor-core QK/PV, bank-conflict-free layouts). Thanks to the fla and CUTLASS communities for public baselines and primitives.
Apache-2.0. third_party/ retains its own licenses
(third_party/cutlass/LICENSE, BSD-3-Clause).