mirror of
https://github.com/elder-plinius/OBLITERATUS.git
synced 2026-08-30 06:30:37 +02:00
fix(qwen): preserve validated runtime on checkpoint reload
This commit is contained in:
@@ -82,9 +82,11 @@ import pathlib
|
||||
import sys
|
||||
import threading
|
||||
import time
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import Mock
|
||||
|
||||
import app
|
||||
from obliteratus.models import loader
|
||||
|
||||
root = pathlib.Path(sys.argv[1])
|
||||
checkpoint = root / "completed"
|
||||
@@ -103,11 +105,24 @@ app.AutoModelForCausalLM.from_pretrained = Mock(
|
||||
side_effect=lambda source, **kwargs: (calls.append((source, kwargs)), model)[1]
|
||||
)
|
||||
app.AutoTokenizer.from_pretrained = Mock(return_value=tokenizer)
|
||||
app.AutoConfig.from_pretrained = Mock(
|
||||
return_value=SimpleNamespace(model_type="gpt2"),
|
||||
)
|
||||
app.dev.supports_device_map_auto = lambda: True
|
||||
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._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 = {}
|
||||
|
||||
quantization = object()
|
||||
loaded_model, loaded_tokenizer = app._reload_local_checkpoint(
|
||||
checkpoint,
|
||||
|
||||
@@ -216,6 +216,10 @@ def test_qwen_excise_route_does_not_touch_inputs_gates_or_lm_head():
|
||||
torch.equal(after[name], before[name])
|
||||
for name in before.keys() - allowed
|
||||
)
|
||||
assert pipeline._effective_refinement_passes == 1
|
||||
method_config = pipeline._build_metadata()["method_config"]
|
||||
assert method_config["refinement_passes"] == 1
|
||||
assert method_config["requested_refinement_passes"] == 2
|
||||
|
||||
|
||||
@pytest.mark.parametrize("architecture", ["qwen3_5_text", "qwen3_5_moe"])
|
||||
|
||||
Reference in New Issue
Block a user