Pull requests / #1372

#1372 DeltaNet recurrence: STRATA_GDN_CHUNKED=1 - the prompt's recurrence in 32-token chunks (opt-in, 1.8x faster, other bits)

open · @sergqwer · 0 コメント · GitHub で見る

BenchmarksMulti-GPUAMD / HIPNVIDIA / CUDA

本文

The prompt path's DeltaNet recurrence (`gdn_rec_kh_kernel`) walks a chunk one token at a time. On an RTX 5090 it costs ~0.41 µs per token and layer: 520-530 ms of a 32K prompt (9%). This PR adds an opt-in (`STRATA_GDN_CHUNKED=1`) that computes the recurrence in chunks of 32 tokens, in the WY form flash-linear-attention uses, all in FP32. The kernel is 1.81x faster per launch, and it takes ~210 ms (3.7%) off a 32K prompt. It sums in another order than the sequential kernel, so its bits differ; that is why it is opt-in. Unset (the default), nothing changes.

### What it does

The math, per value head, with γ_t the gate's running sum inside the chunk and S0 the state before it:
- A[t][i] = β_t e^(γ_t − γ_i) k_t·k_i for i < t.
- T = (I + A)⁻¹.
- D = T · β(v − e^γ S0ᵀk): the delta rule's corrections, one per token.
- Outputs and the state after the chunk follow from D as small matmuls. Every exponent is ≤ 0, because the gates are negative, so nothing overflows.

The two kernels:
- **`gdn_chunk_prep_kernel`**, one block per (chunk, key head), for every chunk of a 2048-token super-block at once. It forms k·k and q·k and solves for T (one lane per token). It writes T, P = e^(γ_t − γ_i) q_t·k_i and γ for the key head's three value heads.
- **`gdn_chunk_scan_kernel`**, one block per (key head, 16 value columns): 128 blocks. Each block walks the super-block's chunks with its state slice in registers. The products are register-tiled from shared memory, and the next chunk's k, q, v, γ and β load (cp.async) while the current one computes.

**Scratch.** The prep's output for a super-block takes 25.6 MB. It is allocated once per device, on the first call, and kept; the two kernels' shared-memory attributes are also set once. A stream-ordered allocation and free per call cost more host time than the kernels saved: a 2K prompt's recurrence took 78 ms against 34 ms.

**Where it runs.** Only on prompt chunks of ≥ 128 tokens (`kGdnChunkedMin`; see below), on a CUDA card with:
- sm_80+ (cp.async);
- at least 128 SMs, because the scan's 128 blocks must fit in one wave: each block walks every chunk, so a second wave would double the time;
- 93.7 KB of opt-in shared memory per block.

`STRATA_EMULATE_CC` is honoured. Anywhere else the existing kernels run: smaller cards (RTX 3090, 4080, 5080), a failed scratch allocation, and HIP/SYCL, where none of this code is compiled.

**Builds and tests.** `kernels.cu` compiles for 75/80/86/89/120. `gdn_rec_parity` now:
- checks the chunked recurrence against `gdn_rec_kh_kernel` (fails above 1e-4 relative) at T = 1..4099, 8192 and 32768, from a non-zero state;
- measures both against a recurrence in FP64 up to T = 8192;
- times it (`--bench`).

### Measured

RTX 5090 (170 SMs), Ryzen 9 9950X3D, IQ2_XS (ISTA GSQ-RCO), `--kv int8`, `--expert-cache 12000`.

**Kernel accuracy** (`gdn_rec_parity`, max |diff| / max |value|, output / state after the chunk, 0 failures):

| T | vs FP64: `gdn_rec_kh_kernel` | vs FP64: chunked | chunked vs `gdn_rec_kh_kernel` |
|---|---|---|---|
| 1 | 1.7e-7 / 9.5e-8 | 2.7e-7 / 9.1e-8 | 3.0e-7 / 1.0e-7 |
| 33 | 2.1e-7 / 1.7e-7 | 5.1e-7 / 4.7e-7 | 5.5e-7 / 4.6e-7 |
| 1000 | 2.6e-7 / 1.6e-7 | 5.6e-7 / 1.6e-7 | 6.4e-7 / 2.5e-7 |
| 4099 | 2.2e-7 / 2.0e-7 | 1.05e-6 / 2.8e-7 | 1.13e-6 / 3.1e-7 |
| 8192 | 3.2e-7 / 1.5e-7 | 8.0e-7 / 8.1e-7 | 7.5e-7 / 8.3e-7 |
| 32768 | – | – | 9.0e-7 / 4.8e-7 |

