fix(qwen): recognize saved text derivatives

This commit is contained in:
Joseph Magly
2026-08-27 20:17:09 -04:00
parent 44f7842c9c
commit 691a5da746
3 changed files with 27 additions and 6 deletions
+6 -4
View File
@@ -113,15 +113,17 @@ app._load_model_to_device(checkpoint, local_files_only=True)
assert calls[0][0] == checkpoint
assert calls[0][1]["local_files_only"] is True
loader.qwen_hybrid_runtime_overrides = Mock(
return_value={"attn_implementation": "sdpa", "device_map": {"": 2}},
app.AutoConfig.from_pretrained.return_value = SimpleNamespace(
model_type="qwen3_5_text",
)
loader._require_qwen_hybrid_kernels = Mock()
loader._estimate_model_memory_gb = Mock(return_value=54.0)
loader._qwen_single_device = Mock(return_value=2)
app._load_model_to_device(checkpoint, local_files_only=True)
assert calls[-1][1]["attn_implementation"] == "sdpa"
assert calls[-1][1]["device_map"] == {"": 2}
assert calls[-1][1]["device_map"] != "auto"
loader.qwen_hybrid_runtime_overrides.reset_mock(return_value=True)
loader.qwen_hybrid_runtime_overrides.return_value = {}
app.AutoConfig.from_pretrained.return_value = SimpleNamespace(model_type="gpt2")
quantization = object()
loaded_model, loaded_tokenizer = app._reload_local_checkpoint(
+15
View File
@@ -573,6 +573,21 @@ def test_qwen_single_device_uses_per_device_free_memory(monkeypatch):
assert loader._qwen_single_device(216.0, "4bit") == 1
def test_qwen_saved_text_derivative_reuses_single_device_runtime(monkeypatch):
config = _config(model_type="qwen3_5_text")
monkeypatch.setattr(loader, "_require_qwen_hybrid_kernels", Mock())
monkeypatch.setattr(loader, "_estimate_model_memory_gb", Mock(return_value=54.0))
monkeypatch.setattr(loader, "_qwen_single_device", Mock(return_value=2))
assert loader.qwen_hybrid_runtime_overrides(
config,
torch.bfloat16,
) == {
"attn_implementation": "sdpa",
"device_map": {"": 2},
}
def test_model_handle_metadata_snapshot_restore_summary_and_cleanup(tmp_path):
model = _model()
nested = SimpleNamespace(