Skip to content

Commit 675cf20

Browse files
solderzzcclaude
andauthored
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

File tree

‎Sources/SwiftLM/Server.swift‎

Lines changed: 78 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -286,6 +286,9 @@ struct MLXServer: AsyncParsableCommand {
286286
@Option(name: .long, help: "Number of MTP tokens to generate per speculation round (default: 3)")
287287
var numMtpTokens: Int = 3
288288

289+
@Option(name: .long, help: "Assistant checkpoint providing MTP heads, for model families that ship them separately instead of in the main checkpoint (Gemma 4). Ignored when the main model carries its own MTP heads (Qwen3.5, DeepSeek V4).")
290+
var mtpAssistantModel: String?
291+
289292
mutating func run() async throws {
290293
// Raise the open-file limit: large sharded models (e.g. Kimi K2.5, 182 safetensor
291294
// shards) + draft model + metallib + dylibs can exhaust the default macOS FD limit of 256.
@@ -683,6 +686,47 @@ struct MLXServer: AsyncParsableCommand {
683686
draftModelRef = nil
684687
}
685688

689+
// ── Load the MTP assistant, for families that ship MTP heads separately ──
690+
// Qwen3.5 and DeepSeek V4 carry their MTP heads inside the main checkpoint and
691+
// conform to MTPLanguageModel directly, so generateMTP already works for them.
692+
// Gemma 4 does not: Google ships the heads as a separate assistant checkpoint,
693+
// and Gemma4AssistantModel conforms to DualModelMTP — MTPLanguageModel plus a
694+
// back-reference to the trunk it drafts for. Loading it here and injecting that
695+
// reference is what makes --mtp mean anything for Gemma 4; without it the flag
696+
// is silently a no-op, because the main model fails the MTPLanguageModel test.
697+
var mtpAssistantModelRef: (any DualModelMTP)? = nil
698+
if self.mtp, let assistantPath = self.mtpAssistantModel {
699+
print("[SwiftLM] Loading MTP assistant: \(assistantPath)")
700+
var assistantConfig: ModelConfiguration
701+
if FileManager.default.fileExists(atPath: assistantPath) {
702+
assistantConfig = ModelConfiguration(directory: URL(filePath: assistantPath))
703+
} else if let local = ModelStorage.validatedContentDirectory(for: assistantPath) {
704+
assistantConfig = ModelConfiguration(directory: local)
705+
} else {
706+
assistantConfig = ModelConfiguration(id: assistantPath)
707+
}
708+
if self.streamExperts { assistantConfig.lazyLoad = true }
709+
let assistantDownloader = HubDownloader(hub: HubApi(downloadBase: cacheRoot))
710+
let assistantContainer = try await LLMModelFactory.shared.loadContainer(
711+
from: assistantDownloader,
712+
using: TransformersTokenizerLoader(),
713+
configuration: assistantConfig
714+
) { _ in }
715+
mtpAssistantModelRef = await assistantContainer.perform { assistantContext in
716+
assistantContext.model as? (any DualModelMTP)
717+
}
718+
if mtpAssistantModelRef == nil {
719+
print("[SwiftLM] ⚠️ \(assistantPath) does not provide MTP heads (not a DualModelMTP).")
720+
print("[SwiftLM] Ignoring --mtp-assistant-model; generation will not use MTP.")
721+
} else {
722+
// The assistant drafts *for* this trunk, so it needs a reference to it.
723+
await container.perform { mainContext in
724+
mtpAssistantModelRef?.mainModelRef = mainContext.model
725+
}
726+
print("[SwiftLM] MTP assistant ready (\(self.numMtpTokens) tokens/round)")
727+
}
728+
}
729+
686730
// ── Load DFlash draft model for block-diffusion speculative decoding ──
687731
let dflashModel: DFlashDraftModel?
688732
let dflashBlockSizeConfig = self.dflashBlockSize
@@ -810,7 +854,8 @@ struct MLXServer: AsyncParsableCommand {
810854
prefillSize: self.prefillSize,
811855
turboKV: self.turboKV,
812856
mtp: self.mtp,
813-
numMtpTokens: self.numMtpTokens
857+
numMtpTokens: self.numMtpTokens,
858+
mtpAssistantModel: self.mtpAssistantModel
814859
)
815860

816861
let parallelSlots = self.parallel
@@ -920,7 +965,8 @@ struct MLXServer: AsyncParsableCommand {
920965
request: request, bodyData: bodyData, config: config, container: container, semaphore: semaphore, stats: stats, promptCache: promptCache,
921966
draftModelRef: draftModelRef, numDraftTokens: numDraftTokensConfig,
922967
dflashModel: dflashModel, dflashBlockSize: dflashBlockSizeConfig,
923-
dflashTargetModel: dflashTargetModel
968+
dflashTargetModel: dflashTargetModel,
969+
mtpAssistant: mtpAssistantModelRef
924970
)
925971
} catch {
926972
let errMsg = String(describing: error).replacingOccurrences(of: "\"", with: "'")
@@ -1091,6 +1137,7 @@ struct ServerConfig: Sendable {
10911137
let turboKV: Bool
10921138
let mtp: Bool
10931139
let numMtpTokens: Int
1140+
let mtpAssistantModel: String?
10941141
}
10951142

10961143
// ── SSD Memory Budget ────────────────────────────────────────────────────────
@@ -1352,7 +1399,8 @@ func handleChatCompletion(
13521399
numDraftTokens: Int = 4,
13531400
dflashModel: DFlashDraftModel? = nil,
13541401
dflashBlockSize: Int? = nil,
1355-
dflashTargetModel: (any DFlashTargetModel)? = nil
1402+
dflashTargetModel: (any DFlashTargetModel)? = nil,
1403+
mtpAssistant: (any DualModelMTP)? = nil
13561404
) async throws -> Response {
13571405
let chatReq = try JSONDecoder().decode(ChatCompletionRequest.self, from: bodyData)
13581406
let isStream = chatReq.stream ?? false
@@ -1626,9 +1674,9 @@ func handleChatCompletion(
16261674
}
16271675
let remainingTokens = lmInput.text.tokens[startIndex...]
16281676
let trimmedInput = LMInput(tokens: remainingTokens)
1629-
if config.mtp, context.model is any MTPLanguageModel {
1677+
if config.mtp, let mtpCtx = mtpContext(main: context, assistant: mtpAssistant) {
16301678
stream = try MLXLMCommon.generateMTP(
1631-
input: trimmedInput, cache: cache, parameters: params, context: context, numMTPTokens: config.numMtpTokens
1679+
input: trimmedInput, cache: cache, parameters: params, context: mtpCtx, numMTPTokens: config.numMtpTokens
16321680
)
16331681
} else {
16341682
stream = try MLXLMCommon.generate(
@@ -1637,9 +1685,9 @@ func handleChatCompletion(
16371685
}
16381686
} else {
16391687
// Cache miss: process the full prompt.
1640-
if config.mtp, context.model is any MTPLanguageModel {
1688+
if config.mtp, let mtpCtx = mtpContext(main: context, assistant: mtpAssistant) {
16411689
stream = try MLXLMCommon.generateMTP(
1642-
input: lmInput, cache: cache, parameters: params, context: context, numMTPTokens: config.numMtpTokens
1690+
input: lmInput, cache: cache, parameters: params, context: mtpCtx, numMTPTokens: config.numMtpTokens
16431691
)
16441692
} else {
16451693
stream = try MLXLMCommon.generate(
@@ -2714,6 +2762,29 @@ func pendingStopPrefixLength(_ text: String, stopSequences: [String]) -> Int {
27142762
return longest
27152763
}
27162764

2765+
/// The context `generateMTP` should run against, and whether MTP applies at all.
2766+
///
2767+
/// Two shapes exist. Qwen3.5 and DeepSeek V4 carry MTP heads inside the main checkpoint,
2768+
/// so the main context is already an `MTPLanguageModel` and is used as-is. Gemma 4 ships
2769+
/// the heads as a separate assistant checkpoint: there the *assistant* is the
2770+
/// `MTPLanguageModel`, so it is swapped into a derived context while the tokenizer,
2771+
/// processor and configuration — and the KV cache passed alongside — stay the trunk's.
2772+
/// That mirrors the reference usage in Gemma4MTPBench, which passes the assistant as the
2773+
/// model and the main model's cache.
2774+
///
2775+
/// Returns nil when MTP does not apply, so callers fall through to plain generation.
2776+
func mtpContext(main: ModelContext, assistant: (any DualModelMTP)?) -> ModelContext? {
2777+
if let assistant, let assistantModel = assistant as? (any LanguageModel) {
2778+
return ModelContext(
2779+
configuration: main.configuration,
2780+
model: assistantModel,
2781+
processor: main.processor,
2782+
tokenizer: main.tokenizer
2783+
)
2784+
}
2785+
return main.model is any MTPLanguageModel ? main : nil
2786+
}
2787+
27172788
/// Trims `text` at the earliest stop sequence it contains.
27182789
///
27192790
/// Earliest *in the text*, not first in the caller's list: returning whichever entry

‎mlx-swift-lm‎

‎tests/test-contract.sh‎

Lines changed: 31 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -210,6 +210,37 @@ else
210210
skip "node not available on this machine"
211211
fi
212212

213+
# ── 8. A prompt long enough to be prefilled in chunks ────────────────────────
214+
# Below prefillStepSize (512 tokens) the generator's prepare() returns the prompt
215+
# tokens without ever forwarding them, so a whole code path — chunked prefill —
216+
# goes unexercised. Every other prompt in this repo's tests is under 80 characters,
217+
# which is how a dual-model MTP crash on any real-sized prompt reached a merge
218+
# queue with green CI (SharpAI/mlx-swift-lm#46). This closes the gap for the
219+
# ordinary generate path only — CI runs no --mtp job, so the speculative variant
220+
# of the same path stays uncovered until one exists (#128).
221+
log "Test 8: prompt exceeding the prefill chunk size"
222+
LONG_PROMPT=$(python3 -c '
223+
# ~3000 tokens: comfortably past 512 even with an efficient tokenizer.
224+
print(("The quick brown fox jumps over the lazy dog near the river bank. " * 190) + "Reply with the single word: done.")')
225+
LONG_BODY=$(python3 -c '
226+
import json, sys
227+
print(json.dumps({"messages": [{"role": "user", "content": sys.argv[1]}],
228+
"max_tokens": 16, "stream": False}))' "$LONG_PROMPT")
229+
LONG_RESP=$(curl -sf --max-time 300 "$URL/v1/chat/completions" \
230+
-H 'Content-Type: application/json' -d "$LONG_BODY" 2>/dev/null || true)
231+
if [ -z "$LONG_RESP" ]; then
232+
# A crash in the prefill path kills the connection rather than returning an error
233+
# body, so an empty response is the signal we are actually looking for here.
234+
fail "no response to a chunk-prefilled prompt — server may have died; see /tmp/SwiftLM-test-contract.log"
235+
elif echo "$LONG_RESP" | python3 -c '
236+
import json, sys
237+
d = json.load(sys.stdin)
238+
sys.exit(0 if d["choices"][0]["message"]["content"].strip() else 1)' 2>/dev/null; then
239+
pass "chunk-prefilled prompt produced content"
240+
else
241+
fail "chunk-prefilled prompt returned no content: $(echo "$LONG_RESP" | head -c 120)"
242+
fi
243+
213244
log "═══════════════════════════════════════"
214245
log "Results: $PASS passed, $FAIL failed, $SKIP skipped, $TOTAL total"
215246
log "═══════════════════════════════════════"

0 commit comments

Comments
 (0)