- Both errors are FP32 rounding. The chunked one is up to 5.4x farther from FP64 than the sequential kernel (closer at a few T), never above 1.1e-6.
- Separately, the chunked recurrence gave the same bits on 199 reruns at each of 7 engine shapes (130 to 8192 tokens). Its scratch was filled with NaN before every other run, so it never reads anything the prep has not written.

**End to end**, first-token KL to the default. The noise arm is `STRATA_PROMPT_ATTN_V1=1`: upstream's first prompt-attention kernel, which has the same accuracy but sums in another FP32 order. The last column says whether the 32 greedy tokens match the default's.

| prompt | chunked | noise arm | top-1 | tokens = default (chunked / noise) |
|---|---|---|---|---|
| 600 | 0.0051 | 0.0032 | same | yes / yes |
| 2K | 0.0152 | 0.0132 | same | yes / no |
| 8K | 0.0016 | 0.0009 | same | no / no |
| 32K (4 chunks) | 0.0113 | 0.0067 | same | yes / no |

The switch moves the output by 1.15-1.83x what the noise arm does. The top token is the same everywhere. The greedy tokens match the default's on 3 of 4 prompts with the switch, and on 1 of 4 with the noise arm.

**Speed per launch** (`gdn_rec_parity --bench`, one layer, 48 value heads, medians of alternating runs):

| tokens | `gdn_rec_kh_kernel` | chunked | |
|---|---|---|---|
| 32 | 18.4 µs | 31.6 µs | 0.58x |
| 96 | 45.0 µs | 43.8 µs | 1.03x |
| 128 | 60.4 µs | 50.2 µs | 1.20x |
| 512 | 217 µs | 130 µs | 1.67x |
| 2048 | 0.85 ms | 0.47 ms | 1.81x |
| 8192 | 3.38 ms | 1.86 ms | 1.81x |
| 32768 | 13.51 ms | 7.45 ms | 1.82x |

Below ~100 tokens the two launches cost more than they save, so the threshold is 128 tokens. Since the switch is opt-in, it then applies to nearly every prompt chunk.

**In the engine**, the recurrence phase (`STRATA_PREFILL_TIMING=1`, 3 to 6 runs each):

| prompt | default | chunked |
|---|---|---|
| 600 | 10 ms | 6-7 ms |
| 2K | 34 ms | 21 ms |
| 8K | 130-133 ms | 77-87 ms |
| 32K (4 chunks of 8K) | 520-532 ms | 306-334 ms |

**Prompt time**, 3 interleaved pairs in each order:

| | default | chunked |
|---|---|---|
| 600, default first | 593 / 606 / 604 ms | 602 / 612 / 621 ms |
| 600, chunked first | 623 / 604 / 599 ms | 601 / 607 / 596 ms |
| 2K, default first | 875 / 865 / 873 ms | 884 / 887 / 860 ms |
| 2K, chunked first | 873 / 884 / 861 ms | 872 / 859 / 848 ms |
| 8K, default first | 1517 / 1515 / 1539 ms | 1610 / 1525 / 1617 ms |
| 8K, chunked first | 1510 / 1512 / 1513 ms | 1478 / 1460 / 1459 ms |
| 32K, default first | 5749 / 5597 / 5590 ms | 5830 / 5388 / 5381 ms |

- **Order effect.** When the default ran first, the second run was also slower in the phases the recurrence does not touch: the GPU timeline minus the recurrence was 0.1-9.9% higher. With the order reversed, those phases were even (−3.1% to +1.6%).
- **Where those phases were even** (8K chunked first, 32K pairs 2 and 3), the prompt is faster by about the recurrence's saving: −32 to −54 ms at 8K (2.1-3.6%) and −209 ms at 32K (3.7%).
- **At 600 and 2K**, the saving (4 and 13 ms) is below the machine's noise.

**Identity.** With the switch unset, the first-token logits and the 32 greedy tokens equal v0.1.40.2 on IQ2_XS after 2K, 8K and 32K (identical bytes; 2K repeated on the final build).

Our NVFP4 fork has shipped these kernels since 0.1.40-nvfp4.3. This PR adds the per-device scratch, which removes the per-call cost that had kept the fork's threshold at 16K tokens.

🤖 Generated with [Claude Code](https://claude.com/claude-code)

https://claude.ai/code/session_01VZy1yKaDDiA8a7svdwaHio

関連リンク

インストール・モデル・リリースへの站内リンク。