From 691a5da746226ddc630ec1e1d419aa54e3cec7cb Mon Sep 17 00:00:00 2001 From: Joseph Magly <1159087+jmagly@users.noreply.github.com> Date: Thu, 27 Aug 2026 20:17:09 -0400 Subject: [PATCH] fix(qwen): recognize saved text derivatives --- obliteratus/models/loader.py | 8 ++++++-- tests/test_app_model_lifecycle.py | 10 ++++++---- tests/test_loader_boundaries.py | 15 +++++++++++++++ 3 files changed, 27 insertions(+), 6 deletions(-) diff --git a/obliteratus/models/loader.py b/obliteratus/models/loader.py index 4103813..c6eb0c4 100644 --- a/obliteratus/models/loader.py +++ b/obliteratus/models/loader.py @@ -666,7 +666,7 @@ def qwen_hybrid_runtime_overrides( contract. Generic Accelerate sharding is not valid for the recurrent DeltaNet execution path, even when aggregate VRAM is sufficient. """ - if getattr(config, "model_type", "") != "qwen3_5": + if getattr(config, "model_type", "") not in {"qwen3_5", "qwen3_5_text"}: return {} import transformers @@ -876,7 +876,11 @@ def load_model( if task == "classification": config.num_labels = num_labels load_kwargs["config"] = config - is_qwen_hybrid = task == "causal_lm" and getattr(config, "model_type", "") == "qwen3_5" + is_qwen_hybrid = task == "causal_lm" and getattr( + config, + "model_type", + "", + ) in {"qwen3_5", "qwen3_5_text"} qwen_device_index = None if is_qwen_hybrid: qwen_overrides = qwen_hybrid_runtime_overrides( diff --git a/tests/test_app_model_lifecycle.py b/tests/test_app_model_lifecycle.py index f36d6f7..408a103 100644 --- a/tests/test_app_model_lifecycle.py +++ b/tests/test_app_model_lifecycle.py @@ -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( diff --git a/tests/test_loader_boundaries.py b/tests/test_loader_boundaries.py index 831f713..d63140d 100644 --- a/tests/test_loader_boundaries.py +++ b/tests/test_loader_boundaries.py @@ -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(