Pull requests / #1253

#1253 Enable serial multi-GPU batch MTP

open · @ilumn · 0 Kommentare · Auf GitHub

BenchmarksMulti-GPUAMD / HIPNVIDIA / CUDAModels & quantsWindows

Beschreibung

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.

Mehr auf der Site

Links zu Install, Modellen, Releases.