Pull requests / #1253
#1253 Enable serial multi-GPU batch MTP
open · @ilumn · 0 comments · View on GitHub
BenchmarksMulti-GPUAMD / HIPNVIDIA / CUDAModels & quantsWindows
Description
This enables per-slot multi-token prediction (MTP) when the target model is split across CUDA GPUs. Previously, `--layer-split 24 --batch 4 --batch-mtp` disabled batch MTP and used ordinary batching. The enabled path verifies one proposal per slot through the stages serially, with `--batch-groups 1`, and commits the same accepted prefix on every stage. A slot drafter needs the target model's final residual and output head, which reside on the last GPU in a layer split. Slot drafters now bind that GPU's weights and head, retain their residual buffers and draft-KV copies there, and use the first stage's PLE history. Each target stage supports eight-row verification windows with bounded graph caching. Allocations, idle waits and cleanup select the GPU that owns the corresponding buffers. Shared draft-vocabulary metadata supports the top-two diagnostic. This depends on [#1242](https://github.com/Niko1221/Strata/pull/1242), which isolates concurrent slots' image positions and corrects shared draft-head metadata. That dependency is included in this main-based PR. The serial multi-GPU feature commit is `0de9365f18a2ddcc7e428cc70a639ffcdfcc3b55` (7 files, 210 additions, 25 deletions), based on dependency commit `8fc9490e9cb62e4d3c3ab5f09c1ca8d42f4a8623`; validation used main `82f46a8c8f475f001ad76d92f58f4a4f8ffb0253`. ## Validation on Talos CUDA 13 / GNU 13.3 Release build, 2× RTX 5090 32 GiB, Swift IQ3_XXS, native Q5_K head, 106299-token draft subset, 8K allocated context, split 24, visible device order 1 then 0, 3 GiB reserve per GPU. Exactness controls: `STRATA_IQ_MT_MIN=1 --pcie-frac 0 --adapt-every 1000000 --no-prefill-borrow`. Requests generated up to 64 tokens; long prompt 1210 tokens. Independent review covered this serial implementation and its validation evidence. | Check | Coverage | Result | |---|---|---| | CUDA build, whitespace, Python syntax | Full native build and new harness | Pass | | CPU regressions | 27 checks | Pass | | Host native regressions | 13 on GPU1, 3 on GPU0 | Pass | | Dual batch 2 / spec 2 | Four speculative rows across both stages | Exact solo parity | | Five-slot rotation | Eight-row windows, greedy and seed 1234 / temperature 0.7 / top_k 20 | Exact solo parity | | Q4 subset head, top-two | `--mtp-q4 head`, `STRATA_MTP_TOP2=1` | Exact solo parity | | Single-GPU regression | Four-slot batch MTP | Exact solo parity | | Plain split/fallback | Ordinary batching; groups 2 disable MTP with expected warning | Exact solo parity | | Staggered lifecycle | Long admission, completed next turn, completed/cancelled slot → solo, reuse; greedy/seeded | 7/7 each | | Synthetic mixed images/text | Different 4×6/6×4 grids, malformed admission, cancel/reuse; greedy/seeded | 10/10 | | Real mixed images/text | Existing CPU encoder, 11×5/5×11 grids, same lifecycle; greedy/seeded | 10/10 | | Equivalent-mode performance | 18 runs, 3 repetitions per mode/profile, identical per-stage cache counts and cross-mode tokens | Pass | Speculative runs were checked for grouped-row captures and no disabled-option warning; lifecycle logs prove cache and turn-checkpoint restore paths ran. ## Performance and limits Four fixed prompts, 64 tokens each, same binary/settings, 3 repetitions per batch mode/profile. Rates below are median tokens/s. Normal MTP solo controls ran in each invocation; its rates separate decode from prompt reading. Batch overall includes admissions; the second batch rate starts at the first batch token and still includes overlapping admissions. | Experts per stage | Normal MTP solo decode / with prefill | Plain serial overall / from first batch token | Plain pipeline overall / from first batch token | Batch MTP overall / from first batch token | |---|---|---|---|---| | 4000 | 70.9 / 55.8 | 55.6 / 59.5 | 61.0 / 66.2 | 54.7 / 58.7 | | 11000 | 180.1 / 148.0 | 150.1 / 164.4 | 176.7 / 196.8 | 148.0 / 161.8 | Batch MTP was slower than ordinary pipelined batching in both measured profiles. It adds opt-in compatibility for split models; the default is unchanged. These 64-token measurements do not establish long-context or sustained server-load performance. Explicit fallbacks remain for `--batch-groups > 1`, same-device split stages, helper caches, remote expert optimization and peer devices. `--pipeline-windows` remains off with batch slots. Validated scope is two CUDA GPUs and the existing model/runtime above. HIP, more than two GPUs, more than five slots, worst-case context/memory pressure, other model formats and speculative pipeline groups are unrun. The existing idle split BYIELD path is outside this validation; the new harness checks supported staggered admission. Penalty sampling and coupled draft sampling were not validated. With a compatible two-GPU model config, the included tools reproduce rotating batches and lifecycle checks: ```bash CUDA_VISIBLE_DEVICES=1,0 python tools/batch_test.py --exe build/strata --config strata-model.json --batch 5 --n 5 --max-new 64 --extra "--batch-mtp --batch-groups 1 --layer-split 24 --trim-stage-weights --vision --pcie-frac 0 --adapt-every 1000000 --no-prefill-borrow" CUDA_VISIBLE_DEVICES=1,0 python tools/batch_mtp_split_test.py --exe build/strata --config strata-model.json --keys "temperature=0.7 top_k=20 seed=1234" --extra "--batch-mtp --batch-groups 1 --layer-split 24 --trim-stage-weights --vision --pcie-frac 0 --adapt-every 1000000 --no-prefill-borrow --prompt-cache 4096" ``` Mixed-image/performance harnesses, exact local command records and logs are retained separately in the validation bundle. ## Longer-run exactness investigation A separate experimental implementation overlaps speculative slot groups across the GPU stages. Its 32-request workload at concurrency 4 and 512 generated tokens per request produced intermittent exact-token mismatches, including with overlap disabled while retaining the experiment's private verifier/drafter structure. Its benchmark matrix is paused. The cause and whether the published serial implementation is affected remain unresolved; the scoped results above do not establish long-run exactness for that experiment. The existing [QFUSE activation-quantization issue (#1139)](https://github.com/Niko1221/Strata/issues/1139) was investigated as a possible explanation: batched GDN can consume stale quantized activations when `STRATA_QFUSE=1`. The [QFUSE correction in the concurrent-generation PR (#1209)](https://github.com/Niko1221/Strata/pull/1209) addresses that path, with [runtime validation reported on gfx1151](https://github.com/Niko1221/Strata/issues/1139#issuecomment-6025611893). The Talos CUDA diagnostic configurations leave QFUSE disabled, so this specific mechanism is inactive there. Other underlying causes remain under investigation.
Related on strata.com
Editorial links to help you install, pick models, or read release notes — not part of the upstream thread.