Commit 675cf20
feat: make --mtp work for model families that ship MTP heads separately (#137)
* feat: make --mtp work for model families that ship MTP heads separately
--mtp has been silently a no-op for Gemma 4. The gate is `context.model is any
MTPLanguageModel`, and only Qwen35Model, Qwen35TextModel and DeepseekV4Model
conform — those carry their MTP heads inside the main checkpoint. Gemma 4 does
not: Google ships the heads as a separate assistant checkpoint, and
Gemma4AssistantModel conforms to DualModelMTP (MTPLanguageModel plus a
back-reference to the trunk it drafts for). Nothing in Sources/ ever set that
reference except Gemma4MTPBench, which is not a target in Package.swift and so
cannot build — leaving the whole path unreachable.
--mtp-assistant-model loads the assistant, injects mainModelRef, and routes
through the existing generateMTP call. Rather than adding a second generation
branch, mtpContext() picks which context generateMTP should run against: the
main context for in-checkpoint MTP, or a derived context whose model is the
assistant while tokenizer, processor and configuration — and the KV cache
passed alongside — stay the trunk's. That mirrors the reference usage in
Gemma4MTPBench and keeps one code path, so the prompt cache is unaffected.
An explicit flag rather than an id table: the table in #109 maps gemma-4-e4b-it
to the E2B assistant and gemma-4-31b-it to the 26B one, which look like slips,
and a wrong guess here silently drafts from the wrong model.
Measured, and the result is not favourable yet. Output is correct — identical
prefixes to baseline — but throughput is worse on both pairs available here:
gemma-4-e2b-it-4bit + E2B assistant: 136.8 → 117.2 tok/s
gemma-4-26b-a4b-4bit + 26B assistant: 74.1 → 63.6 tok/s
and flat across --num-mtp-tokens 1/2/3 (63.3 / 64.1 / 63.6 on the 26B pair).
Invariance to draft depth points at a fixed per-round cost rather than draft
token cost, which is what the unlanded maxSharedKV=16 cap in #109 targets. Both
assistants also ship bf16 against 4-bit trunks, so each drafted token costs
more than the trunk token it replaces.
So this makes the flag mean something and gives the perf work something to be
measured against; it is not a speedup on its own. MTP stays opt-in and off by
default, and with no --mtp-assistant-model the behaviour is byte-identical to
before.
259 tests pass.
Refs #109.
Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
* fix: bump mlx-swift-lm to pick up the dual-model MTP prefill fix
Points at SharpAI/mlx-swift-lm#46, which makes the Gemma 4 assistant's
callAsFunction delegate to the trunk. Without it this PR's feature aborts on
any prompt over prefillStepSize (512 tokens) with
Fatal error: Layer 0 is a KV-shared layer but received no sharedKV
because MTPTokenIterator.prepare() prefills through context.model, which this
PR makes the assistant — and an assistant checkpoint is entirely KV-shared
layers that cannot run without sharedKV from the trunk.
To be re-pointed at main once #46 lands, since the squash rewrites the SHA.
Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
* test: exercise chunked prefill with a prompt over the 512-token boundary
Every prompt in this repo's test suite is under 80 characters, so prepare()
always returned prompt tokens without forwarding them and chunked prefill was
never run. That gap is how a dual-model MTP crash on any real-sized prompt
reached a green CI (SharpAI/mlx-swift-lm#46) — the failure needed only a
prompt past prefillStepSize to appear, and nothing in CI supplied one.
Adds one ~2700-token request to the contract suite. An empty response is
treated as a failure, not an error case: a crash in prefill drops the
connection rather than returning an error body, which is precisely the
signature being watched for.
This covers the ordinary generate path only. CI runs no --mtp job, so the
speculative variant of the same code path remains uncovered (#128).
Verified locally: server logs prompt=2697t for the new case, suite 10 passed
0 failed 2 skipped.
Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
* chore: re-point mlx-swift-lm at the merged prefill fix on main
SharpAI/mlx-swift-lm#46 landed as squash commit 6a2c179, which replaces the
branch SHA the previous bump pointed at. The tree is byte-identical to the
interim pointer, so the CI already run against this PR still applies — only
the commit identity changes, from a now-deleted branch to main.
Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
---------
Co-authored-by: Claude Opus 5 <noreply@anthropic.com>1 parent b465017 commit 675cf20
3 files changed
Lines changed: 110 additions & 8 deletions
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
| |||
286 | 286 | | |
287 | 287 | | |
288 | 288 | | |
| 289 | + | |
| 290 | + | |
| 291 | + | |
289 | 292 | | |
290 | 293 | | |
291 | 294 | | |
| |||
683 | 686 | | |
684 | 687 | | |
685 | 688 | | |
| 689 | + | |
| 690 | + | |
| 691 | + | |
| 692 | + | |
| 693 | + | |
| 694 | + | |
| 695 | + | |
| 696 | + | |
| 697 | + | |
| 698 | + | |
| 699 | + | |
| 700 | + | |
| 701 | + | |
| 702 | + | |
| 703 | + | |
| 704 | + | |
| 705 | + | |
| 706 | + | |
| 707 | + | |
| 708 | + | |
| 709 | + | |
| 710 | + | |
| 711 | + | |
| 712 | + | |
| 713 | + | |
| 714 | + | |
| 715 | + | |
| 716 | + | |
| 717 | + | |
| 718 | + | |
| 719 | + | |
| 720 | + | |
| 721 | + | |
| 722 | + | |
| 723 | + | |
| 724 | + | |
| 725 | + | |
| 726 | + | |
| 727 | + | |
| 728 | + | |
| 729 | + | |
686 | 730 | | |
687 | 731 | | |
688 | 732 | | |
| |||
810 | 854 | | |
811 | 855 | | |
812 | 856 | | |
813 | | - | |
| 857 | + | |
| 858 | + | |
814 | 859 | | |
815 | 860 | | |
816 | 861 | | |
| |||
920 | 965 | | |
921 | 966 | | |
922 | 967 | | |
923 | | - | |
| 968 | + | |
| 969 | + | |
924 | 970 | | |
925 | 971 | | |
926 | 972 | | |
| |||
1091 | 1137 | | |
1092 | 1138 | | |
1093 | 1139 | | |
| 1140 | + | |
1094 | 1141 | | |
1095 | 1142 | | |
1096 | 1143 | | |
| |||
1352 | 1399 | | |
1353 | 1400 | | |
1354 | 1401 | | |
1355 | | - | |
| 1402 | + | |
| 1403 | + | |
1356 | 1404 | | |
1357 | 1405 | | |
1358 | 1406 | | |
| |||
1626 | 1674 | | |
1627 | 1675 | | |
1628 | 1676 | | |
1629 | | - | |
| 1677 | + | |
1630 | 1678 | | |
1631 | | - | |
| 1679 | + | |
1632 | 1680 | | |
1633 | 1681 | | |
1634 | 1682 | | |
| |||
1637 | 1685 | | |
1638 | 1686 | | |
1639 | 1687 | | |
1640 | | - | |
| 1688 | + | |
1641 | 1689 | | |
1642 | | - | |
| 1690 | + | |
1643 | 1691 | | |
1644 | 1692 | | |
1645 | 1693 | | |
| |||
2714 | 2762 | | |
2715 | 2763 | | |
2716 | 2764 | | |
| 2765 | + | |
| 2766 | + | |
| 2767 | + | |
| 2768 | + | |
| 2769 | + | |
| 2770 | + | |
| 2771 | + | |
| 2772 | + | |
| 2773 | + | |
| 2774 | + | |
| 2775 | + | |
| 2776 | + | |
| 2777 | + | |
| 2778 | + | |
| 2779 | + | |
| 2780 | + | |
| 2781 | + | |
| 2782 | + | |
| 2783 | + | |
| 2784 | + | |
| 2785 | + | |
| 2786 | + | |
| 2787 | + | |
2717 | 2788 | | |
2718 | 2789 | | |
2719 | 2790 | | |
| |||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
| |||
210 | 210 | | |
211 | 211 | | |
212 | 212 | | |
| 213 | + | |
| 214 | + | |
| 215 | + | |
| 216 | + | |
| 217 | + | |
| 218 | + | |
| 219 | + | |
| 220 | + | |
| 221 | + | |
| 222 | + | |
| 223 | + | |
| 224 | + | |
| 225 | + | |
| 226 | + | |
| 227 | + | |
| 228 | + | |
| 229 | + | |
| 230 | + | |
| 231 | + | |
| 232 | + | |
| 233 | + | |
| 234 | + | |
| 235 | + | |
| 236 | + | |
| 237 | + | |
| 238 | + | |
| 239 | + | |
| 240 | + | |
| 241 | + | |
| 242 | + | |
| 243 | + | |
213 | 244 | | |
214 | 245 | | |
215 | 246 | | |
| |||
0 commit comments