Pull requests / #1107

#1107 prefill: gather experts in groups on short prompts too

open · @brenoperucchi · 0 comentarios · En GitHub

BenchmarksSetup & installNVIDIA / CUDAModels & quantsWindows

Descripción

This replaces #1033, which GitHub closed when `main` was rewritten. It is rebased on v0.1.40.2 and groups the staged experts as well, which raises the short-prompt gain from 12-14% (#1033 alone) to 15-20%.

#372 (in 0.1.38) gathers an MMQ group's experts in one launch, but only in the streamed walk. A chunk below `stream_all_min()` (1024 tokens) takes the staged walk, where every expert gets its own `gather_native` launch, and a staged one also gets its own wait and release event. @BlueKingMuch noticed this in #372: the 923-token tail of a 64K prompt was not affected.

On our box most short requests are 150-600 tokens. On 0.1.39, `STRATA_PREFILL_TIMING=1` showed a ~450-token prompt spending 50-55% of its GPU timeline in the "dequant" phase, which on the MMQ path is the gather. About three quarters of its experts are already in the GPU cache. The other quarter (3,100-3,800 per prompt here) comes through the 8-slot staging ring.

The staged walk was left out of #372 because its ring is smaller than a group of 16, so a group cannot hold its staged slots. This patch gives the staged walk a 32-slot ring (two groups) and lets it stage at most 15 entries ahead. A slot is then refilled only after the group that read it has been gathered, with one wait on the group's last copy and one event releasing all of its slots, as in the streamed walk. Resident experts join the same groups and need no wait. With #789's stream-ahead (counting transfers in flight), a slot is reused 32 transfers after it was filled, when at least 24 have been computed, so its group has been flushed and released.

`stage_slots()` counts those 32 slots everywhere the ring is counted or taken, so a short prompt borrows every byte it writes. My first version allocated 32 slots but still counted 8, and the extra slots ran into live expert-cache slots. The engine then hung on the first request without an error. Long prompts are unchanged, since their ring is already larger than 32. Products, their order, the MMQ tail memsets and the streamed walk are as before, and `STRATA_PREFILL_GROUP_GATHER=0` still turns grouping off everywhere. One file, `src/prefill/prefill.cpp`.

Tested on an RTX 5090 32 GB, Ryzen 9 5950X, 96 GB DDR4, Windows 11 (WDDM), CUDA 13.0, with the Swift IQ3_XXS native pack, `--expert-cache auto`, `--prefill auto:32768`, `--kv int8`, `--max-context 32768`. Both arms are v0.1.40.2 (e8ca9af) built here with the setup.py flags plus `STRATA_PORTABLE=ON`, one with this commit and one without.

For short prompts I used the `gprobe` size of strata-bench (https://github.com/brenoperucchi/strata-bench, tag `core-v0.1.1`): chat requests with a ~100-token system prompt (reused from the prompt cache) and a different excerpt of `src/program/generate.cpp` each time, 1 token out, 30 requests per size, median of the last 28. Two passes per arm, alternating, no timing switches:

| prompt read | v0.1.40.2 | this PR |
|---|---:|---:|
| ~150 tokens | 323 ms | 258 ms (-20%) |
| ~300 tokens | 397 ms | 327 ms (-18%) |
| ~450 tokens | 450 ms | 383 ms (-15%) |

Longer prompts take the streamed walk and read at 3,164 tok/s (2.6K) and 6,464 tok/s (14.7K) on v0.1.40.2, against 3,157 and 6,463 tok/s here. Decode moves by about 10% between passes on both builds.

To check that the results do not change, I used `STRATA_STATE_HASH`, which prints a fingerprint of the whole session state (GDN state, PLE, tail, pooled rows, K/V, MTP) after a prompt has been read. Without `--short-read` every prompt goes through the batched prefill. I compared v0.1.40.2 with this branch, 6 prompts at each of 256, 512, 900, 2048 and 4096 tokens (the first three take the staged walk, the last two the streamed walk): all 30 fingerprints are identical. The unpatched build run a second time also matches all 30. The 84 one-token answers of the `gprobe` runs are identical between the two builds as well.

@blange48 tested #1033 (the resident-only half) and this version on a 4-GPU layer split and saw identical batch output and about -1% to -2% TTFT on short prompts. On a split the stage hand-offs take most of a short prompt's time, which would explain the smaller number.

En el sitio

Enlaces a install, modelos, releases.