Pull requests / #413

#413 prefill: DeltaNet recurrence with the three value heads of a key head in one thread (bitwise; 1.4x on a 4080 SUPER, 1.3x on a 3090)

closed · @BlueKingMuch · 0 コメント · GitHub で見る

BenchmarksSetup & installMulti-GPUNVIDIA / CUDAModels & quantsWindowsLinux

本文

Measurements on other GPUs are still welcome, see the end of "Other cards".

The change gives the same bits on every card, but how much faster it is depends on how many SMs a card has. I measured that on an RTX 4080 SUPER (80 SMs) with the engine's grids on 32 to 80 of its SMs, and @nsandrus measured the same on an RTX 3090 (82 SMs, Ampere). On both cards it is faster or a draw at every SM count we could measure (32 to 82), so it runs on every CUDA card from Ampere on that holds its 64 blocks at once (32 SMs and more).

On the 4080 SUPER the prompt path's DeltaNet recurrence gets 1.44x faster and its output norm 1.52x: +2 % prefill tokens/s, measured on 0.1.32 and before that on 0.1.31. On the 3090: 1.28x and 1.46x.

Below: what changes, what I measured, how the SM count enters, and a 5-minute test that needs no model.

## What changes

- **`gdn_rec_kh_kernel`:** `gdn_rec_cols_pipe_kernel` gives each warp 32 value columns of one row group, so all 32 lanes read the same 32 q and 32 k values of the token from shared memory (the k twice: for the kv sum and for the state update), with five `__syncthreads` per token. 
The new kernel runs column c and row group rg of the three value heads that share a key head (`head % 16`: heads qh, qh + 16, qh + 32) in one thread: every q/k value it loads feeds three heads, and the token's k row goes into registers once for both of its uses. 64 blocks of 128 threads (16 key heads x 4 column blocks) instead of 192. 
The inputs (the q/k rows of the key head, the block's v columns of the three heads, gate and beta) come in blocks of 8 tokens, copied by cp.async while the block before computes; the two cross-row-group sums have their own arrays, so a token needs 2 `__syncthreads` instead of 5. 216 registers, 26 KB static shared memory, no spills (nvcc 13.3, sm_89).
- **Where it runs:** CUDA with compute capability 8.0 or newer (cp.async), on a card that holds all 64 blocks at once (`cudaOccupancyMaxActiveBlocksPerMultiprocessor` x SMs >= 64; with 2 blocks per SM that is 32 SMs or more). Asked once per device, per call from the current device. 
Everywhere else, and with `STRATA_GDN_KEYHEAD=0`, the kernel as before. `STRATA_GDN_PIPELINE=0` and `STRATA_GDN_REC_HEADS` keep their meaning.
- **`gdn_out_norm_kernel`** no longer stores the normalized value in FP32: nothing reads it, the out projection takes the FP16 copy (`y_h`). `y` stays the FP32 scratch for the recurrence output (the header says so now).

## Exactness

- Per value head and column the same arithmetic in the same order as `gdn_rec_cols_pipe_kernel`: the kv and output sums as the same `fmaf` chains over each row group's 32 rows, the four row-group partials added as `r0 + r1 + r2 + r3`, delta and the state update unchanged.
- `gdn_rec_parity`, a new GPU ctest next to `iq_multi_parity` (synthetic inputs of the model's shape, no model): both recurrences and the output norm are copied from `kernels.cu` verbatim (the norm also as it was before this PR), and nvcc 13.3 emits the same SASS for the copies as for the engine's kernels (compared with cuobjdump). Output, state after the chunk and FP16 output are compared bit for bit for T = 1, 2, 3, 4, 5, 7, 8, 9, 15, 16, 17, 31, 33, 257, 1000 and 4099 tokens from a non-zero starting state. On this branch: 0 failures.
- End to end: the greedy answers after each prompt below are the same tokens with and without `STRATA_GDN_KEYHEAD=0`.

## Measured on an RTX 4080 SUPER

RTX 4080 SUPER 32 GB (sm_89, 80 SMs), Ryzen 7 5800X3D, 64 GB DDR4, PCIe 3.0 x16, Windows 11, CUDA 13.3. The numbers below are from 0.1.32 (c499bd1) with this PR only. 

I measured the change on 0.1.31 first and repeated everything after rebasing it onto 0.1.32 (no conflict): on 0.1.31 it gave 1.42x for the recurrence, 1.51x for the norm and +2.1 % / +2.0 % prefill tokens/s on the same prompts.

**The kernels alone** (`gdn_rec_parity --bench`: one layer, all 48 value heads, the two variants' launches alternating, medians of 20 runs at 8192 tokens and 10 at 32768):

| | 8192 tokens | 32768 tokens |
|---|---|---|
| recurrence, before | 4.59 ms (0.560 µs/token) | 18.82 ms (0.574 µs/token) |
| recurrence, this PR | 3.25 ms (0.397 µs/token), 1.41x | 13.04 ms (0.398 µs/token), 1.44x |
| output norm, before | 1.13 ms | 4.65 ms |
| output norm, this PR | 0.80 ms, 1.40x | 3.05 ms, 1.52x |
| both, 36 DeltaNet layers per chunk | 206 -> 146 ms | 845 -> 579 ms |

**In the engine:** IQ3_S, `--prefill 32768`, `--max-context 262144`, int8 KV, 11055 experts in VRAM; prompts of 32000 and 64000 tokens of source code, fresh (nothing cached), two runs each, `STRATA_PREFILL_TIMING=1`; this PR against `STRATA_GDN_KEYHEAD=0` in the same build:

| | before (`STRATA_GDN_KEYHEAD=0`) | this PR |
|---|---|---|
| 32000-token prompts | 2945 t/s (runs 2879, 3011) | 3012 t/s (runs 2942, 3082), +2.3 % |
| 64000-token prompts | 2799 t/s (runs 2794, 2804) | 2862 t/s (runs 2857, 2866), +2.2 % |
| time to first token, 64000 tokens | 22.9 s | 22.4 s |
| phase `gdn recurrence` (recurrence + norm), all six chunks | 4647 ms, 7.6 % of the GPU time | 3451 ms, 5.8 % |
| GPU timeline, all prompts | 60.9 s | 59.4 s |

## Other cards: measured with the engine's grids on fewer SMs

Each block of either kernel walks the whole chunk, so the busiest SM sets the time. The kernel before puts ceil(192 / SMs) of its blocks on it, the new one ceil(64 / SMs): one per SM from 64 SMs up, two on some SMs below that.

`gdn_rec_parity --bench` now runs both kernels with the engine's grids on N of the card's SMs, from all of them down to 32, while a sleeping kernel holds the others (each of its blocks takes all the shared memory a block may have, so no other block fits beside it). A probe kernel with the same grid, block size and blocks per SM shows where the blocks land: none on a held SM, and how many share the busiest one. On the 4080 SUPER, at 8192 tokens, medians of 15 alternating runs:

| SMs (measured at) | for example | busiest SM, before / new | 4080 SUPER: µs per token, before -> new | 4080 SUPER | 3090 (@nsandrus) |
|---|---|---|---|---|---|
| 64-95 (64, 66, 70, 72, 76, 80) | 4080 / 4080 SUPER, 4070 Ti SUPER, 5080, 5070 Ti, 3090, 3080, A5000 | 3 / 1 | 0.556-0.562 -> 0.395-0.398 | 1.40-1.42x | 1.28-1.29x |
| 48-63 (48, 52, 56, 60, 63) | 4070 Ti, 4070 SUPER, 5070, 3070 Ti | 4 / 2 | 0.607-0.619 -> 0.593-0.594 | 1.02-1.04x | 1.03-1.06x |
| 39-47 (39, 40, 44, 46, 47) | 4070, 3070 | 5 / 2 | 0.760-0.777 -> 0.591-0.593 | 1.28-1.31x | 1.28-1.30x |
| 32-38 (32, 34, 36, 38) | 4060 Ti, 5060 Ti, 3060 Ti | 6 / 2 | 0.912-0.928 -> 0.592-0.594 | 1.53-1.57x | 1.54-1.55x |
| below 32 | 4060, 5060, 3060 | the kernel as before (the 64 blocks do not fit at once) | | | |

At all 80 SMs this gives the direct measurement again (0.560 -> 0.395 µs, 1.42x). Only the busiest SM counts: at 63 SMs, 3 SMs carry 4 blocks of the kernel before and 1 SM carries 2 of the new one, and the times are those at 48 SMs, where every SM carries 4 (before) and 16 carry 2 (new). The estimate the bench prints before this part (both kernels at b blocks on every SM, the table in the first version of this PR) came out 1-6 % above these rows, most at 64-95 SMs (1.49x): it puts b blocks on all 80 SMs, more than the real grid gives the card.

L2, memory and clock stay this card's, so @nsandrus ran the same on an RTX 3090 (82 SMs, Ampere, Linux; the last column above, measured on 32 to 82 of its SMs): bitwise again, and 1.46x / 1.48x for the norm. The new kernel's time follows the clock (what is left in it is the latency of a chain of dependent FMAs): at 1 block per SM it took 1.45 times as long there as here, about the ratio of the boost clocks, the kernel before only 1.33 times. So the 3090 gains less where the new kernel has one block per SM, and as much as this card where the kernel before has five or six.

**No rule beyond fitting:** with 48 to 63 SMs the new kernel has 2 blocks on its busiest SM against at most 4 of the kernel before, a draw on both cards (1.02-1.04x here, 1.03-1.06x on the 3090) and slower on neither. An earlier version of this PR (8badd17) kept the kernel before there, expecting Ampere to come out slower; the 3090's bands showed otherwise, and f73ad58 takes it back. The output norm change applies on every card.

**Not measured:** 96 SMs and more (4090, 5090, RTX 6000 Ada: 2 / 1 per SM), which a card with 80 SMs cannot stand in for; the estimate gives 1.31x there. And no Blackwell card yet. A run of the test below on such a card would be the most useful.

## How to measure (about 5 minutes, no model needed)

Trying it changes no output: both kernels give the same bits.

```
git fetch https://github.com/BlueKingMuch/Strata gdn-keyhead:gdn-keyhead
git checkout gdn-keyhead
cmake --build build --target gdn_rec_parity
build\gdn_rec_parity.exe --bench          (Linux: ./build/gdn_rec_parity --bench)
git checkout main
```

On Windows, run `cmake` from an "x64 Native Tools Command Prompt for VS 2022" in the Strata folder (the build folder SETUP.bat made). This builds only the test program; the engine stays as it is. Please post the first line (card and compute capability), everything from `--bench:` on, and the last line (`gdn_rec_parity: 0 failures`). The last part of the bench (`the engine's grids on N of this card's SMs`) takes about 20 seconds.

If you run long prompts anyway: the engine on this branch once as it is and once with `STRATA_GDN_KEYHEAD=0` (`STRATA_PREFILL_TIMING=1` adds a line per chunk with the `gdn recurrence` phase), and the prompt-processing speed of a fresh long prompt. Starting Strata on this branch compiles the engine once.

## What else I tried

All bitwise, all in earlier versions of `gdn_rec_parity` (on 0.1.31), at 32768 tokens on the 4080 SUPER against the kernel before:

- One head per thread, blocks of 4 or 8 tokens staged with cp.async, 2 syncs per token: 1.18x / 1.22x.
- Three heads per thread with blocks of 4 tokens: 1.31x (the staging work per block is spread over fewer tokens).
- Three heads per thread with one `__syncthreads` per token (the next token's kv sums run in the loop of this token's output sums, the cross-warp sums double-buffered): 1.38x, no better than two syncs.

What is left is the chain itself: per row group and token 32 dependent FMAs for the kv sum and 32 for the output, with one warp per SM partition, so their latency shows. A chunked (WY) formulation on tensor cores would remove the chain but would not give the same bits, so it is not part of this.

To reproduce my numbers: `STRATA_GDN_KEYHEAD=0` against unset, `STRATA_PREFILL_TIMING=1` for the phases, `gdn_rec_parity --bench` for the kernels alone.

関連リンク

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