Skip to content

Commit 398252e

Browse files
simbaclaude
authored andcommitted
fix(load): only fall back to text-only for VLM checkpoint mismatches
Review on #173: `catch where !self.vision` caught every error, so a download failure or a CancellationError triggered a second, LLM load attempt and then surfaced a confusing error. The fallback now requires isVLMCheckpointMismatch(error): a DecodingError, an MLXNN UpdateError (unhandledKeys, keyNotFound, mismatchedSize, ...), or a ModelFactoryError for an unsupported model/processor type or an undecodable/invalid config. Every other error propagates unchanged. VLMFallbackTests covers the fallback cases (missing image_mean, unhandled weight keys, unsupported model type) and the pass-through cases (cancellation, URLError, missing config file). Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
1 parent 5eecf22 commit 398252e

2 files changed

Lines changed: 72 additions & 5 deletions

File tree

‎Sources/SwiftLM/Server.swift‎

Lines changed: 28 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -205,6 +205,28 @@ private struct TransformersTokenizerBridge: MLXLMCommon.Tokenizer, Sendable {
205205

206206
/// Returns `nil` when the value must be dropped (JSON `null` / NSNull), otherwise a
207207
/// structure with every nested null removed. See `TransformersTokenizerBridge.applyChatTemplate`.
208+
/// True when a VLM load failed because the checkpoint doesn't match the VLM code:
209+
/// its config doesn't decode, its weights don't line up with the module tree, or the
210+
/// factory doesn't know the model/processor type. Only these justify retrying an
211+
/// auto-detected VLM as a text-only LLM. Anything else (cancellation, download, I/O)
212+
/// would fail the same way again and just hide the real error.
213+
func isVLMCheckpointMismatch(_ error: any Error) -> Bool {
214+
switch error {
215+
case is DecodingError, is UpdateError:
216+
return true
217+
case let factoryError as ModelFactoryError:
218+
switch factoryError {
219+
case .unsupportedModelType, .unsupportedProcessorType, .configurationDecodingError,
220+
.invalidConfiguration:
221+
return true
222+
default:
223+
return false
224+
}
225+
default:
226+
return false
227+
}
228+
}
229+
208230
func sanitizeForJinja(_ value: any Sendable) -> (any Sendable)? {
209231
if value is NSNull { return nil }
210232
let mirror = Mirror(reflecting: value)
@@ -904,11 +926,12 @@ struct MLXServer: AsyncParsableCommand {
904926
) { progress in
905927
tracker.printProgress(progress)
906928
}
907-
} catch where !self.vision {
908-
// Vision was only auto-detected. A checkpoint whose vision side
909-
// doesn't load (e.g. a preprocessor_config.json without
910-
// image_mean) can still serve text, so fall back rather than exit.
911-
// With an explicit --vision, the error still propagates.
929+
} catch where !self.vision && isVLMCheckpointMismatch(error) {
930+
// Vision was only auto-detected, and the vision side of the checkpoint
931+
// doesn't match the VLM code (e.g. a preprocessor_config.json without
932+
// image_mean). The text model can still serve, so fall back rather than
933+
// exit. Cancellation, network and I/O errors still propagate, as does
934+
// any error under an explicit --vision.
912935
print("[SwiftLM] ⚠️ Auto-detected VLM failed to load (\(error)); loading as a text-only LLM. Pass --vision to make this fatal.")
913936
isVision = false
914937
container = try await LLMModelFactory.shared.loadContainer(
Lines changed: 44 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,44 @@
1+
import XCTest
2+
import Foundation
3+
import MLXLMCommon
4+
import MLXNN
5+
@testable import SwiftLM
6+
7+
/// An auto-detected VLM that fails to load falls back to text-only only for
8+
/// checkpoint mismatches. Other failures must surface as themselves.
9+
final class VLMFallbackTests: XCTestCase {
10+
11+
private struct Probe: Decodable { let image_mean: [Double] }
12+
13+
func testConfigDecodingErrorFallsBack() {
14+
// The real case: unsloth/Qwen3.6-35B-A3B's preprocessor_config.json has no image_mean.
15+
do {
16+
_ = try JSONDecoder().decode(Probe.self, from: Data("{}".utf8))
17+
XCTFail("expected a DecodingError")
18+
} catch {
19+
XCTAssertTrue(isVLMCheckpointMismatch(error))
20+
}
21+
}
22+
23+
func testWeightMismatchFallsBack() {
24+
let error = UpdateError.unhandledKeys(path: [], modules: ["Vision"], keys: ["pre_projection"])
25+
XCTAssertTrue(isVLMCheckpointMismatch(error))
26+
}
27+
28+
func testUnsupportedModelTypeFallsBack() {
29+
XCTAssertTrue(isVLMCheckpointMismatch(ModelFactoryError.unsupportedModelType("qwen4_exp")))
30+
}
31+
32+
func testCancellationDoesNotFallBack() {
33+
XCTAssertFalse(isVLMCheckpointMismatch(CancellationError()))
34+
}
35+
36+
func testNetworkErrorDoesNotFallBack() {
37+
XCTAssertFalse(isVLMCheckpointMismatch(URLError(.notConnectedToInternet)))
38+
}
39+
40+
func testMissingConfigFileDoesNotFallBack() {
41+
let io = CocoaError(.fileReadNoSuchFile)
42+
XCTAssertFalse(isVLMCheckpointMismatch(ModelFactoryError.configurationFileError("config.json", "m", io)))
43+
}
44+
}

0 commit comments

Comments
 (0)