test: enforce model and CLI boundary contracts

This commit is contained in:
Joseph Magly
2026-08-14 12:51:09 -04:00
parent 761ee6127c
commit 6add02edd0
13 changed files with 1492 additions and 52 deletions
+460
View File
@@ -0,0 +1,460 @@
"""Behavioral CLI boundaries without network, accelerators, or model downloads."""
from __future__ import annotations
import json
import os
from pathlib import Path
from types import SimpleNamespace
from unittest.mock import MagicMock, Mock
import pytest
from obliteratus import cli
def ns(**values):
return SimpleNamespace(**values)
@pytest.mark.parametrize(
("argv", "target"),
[
(["gpu-calc", "--params", "1"], "_cmd_gpu_calc"),
(["run", "config.yml"], "_cmd_run"),
(["interactive"], "_cmd_interactive"),
(["models"], "_cmd_models"),
(["presets"], "_cmd_presets"),
(["info", "local/model"], "_cmd_info"),
(["strategies"], "_cmd_strategies"),
(["report", "results.json"], "_cmd_report"),
(["aggregate"], "_cmd_aggregate"),
(["ui"], "_cmd_ui"),
(["recommend", "local/model"], "_cmd_recommend"),
(["tourney", "local/model"], "_cmd_tourney"),
(["self-improve", "local/model", "--audit", "audit.json", "--output-dir", "out"], "_cmd_self_improve"),
(["abliterate", "local/model"], "_cmd_abliterate"),
],
)
def test_main_routes_local_commands(monkeypatch, argv, target):
command = Mock()
monkeypatch.setattr(cli, target, command)
cli.main(argv)
command.assert_called_once()
@pytest.mark.parametrize(
("argv", "target"),
[
(["run", "config.yml", "--remote", "gpu"], "_cmd_remote_run"),
(["tourney", "model", "--remote", "gpu"], "_cmd_remote_tourney"),
(["obliterate", "model", "--remote", "gpu"], "_cmd_remote_abliterate"),
],
)
def test_main_routes_remote_commands(monkeypatch, argv, target):
command = Mock()
monkeypatch.setattr(cli, target, command)
cli.main(argv)
command.assert_called_once()
def test_version_is_stable_and_does_not_dispatch(capsys):
from obliteratus import __version__
with pytest.raises(SystemExit) as exc:
cli.main(["--version"])
assert exc.value.code == 0
assert capsys.readouterr().out.endswith(f"obliteratus {__version__}\n")
def test_gpu_selection_contract(monkeypatch):
monkeypatch.delenv("CUDA_VISIBLE_DEVICES", raising=False)
cli._apply_gpu_selection(ns(gpus=None, remote=None))
cli._apply_gpu_selection(ns(gpus="all", remote=None))
cli._apply_gpu_selection(ns(gpus="9", remote="host"))
assert "CUDA_VISIBLE_DEVICES" not in os.environ
cli._apply_gpu_selection(ns(gpus=" 2, 0 ", remote=None))
assert os.environ["CUDA_VISIBLE_DEVICES"] == "2,0"
with pytest.raises(SystemExit) as exc:
cli._apply_gpu_selection(ns(gpus="gpu0", remote=None))
assert exc.value.code == 1
def test_ui_and_interactive_forwarding(monkeypatch):
import obliteratus.interactive
import obliteratus.local_ui
launch = Mock()
interactive = Mock()
monkeypatch.setattr(obliteratus.local_ui, "launch_local_ui", launch)
monkeypatch.setattr(obliteratus.interactive, "run_interactive", interactive)
cli._cmd_ui(ns(auth="user:pass", host="127.0.0.1", port=9000, share=True, no_browser=True, quiet=True))
launch.assert_called_once_with(
host="127.0.0.1",
port=9000,
share=True,
open_browser=False,
auth=("user", "pass"),
quiet=True,
)
cli._cmd_interactive()
interactive.assert_called_once_with()
def test_models_presets_and_strategies_render(monkeypatch):
import obliteratus.presets
import obliteratus.strategies
import obliteratus.study_presets
model = ns(
name="Tiny",
hf_id="local/tiny",
params="1B",
tier="tiny",
recommended_dtype="float32",
recommended_quantization=None,
description="fixture",
)
preset = ns(key="quick", name="Quick", strategies=[{"name": "heads"}], max_samples=2, description="fixture")
strategy = type("Strategy", (), {"__doc__": "First line.\nSecond line."})
monkeypatch.setattr(obliteratus.presets, "list_all_presets", lambda: [model])
monkeypatch.setattr(obliteratus.presets, "get_presets_by_tier", lambda _tier: [model])
monkeypatch.setattr(obliteratus.study_presets, "list_study_presets", lambda: [preset])
monkeypatch.setattr(obliteratus.strategies, "STRATEGY_REGISTRY", {"fixture": strategy})
console = Mock()
monkeypatch.setattr(cli, "console", console)
cli._cmd_models(ns(tier=None))
cli._cmd_models(ns(tier="tiny"))
cli._cmd_presets()
cli._cmd_strategies()
assert console.print.call_count >= 8
def test_run_local_overrides_and_remote_config(monkeypatch, tmp_path):
import obliteratus.config
import obliteratus.remote
import obliteratus.runner
config_path = tmp_path / "study.yml"
config_path.write_text("model: {name: fixture}\n")
config = ns(remote=None, output_dir="old")
monkeypatch.setattr(obliteratus.config.StudyConfig, "from_yaml", Mock(return_value=config))
monkeypatch.setattr(obliteratus.config.StudyConfig, "from_dict", Mock(return_value=config))
run = Mock()
monkeypatch.setattr(obliteratus.runner, "run_study", run)
cli._cmd_run(ns(config=str(config_path), preset="quick", output_dir="new"))
assert config.output_dir == "new"
run.assert_called_once_with(config)
config.remote = ns(
host="host", user="user", port=2200, ssh_key="key", remote_dir="/work",
python="python", sync_results=True, gpus="0",
)
runner = MagicMock()
runner.run_config.return_value = "/local/results"
monkeypatch.setattr(obliteratus.remote, "RemoteRunner", Mock(return_value=runner))
cli._cmd_run(ns(config=str(config_path), preset=None, output_dir=None))
runner.run_config.assert_called_once()
runner.run_config.return_value = None
with pytest.raises(SystemExit) as exc:
cli._cmd_run(ns(config=str(config_path), preset=None, output_dir=None))
assert exc.value.code == 1
def test_info_prints_stable_summary(monkeypatch):
import obliteratus.models.loader
handle = ns(summary=lambda: {"total_params": 1234, "architecture": "fixture"})
load = Mock(return_value=handle)
monkeypatch.setattr(obliteratus.models.loader, "load_model", load)
console = Mock()
monkeypatch.setattr(cli, "console", console)
cli._cmd_info(ns(model="local/model", task="causal_lm", device="cpu", dtype="float32"))
load.assert_called_once_with(model_name="local/model", task="causal_lm", device="cpu", dtype="float32")
assert console.print.call_count == 3
def test_report_success_and_plot_failure(monkeypatch, tmp_path):
import obliteratus.reporting.report
path = tmp_path / "results.json"
path.write_text(json.dumps({
"model_name": "fixture",
"baseline_metrics": {"score": 1.0},
"results": [{
"strategy": "heads", "component": "0", "description": "fixture",
"metrics": {"score": 0.5}, "metadata": {"source": "test"},
}],
}))
report = MagicMock()
monkeypatch.setattr(obliteratus.reporting.report, "AblationReport", Mock(return_value=report))
cli._cmd_report(ns(results_json=str(path), output_dir=str(tmp_path / "plots")))
report.add_baseline.assert_called_once_with({"score": 1.0})
report.add_result.assert_called_once()
report.plot_impact.side_effect = RuntimeError("no renderer")
cli._cmd_report(ns(results_json=str(path), output_dir=None))
def test_aggregate_formats_filters_and_empty_result(monkeypatch):
import obliteratus.community
records = [object()]
aggregated = {
"org/model": {
"keep": {"n_runs": 2, "refusal_rate": {"mean": 0.1}, "perplexity": {"mean": 4.5}},
"drop": {"n_runs": 1},
},
}
monkeypatch.setattr(obliteratus.community, "load_contributions", lambda _dir: records)
monkeypatch.setattr(obliteratus.community, "aggregate_results", lambda _records: json.loads(json.dumps(aggregated)))
latex = Mock(return_value="TABLE")
monkeypatch.setattr(obliteratus.community, "generate_latex_table", latex)
cli._cmd_aggregate(ns(dir="results", min_runs=2, format="summary", metric="refusal_rate"))
cli._cmd_aggregate(ns(dir="results", min_runs=1, format="latex", metric="score"))
latex.assert_called_once()
monkeypatch.setattr(obliteratus.community, "aggregate_results", lambda _records: {"model": {"one": {"n_runs": 1}}})
cli._cmd_aggregate(ns(dir="results", min_runs=3, format="summary", metric="score"))
def test_recommend_handles_config_fallback_telemetry_and_insights(monkeypatch):
import obliteratus.adaptive_defaults
import obliteratus.architecture_profiles
import transformers
profile = ns(
profile_label="Fixture", arch_class=ns(value="dense"), reasoning_class=ns(value="general"),
total_params_b=1.0, num_layers=2, hidden_size=8, recommended_method="advanced",
method_overrides={"n_directions": 2},
)
monkeypatch.setattr(transformers.AutoConfig, "from_pretrained", Mock(return_value=ns(num_hidden_layers=2, hidden_size=8)))
detect = Mock(return_value=profile)
monkeypatch.setattr(obliteratus.architecture_profiles, "detect_architecture", detect)
monkeypatch.setattr(obliteratus.architecture_profiles, "enhance_profile_with_telemetry", lambda p: (p, {"method": "advanced"}))
monkeypatch.setattr(obliteratus.adaptive_defaults, "format_recommendation", lambda _rec: "recommendation")
monkeypatch.setattr(obliteratus.adaptive_defaults, "get_global_insights", lambda: {
"total_records": 3,
"overall_best_methods": [{"method": "advanced", "mean_score": 0.8, "n_runs": 3}],
"architecture_breakdown": {"dense": {"best_method": "advanced", "best_score": 0.8, "n_methods_tested": 2, "total_runs": 3}},
})
cli._cmd_recommend(ns(model="local/model", insights=True))
detect.assert_called_once()
transformers.AutoConfig.from_pretrained.side_effect = OSError("offline")
monkeypatch.setattr(obliteratus.architecture_profiles, "enhance_profile_with_telemetry", lambda p: (p, None))
cli._cmd_recommend(ns(model="local/model", insights=False))
def test_tourney_callbacks_and_winner(monkeypatch):
import obliteratus.tourney
result = ns(
winner=ns(method="advanced", score=0.9, metrics={"refusal_rate": 0.1, "coherence": 0.8}),
hub_repo="org/winner",
)
runner = MagicMock()
runner.run.return_value = result
factory = Mock(return_value=runner)
monkeypatch.setattr(obliteratus.tourney, "TourneyRunner", factory)
cli._cmd_tourney(ns(
model="model", hub_org="org", hub_repo=None, device="cpu", dtype="float32",
dataset="builtin", quantization=None, methods=["advanced"], output_dir="out",
))
callbacks = factory.call_args.kwargs
callbacks["on_log"]("message")
callbacks["on_round"](ns(round_num=1, advanced_to=[1], eliminated=[2]))
def test_gpu_calculator_validation_profile_and_moe(monkeypatch):
import obliteratus.model_profile
console = Mock()
monkeypatch.setattr(cli, "console", console)
cli._cmd_gpu_calc(ns(model=None, params=10.0, active_params=2.0, dtype="float16", gpu_mem=24.0))
profile = ns(total_params_b=7.0, active_params_b=3.0)
monkeypatch.setattr(obliteratus.model_profile, "profile_model", Mock(return_value=profile))
cli._cmd_gpu_calc(ns(model="local/model", params=None, active_params=None, dtype="int8", gpu_mem=16.0))
with pytest.raises(SystemExit):
cli._cmd_gpu_calc(ns(model=None, params=None, active_params=None, dtype="float16", gpu_mem=16.0))
with pytest.raises(SystemExit):
cli._cmd_gpu_calc(ns(model=None, params=1.0, active_params=None, dtype="float16", gpu_mem=1.0))
monkeypatch.setattr(obliteratus.model_profile, "profile_model", Mock(side_effect=OSError("offline")))
with pytest.raises(SystemExit):
cli._cmd_gpu_calc(ns(model="x", params=None, active_params=None, dtype="float16", gpu_mem=16.0))
def test_parameter_estimators_cover_dense_moe_and_invalid():
assert cli._estimate_total_params_b(ns(num_parameters=2_000_000_000)) == 2.0
dense = ns(hidden_size=128, num_hidden_layers=2, vocab_size=1000, intermediate_size=512)
total = cli._estimate_total_params_b(dense)
assert total > 0
assert cli._estimate_active_params_b(dense, total) == total
moe = ns(
hidden_size=4096, num_hidden_layers=32, vocab_size=100_000, intermediate_size=14_336,
moe_intermediate_size=2048, num_local_experts=8, num_experts_per_tok=2,
)
moe_total = cli._estimate_total_params_b(moe)
assert 0.1 <= cli._estimate_active_params_b(moe, moe_total) < moe_total
with pytest.raises(SystemExit):
cli._estimate_total_params_b(ns(hidden_size=0, num_hidden_layers=0, vocab_size=0))
def _remote_args(**overrides):
values = {
"remote": "user@host", "ssh_port": 22, "ssh_key": None, "remote_dir": "/work",
"remote_python": "python3", "no_sync": False, "gpus": "0", "model": "model",
"output_dir": "out", "method": "advanced", "device": "cuda", "dtype": "float16",
"quantization": "4bit", "n_directions": 2, "direction_method": "svd",
"regularization": 0.2, "refinement_passes": 2, "min_layer_fraction": 0.1,
"max_layer_fraction": 0.9, "harmless_pc_count": 1, "shield_concept_count": 2,
"shield_ridge": 0.1, "shield_residualize": True, "shield_layer_penalty": 0.2,
"projection_target": "all", "projection_row_fraction": 0.5, "large_model": True,
"verify_sample_size": 5, "config": "config.yml", "preset": "quick", "methods": ["advanced"],
"hub_org": "org", "hub_repo": None, "dataset": "builtin",
}
values.update(overrides)
return ns(**values)
def test_remote_runner_factory_and_commands(monkeypatch):
import obliteratus.remote
config = object()
runner = MagicMock()
monkeypatch.setattr(obliteratus.remote.RemoteConfig, "from_cli_args", Mock(return_value=config))
monkeypatch.setattr(obliteratus.remote, "RemoteRunner", Mock(return_value=runner))
args = _remote_args()
assert cli._make_remote_runner(args) is runner
monkeypatch.setattr(cli, "_make_remote_runner", lambda _args: runner)
runner.run_obliterate.return_value = "results"
cli._cmd_remote_abliterate(args)
assert runner.run_obliterate.call_args.kwargs["projection_row_fraction"] == 0.5
runner.run_config.return_value = "results"
cli._cmd_remote_run(args)
runner.run_tourney.return_value = "results"
cli._cmd_remote_tourney(args)
runner.run_obliterate.return_value = None
with pytest.raises(SystemExit):
cli._cmd_remote_abliterate(args)
runner.run_config.return_value = None
with pytest.raises(SystemExit):
cli._cmd_remote_run(args)
runner.run_tourney.return_value = None
with pytest.raises(SystemExit):
cli._cmd_remote_tourney(args)
def test_abliterate_pipeline_callbacks_residue_and_contribution(monkeypatch, tmp_path):
import obliteratus.abliterate
import obliteratus.community
import obliteratus.hard_negative
import obliteratus.telemetry
import rich.live
stages = [ns(key=f"s{i}", name=f"Stage {i}") for i in range(6)]
monkeypatch.setattr(obliteratus.abliterate, "STAGES", stages)
monkeypatch.setattr(obliteratus.abliterate, "METHODS", {"advanced": {"label": "Advanced"}})
monkeypatch.setattr(
obliteratus.hard_negative,
"build_weighted_prompt_pairs",
lambda **_kwargs: (["harm"], ["safe"], {"residue_examples": 1, "residue_added_pairs": 2, "total_pairs": 3}),
)
monkeypatch.setattr(obliteratus.community, "save_contribution", lambda *_args, **_kwargs: "contribution.json")
telemetry = Mock()
monkeypatch.setattr(obliteratus.telemetry, "maybe_send_pipeline_report", telemetry)
result_path = tmp_path / "result"
result_path.mkdir()
pipeline = MagicMock()
def run():
kwargs = factory.call_args.kwargs
kwargs["on_log"]("working")
kwargs["on_stage"](ns(stage="s0", status="running", message="work"))
kwargs["on_stage"](ns(stage="s0", status="done", message="done"))
return str(result_path)
pipeline.run.side_effect = run
factory = Mock(return_value=pipeline)
monkeypatch.setattr(obliteratus.abliterate, "AbliterationPipeline", factory)
class FakeLive:
def __init__(self, *_args, **_kwargs):
self.update = Mock()
def __enter__(self):
return self
def __exit__(self, *_args):
return False
monkeypatch.setattr(rich.live, "Live", FakeLive)
args = ns(
model="org/model", output_dir=None, device="cpu", dtype="float32", method="advanced",
n_directions=2, direction_method="svd", regularization=0.2, refinement_passes=1,
min_layer_fraction=0.1, max_layer_fraction=0.9, harmless_pc_count=1,
shield_concept_count=2, shield_ridge=0.1, shield_residualize=True,
shield_layer_penalty=0.2, projection_target="all", projection_row_fraction=0.5,
quantization=None, large_model=False, verify_sample_size=3,
residue_file=["audit.json"], dataset="builtin", residue_weight=2, residue_max=4,
contribute=True, contribute_notes="fixture",
)
cli._cmd_abliterate(args)
assert (result_path / "hard_negative_residue.json").is_file()
telemetry.assert_called_once_with(pipeline)
def test_self_improve_dry_run_and_pipeline(monkeypatch, tmp_path):
import obliteratus.abliterate
import obliteratus.hard_negative
import obliteratus.model_profile
profile = ns(
total_params_b=1.0, total_params=1_000_000_000, active_params_b=1.0,
source="fixture", num_layers=2, hidden_size=8,
to_json=lambda: {
"model_id": "model", "total_params": 1_000_000_000, "total_params_b": 1.0,
"active_params_b": 1.0, "num_layers": 2, "hidden_size": 8, "source": "fixture",
},
)
monkeypatch.setattr(obliteratus.model_profile, "profile_model", lambda *_args, **_kwargs: profile)
monkeypatch.setattr(obliteratus.model_profile, "default_self_improve_params", lambda _profile: {
"n_directions": 2, "regularization": 0.2, "refinement_passes": 1,
"residue_weight": 3, "verify_sample_size": 5, "note": "fixture",
})
monkeypatch.setattr(obliteratus.hard_negative, "load_residue_file", lambda _path: ["residue"])
def save_residue(_items, path):
path = Path(path)
path.parent.mkdir(parents=True, exist_ok=True)
path.write_text("[]")
monkeypatch.setattr(obliteratus.hard_negative, "save_residue_file", save_residue)
monkeypatch.setattr(obliteratus.hard_negative, "build_weighted_prompt_pairs", lambda **_kwargs: (
["harm"], ["safe"], {"residue_examples": 1, "residue_added_pairs": 3, "total_pairs": 4},
))
output = tmp_path / "out"
args = ns(
model="model", dtype="float32", params_b=None, no_param_auto_scale=False,
n_directions=None, regularization=None, refinement_passes=None, residue_weight=None,
verify_sample_size=None, audit=["audit.json"], residue_out=None, output_dir=str(output),
dataset="builtin", residue_max=None, projection_row_fraction=None, method="advanced",
direction_method="svd", min_layer_fraction=None, max_layer_fraction=None,
harmless_pc_count=None, shield_concept_count=None, shield_ridge=None,
shield_residualize=None, shield_layer_penalty=None, projection_target=None,
device="cpu", dry_run=True,
)
cli._cmd_self_improve(args)
assert (output / "self_improve_plan.json").is_file()
result = tmp_path / "candidate"
result.mkdir()
pipeline = MagicMock()
pipeline.run.return_value = str(result)
monkeypatch.setattr(obliteratus.abliterate, "AbliterationPipeline", Mock(return_value=pipeline))
args.dry_run = False
cli._cmd_self_improve(args)
assert (result / "hard_negative_residue.json").is_file()
+62 -7
View File
@@ -13,35 +13,42 @@ MODULE = importlib.util.module_from_spec(SPEC)
SPEC.loader.exec_module(MODULE)
def _report(line: float = 49.0, branch: float = 36.0) -> dict[str, object]:
def _report(line: float = 55.0, branch: float = 42.0) -> dict[str, object]:
return {
"totals": {
"percent_statements_covered": line,
"percent_branches_covered": branch,
},
"files": {
"obliteratus/example.py": {
"summary": {"percent_statements_covered": 70.0},
"executed_lines": [1, 2, 4],
"missing_lines": [3],
},
},
}
def test_validate_coverage_accepts_exact_floors():
assert MODULE.validate_coverage(
_report(), min_line=49.0, min_branch=36.0,
_report(), min_line=55.0, min_branch=42.0,
) == []
def test_validate_coverage_reports_each_regression():
failures = MODULE.validate_coverage(
_report(line=48.9, branch=35.9), min_line=49.0, min_branch=36.0,
_report(line=54.9, branch=41.9), min_line=55.0, min_branch=42.0,
)
assert failures == [
"line coverage 48.90% is below the 49.00% floor",
"branch coverage 35.90% is below the 36.00% floor",
"line coverage 54.90% is below the 55.00% floor",
"branch coverage 41.90% is below the 42.00% floor",
]
def test_validate_coverage_rejects_malformed_totals():
assert MODULE.validate_coverage(
{}, min_line=49.0, min_branch=36.0,
{}, min_line=55.0, min_branch=42.0,
) == ["coverage report is missing the totals object"]
@@ -51,8 +58,56 @@ def test_validate_coverage_rejects_non_numeric_metrics():
report["totals"]["percent_branches_covered"] = float("nan")
assert MODULE.validate_coverage(
report, min_line=49.0, min_branch=36.0,
report, min_line=55.0, min_branch=42.0,
) == [
"coverage report is missing numeric percent_statements_covered",
"coverage report is missing numeric percent_branches_covered",
]
def test_validate_coverage_enforces_critical_file_floors():
assert MODULE.validate_coverage(
_report(),
min_line=55.0,
min_branch=42.0,
file_floors={"obliteratus/example.py": 70.0},
) == []
report = _report()
report["files"]["obliteratus/example.py"]["summary"]["percent_statements_covered"] = 69.9
assert MODULE.validate_coverage(
report,
min_line=55.0,
min_branch=42.0,
file_floors={"obliteratus/example.py": 70.0, "missing.py": 70.0},
) == [
"critical file obliteratus/example.py coverage 69.90% is below the 70.00% floor",
"coverage report is missing critical file missing.py",
]
def test_parse_changed_lines_and_measurement_ignore_non_executable_lines():
diff = """diff --git a/obliteratus/example.py b/obliteratus/example.py
+++ b/obliteratus/example.py
@@ -1,2 +1,4 @@
diff --git a/tests/test_example.py b/tests/test_example.py
+++ b/tests/test_example.py
@@ -0,0 +1,2 @@
"""
changed = MODULE.parse_changed_lines(diff)
assert changed == {
"obliteratus/example.py": {1, 2, 3, 4},
"tests/test_example.py": {1, 2},
}
assert MODULE.changed_line_coverage(_report(), changed) == (3, 4, 75.0)
assert MODULE.validate_changed_coverage(
_report(), changed, minimum=90.0,
) == ["changed-line coverage 75.00% (3/4) is below the 90.00% floor"]
def test_changed_line_gate_passes_when_diff_has_no_measured_source():
assert MODULE.changed_line_coverage(_report(), {"tests/test_example.py": {1}}) == (
0,
0,
100.0,
)
+209
View File
@@ -0,0 +1,209 @@
"""Deterministic contracts for accelerator detection and fallback behavior."""
from __future__ import annotations
from types import SimpleNamespace
from unittest.mock import Mock
import pytest
import torch
from obliteratus import device
@pytest.mark.parametrize(
("cuda", "mps", "expected"),
[(True, False, "cuda"), (False, True, "mps"), (False, False, "cpu")],
)
def test_auto_device_preference(monkeypatch, cuda, mps, expected):
monkeypatch.setattr(device, "is_cuda", lambda: cuda)
monkeypatch.setattr(device, "is_mps", lambda: mps)
assert device.get_device() == expected
assert device.is_gpu_available() is (cuda or mps)
def test_explicit_device_validation(monkeypatch):
monkeypatch.setattr(device, "is_cuda", lambda: False)
monkeypatch.setattr(device, "is_mps", lambda: False)
assert device.get_device("cpu") == "cpu"
with pytest.raises(RuntimeError, match="CUDA was requested.*device='auto'"):
device.get_device("cuda:0")
with pytest.raises(RuntimeError, match="MPS was requested.*device='auto'"):
device.get_device("mps")
with pytest.raises(ValueError, match="Unknown device 'tpu'"):
device.get_device("tpu")
with pytest.raises(ValueError, match="Unknown device 'cuda:gpu'"):
device.get_device("cuda:gpu")
monkeypatch.setattr(device, "is_cuda", lambda: True)
assert device.get_device("cuda") == "cuda"
assert device.get_device("cuda:3") == "cuda:3"
monkeypatch.setattr(device, "is_mps", lambda: True)
assert device.get_device("mps") == "mps"
def test_names_and_device_counts(monkeypatch):
monkeypatch.setattr(device, "is_cuda", lambda: True)
monkeypatch.setattr(device.torch.cuda, "get_device_name", lambda _index: "Test GPU")
monkeypatch.setattr(device.torch.cuda, "device_count", lambda: 4)
assert device.get_device_name() == "Test GPU"
assert device.device_count() == 4
monkeypatch.setattr(device, "is_cuda", lambda: False)
monkeypatch.setattr(device, "is_mps", lambda: True)
monkeypatch.setattr(device.platform, "processor", lambda: "M3")
assert device.get_device_name() == "Apple M3 (MPS)"
assert device.device_count() == 1
monkeypatch.setattr(device, "is_mps", lambda: False)
assert device.get_device_name() == "CPU"
assert device.device_count() == 0
def test_system_memory_sources_and_fallback(monkeypatch):
gib = 1024**3
fake_psutil = SimpleNamespace(
virtual_memory=lambda: SimpleNamespace(total=32 * gib, available=12 * gib),
)
monkeypatch.setitem(__import__("sys").modules, "psutil", fake_psutil)
assert device._system_memory_gb() == (32.0, 12.0)
monkeypatch.delitem(__import__("sys").modules, "psutil", raising=False)
real_import = __import__("builtins").__import__
def reject_psutil(name, *args, **kwargs):
if name == "psutil":
raise ImportError
return real_import(name, *args, **kwargs)
monkeypatch.setattr("builtins.__import__", reject_psutil)
monkeypatch.setattr(device.os, "sysconf", lambda name: {"SC_PHYS_PAGES": 4, "SC_PAGE_SIZE": gib}[name])
assert device._system_memory_gb() == (4.0, 2.4)
monkeypatch.setattr(device.os, "sysconf", Mock(side_effect=ValueError))
assert device._system_memory_gb() == (16.0, 8.0)
def test_memory_info_for_cuda_and_cuda_fallback(monkeypatch):
gib = 1024**3
monkeypatch.setattr(device, "is_cuda", lambda: True)
monkeypatch.setattr(device, "get_device_name", lambda: "GPU")
monkeypatch.setattr(device.torch.cuda, "mem_get_info", lambda _index: (6 * gib, 8 * gib))
monkeypatch.setattr(device.torch.cuda, "memory_allocated", lambda _index: 1 * gib)
monkeypatch.setattr(device.torch.cuda, "memory_reserved", lambda _index: 2 * gib)
assert device.get_memory_info(2) == device.MemoryInfo(1, 2, 8, 6, "GPU")
monkeypatch.setattr(device.torch.cuda, "mem_get_info", Mock(side_effect=RuntimeError("unsupported")))
monkeypatch.setattr(
device.torch.cuda,
"get_device_properties",
lambda _index: SimpleNamespace(total_memory=10 * gib),
)
assert device.get_memory_info(2) == device.MemoryInfo(total_gb=10, free_gb=10, device_name="GPU")
def test_memory_info_for_mps_cpu_and_total_free(monkeypatch):
monkeypatch.setattr(device, "is_cuda", lambda: False)
monkeypatch.setattr(device, "is_mps", lambda: True)
monkeypatch.setattr(device, "get_device_name", lambda: "MPS")
monkeypatch.setattr(device, "_system_memory_gb", lambda: (20.0, 8.0))
assert device.get_memory_info() == device.MemoryInfo(6, 0, 14, 8, "MPS")
assert device.get_total_free_gb() == pytest.approx(5.6)
monkeypatch.setattr(device, "is_mps", lambda: False)
monkeypatch.setattr(device, "get_device_name", lambda: "CPU")
assert device.get_memory_info() == device.MemoryInfo(total_gb=20, free_gb=8, device_name="CPU")
assert device.get_total_free_gb() == 0
def test_total_cuda_memory_sums_query_and_fallback(monkeypatch):
gib = 1024**3
monkeypatch.setattr(device, "is_cuda", lambda: True)
monkeypatch.setattr(device.torch.cuda, "device_count", lambda: 2)
monkeypatch.setattr(
device.torch.cuda,
"mem_get_info",
Mock(side_effect=[(3 * gib, 4 * gib), RuntimeError("no query")]),
)
monkeypatch.setattr(
device.torch.cuda,
"get_device_properties",
lambda _index: SimpleNamespace(total_memory=5 * gib),
)
assert device.get_total_free_gb() == 8.0
def test_cache_cleanup_paths_are_best_effort(monkeypatch):
cuda_empty = Mock()
monkeypatch.setattr(device, "is_cuda", lambda: True)
monkeypatch.setattr(device.torch.cuda, "empty_cache", cuda_empty)
device.empty_cache()
cuda_empty.assert_called_once_with()
sync = Mock()
reset = Mock()
monkeypatch.setattr(device.torch.cuda, "empty_cache", Mock(side_effect=RuntimeError("busy")))
monkeypatch.setattr(device.torch.cuda, "synchronize", sync)
monkeypatch.setattr(device.torch.cuda, "reset_peak_memory_stats", reset)
monkeypatch.setattr(device.gc, "collect", Mock())
device.free_gpu_memory()
sync.assert_called_once_with()
reset.assert_called_once_with()
mps_empty = Mock()
mps_sync = Mock(side_effect=RuntimeError("busy"))
monkeypatch.setattr(device, "is_cuda", lambda: False)
monkeypatch.setattr(device, "is_mps", lambda: True)
monkeypatch.setattr(device.torch, "mps", SimpleNamespace(empty_cache=mps_empty, synchronize=mps_sync))
device.empty_cache()
mps_empty.side_effect = RuntimeError("busy")
device.free_gpu_memory()
assert mps_empty.call_count == 2
mps_sync.assert_called_once_with()
def test_seed_dtype_and_capability_contracts(monkeypatch):
manual_seed = Mock()
cuda_seed = Mock()
monkeypatch.setattr(device.torch, "manual_seed", manual_seed)
monkeypatch.setattr(device.torch.cuda, "manual_seed_all", cuda_seed)
monkeypatch.setattr(device, "is_cuda", lambda: True)
device.set_seed_all(7)
manual_seed.assert_called_once_with(7)
cuda_seed.assert_called_once_with(7)
assert device.default_dtype("cpu") is torch.float32
assert device.default_dtype("cuda") is torch.float16
monkeypatch.setattr(device.torch.cuda, "get_device_capability", lambda _index: (8, 0))
assert device.supports_bfloat16("cuda") is True
monkeypatch.setattr(device.torch.cuda, "get_device_capability", lambda _index: (7, 5))
assert device.supports_bfloat16("cuda") is False
monkeypatch.setattr(device, "is_cuda", lambda: False)
assert device.supports_bfloat16("cuda") is False
assert device.supports_bfloat16("cpu") is True
assert device.supports_float64("mps") is False
assert device.supports_float64("cpu") is True
assert device.supports_bitsandbytes("cuda:1") is True
assert device.supports_bitsandbytes("mps") is False
assert device.supports_device_map_auto("cuda") is True
assert device.supports_device_map_auto("cpu") is False
def test_svd_dtype_and_oom_matching():
assert device.safe_svd_dtype(torch.ones(1, dtype=torch.float64)) is torch.float64
assert device.safe_svd_dtype(torch.ones(1, dtype=torch.float16)) is torch.float32
assert device.is_oom_error(torch.cuda.OutOfMemoryError("oom")) is True
assert device.is_oom_error(RuntimeError("MPS backend out of memory")) is True
assert device.is_oom_error(RuntimeError("other")) is False
def test_configure_cuda_allocator(monkeypatch):
monkeypatch.delenv("PYTORCH_CUDA_ALLOC_CONF", raising=False)
monkeypatch.setattr(device, "is_cuda", lambda: True)
device.configure_cuda_alloc()
assert device.os.environ["PYTORCH_CUDA_ALLOC_CONF"] == "expandable_segments:True"
monkeypatch.setenv("PYTORCH_CUDA_ALLOC_CONF", "existing")
device.configure_cuda_alloc()
assert device.os.environ["PYTORCH_CUDA_ALLOC_CONF"] == "existing"
+354
View File
@@ -0,0 +1,354 @@
"""Offline model-loader contracts at provider, device, and quantization boundaries."""
from __future__ import annotations
import builtins
from types import SimpleNamespace
from unittest.mock import MagicMock, Mock
import pytest
import torch
from obliteratus.models import loader
def _config(**overrides):
values = {
"model_type": "gpt2",
"architectures": ["GPT2LMHeadModel"],
"num_hidden_layers": 2,
"num_attention_heads": 4,
"hidden_size": 8,
"intermediate_size": 16,
"vocab_size": 32,
"quantization_config": None,
}
values.update(overrides)
return SimpleNamespace(**values)
def _model():
model = MagicMock()
model.to.return_value = model
model.state_dict.return_value = {"weight": torch.ones(2)}
model.parameters.return_value = iter([torch.ones(2), torch.ones(3)])
return model
@pytest.fixture
def loader_boundary(monkeypatch):
config = _config()
model = _model()
tokenizer = SimpleNamespace(pad_token=None, eos_token="<eos>")
model_class = SimpleNamespace(from_pretrained=Mock(return_value=model))
monkeypatch.setattr(loader.AutoConfig, "from_pretrained", Mock(return_value=config))
monkeypatch.setattr(loader.AutoTokenizer, "from_pretrained", Mock(return_value=tokenizer))
monkeypatch.setitem(loader.TASK_MODEL_MAP, "causal_lm", model_class)
monkeypatch.setitem(loader.TASK_MODEL_MAP, "classification", model_class)
monkeypatch.setattr(loader.dev, "get_device", lambda preference="auto": "cpu" if preference == "auto" else preference)
monkeypatch.setattr(loader.dev, "supports_device_map_auto", lambda _device=None: False)
monkeypatch.setattr(loader.dev, "supports_bitsandbytes", lambda _device=None: False)
monkeypatch.setattr(loader.dev, "supports_bfloat16", lambda _device=None: True)
monkeypatch.setattr(loader.dev, "get_total_free_gb", lambda: 0.0)
monkeypatch.setattr(loader.dev, "empty_cache", Mock())
monkeypatch.setattr(loader.dev, "is_cuda", lambda: False)
return SimpleNamespace(config=config, model=model, tokenizer=tokenizer, model_class=model_class)
@pytest.mark.parametrize(
("kwargs", "message"),
[
({"model_name": ""}, "model_name must be"),
({"model_name": "x", "task": "embedding"}, "Unknown task"),
({"model_name": "x", "dtype": "int9"}, "Unknown dtype"),
({"model_name": "x", "quantization": "3bit"}, "Unknown quantization"),
],
)
def test_invalid_requests_fail_before_provider_access(loader_boundary, monkeypatch, kwargs, message):
config_load = Mock(side_effect=AssertionError("provider should not be called"))
monkeypatch.setattr(loader.AutoConfig, "from_pretrained", config_load)
with pytest.raises(ValueError, match=message):
loader.load_model(**kwargs)
config_load.assert_not_called()
def test_revision_trust_and_offline_flags_reach_every_provider(loader_boundary):
handle = loader.load_model(
"local/model",
revision="deadbeef",
trust_remote_code=True,
local_files_only=True,
skip_snapshot=True,
)
expected = {
"trust_remote_code": True,
"token": None,
"revision": "deadbeef",
"local_files_only": True,
}
loader.AutoConfig.from_pretrained.assert_called_once_with("local/model", **expected)
loader_boundary.model_class.from_pretrained.assert_called_once()
for key, value in expected.items():
assert loader_boundary.model_class.from_pretrained.call_args.kwargs[key] == value
loader.AutoTokenizer.from_pretrained.assert_called_once_with("local/model", **expected)
assert handle.tokenizer.pad_token == "<eos>"
loader_boundary.model.eval.assert_called_once_with()
loader.dev.empty_cache.assert_called_once_with()
def test_hf_token_is_forwarded_without_logging_value(loader_boundary, monkeypatch):
monkeypatch.setenv("HF_TOKEN", "secret-token")
loader.load_model("x", skip_snapshot=True)
assert loader.AutoConfig.from_pretrained.call_args.kwargs["token"] == "secret-token"
assert loader_boundary.model_class.from_pretrained.call_args.kwargs["token"] == "secret-token"
assert loader.AutoTokenizer.from_pretrained.call_args.kwargs["token"] == "secret-token"
def test_device_and_dtype_failures_precede_provider_access(loader_boundary, monkeypatch):
provider = Mock(side_effect=AssertionError("provider should not be called"))
monkeypatch.setattr(loader.AutoConfig, "from_pretrained", provider)
monkeypatch.setattr(loader.dev, "get_device", Mock(side_effect=RuntimeError("CUDA unavailable")))
with pytest.raises(RuntimeError, match="CUDA unavailable"):
loader.load_model("x", device="cuda")
provider.assert_not_called()
monkeypatch.setattr(loader.dev, "get_device", lambda _preference="auto": "mps")
monkeypatch.setattr(loader.dev, "supports_bfloat16", lambda _device=None: False)
with pytest.raises(RuntimeError, match="bfloat16 is not supported on 'mps'"):
loader.load_model("x", device="mps", dtype="bfloat16")
provider.assert_not_called()
def test_config_permission_error_retries_in_temp_cache(loader_boundary, monkeypatch, tmp_path):
monkeypatch.setattr(loader.tempfile, "gettempdir", lambda: str(tmp_path))
loader.AutoConfig.from_pretrained.side_effect = [PermissionError("cache"), loader_boundary.config]
loader.load_model("x", skip_snapshot=True)
assert loader.AutoConfig.from_pretrained.call_count == 2
assert loader.AutoConfig.from_pretrained.call_args.kwargs["cache_dir"] == str(tmp_path / "hf_home" / "hub")
assert (tmp_path / "hf_home" / "hub").is_dir()
@pytest.mark.parametrize("message", ["Gated repo denied", "Access to model is restricted"])
def test_config_gated_repo_failure_is_actionable(loader_boundary, monkeypatch, message):
loader.AutoConfig.from_pretrained.side_effect = OSError(message)
with pytest.raises(RuntimeError, match=r"(?s)Accept the license.*HF_TOKEN"):
loader.load_model("owner/gated", skip_snapshot=True)
def test_non_gated_config_oserror_is_preserved(loader_boundary):
loader.AutoConfig.from_pretrained.side_effect = OSError("offline cache miss")
with pytest.raises(OSError, match="offline cache miss"):
loader.load_model("x", local_files_only=True)
@pytest.mark.parametrize("exc", [ValueError("unknown"), KeyError("model_type")])
def test_malformed_or_unknown_config_has_stable_guidance(loader_boundary, exc):
loader.AutoConfig.from_pretrained.side_effect = exc
with pytest.raises(RuntimeError, match=r"(?s)not recognized by transformers.*trust_remote_code=True"):
loader.load_model("new/model")
def test_task_model_selection_and_gemma_contract(monkeypatch):
causal = object()
classification = object()
image_text = object()
monkeypatch.setitem(loader.TASK_MODEL_MAP, "causal_lm", causal)
monkeypatch.setitem(loader.TASK_MODEL_MAP, "classification", classification)
monkeypatch.setattr(loader, "AutoModelForImageTextToText", image_text)
assert loader._select_model_class("causal_lm", _config()) is causal
assert loader._select_model_class("classification", _config()) is classification
assert loader._select_model_class("causal_lm", _config(model_type="gemma4")) is image_text
assert loader._select_model_class(
"causal_lm",
_config(model_type="unknown", architectures=["Gemma4ForConditionalGeneration"]),
) is image_text
monkeypatch.setattr(loader, "AutoModelForImageTextToText", None)
with pytest.raises(RuntimeError, match="Upgrade transformers"):
loader._select_model_class("causal_lm", _config(model_type="gemma4"))
with pytest.raises(ValueError, match="Unknown task"):
loader._select_model_class("embedding", _config())
def test_model_handle_metadata_snapshot_restore_summary_and_cleanup(tmp_path):
model = _model()
nested = SimpleNamespace(
num_hidden_layers=3,
num_attention_heads=6,
hidden_size=12,
intermediate_size=24,
)
config = _config(
num_hidden_layers=0,
num_attention_heads=0,
hidden_size=0,
intermediate_size=0,
text_config=nested,
)
offload = tmp_path / "offload"
offload.mkdir()
(offload / "weight").write_text("x")
handle = loader.ModelHandle(
model,
SimpleNamespace(),
config,
"x",
"causal_lm",
_offload_dir=str(offload),
_owns_offload_dir=True,
)
assert (handle.num_layers, handle.num_heads, handle.hidden_size, handle.intermediate_size) == (3, 6, 12, 24)
with pytest.raises(RuntimeError, match="call .snapshot"):
handle.restore()
handle.snapshot()
handle.restore()
model.load_state_dict.assert_called_once()
assert handle.summary() == {
"model_name": "x",
"architecture": "gpt2",
"task": "causal_lm",
"num_layers": 3,
"num_heads": 6,
"hidden_size": 12,
"intermediate_size": 24,
"total_params": 5,
}
handle.cleanup()
assert not offload.exists()
assert handle._offload_dir is None
def test_model_memory_estimation_handles_dense_moe_nested_and_unknown():
dense = loader._estimate_model_memory_gb(_config(), torch.float32)
moe = loader._estimate_model_memory_gb(_config(num_local_experts=4), torch.float32)
assert dense > 0
assert moe > dense
nested = _config(hidden_size=0, num_hidden_layers=0, intermediate_size=0, vocab_size=0)
nested.text_config = _config(hidden_size=8, num_hidden_layers=2, intermediate_size=16, vocab_size=32)
assert loader._estimate_model_memory_gb(nested, torch.float16) > 0
assert loader._estimate_model_memory_gb(_config(hidden_size=0), torch.float16) == 0
@pytest.mark.parametrize("quantization", ["4bit", "8bit"])
def test_quantization_rejects_non_cuda_instead_of_silent_degradation(loader_boundary, quantization):
with pytest.raises(RuntimeError, match="requires an available NVIDIA CUDA device"):
loader.load_model("x", quantization=quantization)
loader_boundary.model_class.from_pretrained.assert_not_called()
def test_quantization_requires_bitsandbytes(loader_boundary, monkeypatch):
monkeypatch.setattr(loader.dev, "get_device", lambda _preference="auto": "cuda")
monkeypatch.setattr(loader.dev, "supports_bitsandbytes", lambda _device=None: True)
real_import = builtins.__import__
def reject_bitsandbytes(name, *args, **kwargs):
if name == "bitsandbytes":
raise ImportError
return real_import(name, *args, **kwargs)
monkeypatch.setattr(builtins, "__import__", reject_bitsandbytes)
with pytest.raises(RuntimeError, match="pip install -U bitsandbytes"):
loader.load_model("x", quantization="4bit")
@pytest.mark.parametrize("quantization", ["4bit", "8bit"])
def test_cuda_quantization_builds_explicit_config(loader_boundary, monkeypatch, quantization):
made = []
def fake_bnb(**kwargs):
made.append(kwargs)
return kwargs
monkeypatch.setattr(loader.dev, "get_device", lambda _preference="auto": "cuda")
monkeypatch.setattr(loader.dev, "supports_bitsandbytes", lambda _device=None: True)
monkeypatch.setattr("transformers.BitsAndBytesConfig", fake_bnb)
handle = loader.load_model("x", quantization=quantization, skip_snapshot=True)
kwargs = loader_boundary.model_class.from_pretrained.call_args.kwargs
assert kwargs["device_map"] == "auto"
assert kwargs["quantization_config"] == made[0]
assert made[0][f"load_in_{quantization}"] is True
assert made[0]["llm_int8_enable_fp32_cpu_offload"] is True
handle.cleanup()
def test_native_quantization_wins_and_skips_dtype(loader_boundary):
loader_boundary.config.quantization_config = SimpleNamespace()
handle = loader.load_model("x", quantization="4bit", skip_snapshot=True)
kwargs = loader_boundary.model_class.from_pretrained.call_args.kwargs
assert "torch_dtype" not in kwargs
assert kwargs["device_map"] == "auto"
assert "quantization_config" not in kwargs
handle.cleanup()
def test_cuda_auto_map_has_bounded_memory_and_requested_offload(loader_boundary, monkeypatch, tmp_path):
gib = 1024**3
monkeypatch.setattr(loader.dev, "get_device", lambda _preference="auto": "cuda")
monkeypatch.setattr(loader.dev, "supports_device_map_auto", lambda _device=None: True)
monkeypatch.setattr(loader.dev, "is_cuda", lambda: True)
monkeypatch.setattr(loader.dev, "device_count", lambda: 2)
monkeypatch.setattr(loader.dev, "_system_memory_gb", lambda: (64.0, 40.0))
monkeypatch.setattr(
loader.torch.cuda,
"get_device_properties",
lambda _index: SimpleNamespace(total_memory=20 * gib),
)
handle = loader.load_model("x", offload_folder=str(tmp_path), skip_snapshot=True)
kwargs = loader_boundary.model_class.from_pretrained.call_args.kwargs
assert kwargs["offload_folder"] == str(tmp_path)
assert kwargs["max_memory"] == {0: "17408MiB", 1: "17408MiB", "cpu": "54GiB"}
handle.cleanup()
assert tmp_path.exists(), "cleanup must not remove an operator-owned directory"
def test_model_permission_error_retries_and_explicit_device_moves(loader_boundary, monkeypatch, tmp_path):
monkeypatch.setattr(loader.tempfile, "gettempdir", lambda: str(tmp_path))
loader_boundary.model_class.from_pretrained.side_effect = [PermissionError("cache"), loader_boundary.model]
loader.load_model("x", device="cpu", skip_snapshot=True)
assert loader_boundary.model_class.from_pretrained.call_count == 2
assert loader_boundary.model_class.from_pretrained.call_args.kwargs["cache_dir"].endswith("hf_home/hub")
loader_boundary.model.to.assert_called_once_with("cpu")
def test_model_gated_and_unknown_architecture_errors_are_actionable(loader_boundary):
loader_boundary.model_class.from_pretrained.side_effect = OSError("gated repo")
with pytest.raises(RuntimeError, match="Accept the license"):
loader.load_model("owner/gated")
loader_boundary.model_class.from_pretrained.side_effect = ValueError("does not recognize this architecture")
with pytest.raises(RuntimeError, match="pip install --upgrade transformers"):
loader.load_model("new/model")
loader_boundary.model_class.from_pretrained.side_effect = KeyError("unrelated")
with pytest.raises(KeyError, match="unrelated"):
loader.load_model("x")
def test_tokenizer_permission_retry_and_forced_snapshot(loader_boundary, monkeypatch, tmp_path):
monkeypatch.setattr(loader.tempfile, "gettempdir", lambda: str(tmp_path))
loader.AutoTokenizer.from_pretrained.side_effect = [PermissionError("cache"), loader_boundary.tokenizer]
handle = loader.load_model("x", skip_snapshot=False)
assert loader.AutoTokenizer.from_pretrained.call_count == 2
assert loader.AutoTokenizer.from_pretrained.call_args.kwargs["cache_dir"].endswith("hf_home/hub")
assert handle._original_state is not None
@pytest.mark.parametrize(
("native", "initial_free", "remaining_free", "snapshots"),
[
(True, 10.0, 3.0, 0),
(True, 10.0, 6.0, 1),
(False, 0.000001, 0.0, 0),
(False, 0.0, 0.0, 1),
],
)
def test_automatic_snapshot_memory_policy(
loader_boundary, monkeypatch, native, initial_free, remaining_free, snapshots,
):
if native:
loader_boundary.config.quantization_config = SimpleNamespace()
monkeypatch.setattr(loader.dev, "get_total_free_gb", Mock(side_effect=[initial_free, remaining_free]))
snapshot = Mock()
monkeypatch.setattr(loader.ModelHandle, "snapshot", snapshot)
handle = loader.load_model("x")
assert snapshot.call_count == snapshots
handle.cleanup()
+170
View File
@@ -0,0 +1,170 @@
"""Simulated MLX contracts that run without Apple hardware or MLX packages."""
from __future__ import annotations
import sys
from types import ModuleType, SimpleNamespace
from unittest.mock import Mock
import numpy as np
import pytest
import torch
from obliteratus import mlx_backend
@pytest.fixture
def fake_mlx(monkeypatch):
core = ModuleType("mlx.core")
core.array = np.array
core.eval = Mock()
core.matmul = np.matmul
core.save_safetensors = Mock()
nn = ModuleType("mlx.nn")
package = ModuleType("mlx")
package.core = core
package.nn = nn
lm = ModuleType("mlx_lm")
lm.load = Mock(return_value=(SimpleNamespace(config={"kind": "fixture"}), SimpleNamespace()))
lm.generate = Mock(return_value="generated")
lm.upload_to_hub = Mock()
monkeypatch.setitem(sys.modules, "mlx", package)
monkeypatch.setitem(sys.modules, "mlx.core", core)
monkeypatch.setitem(sys.modules, "mlx.nn", nn)
monkeypatch.setitem(sys.modules, "mlx_lm", lm)
monkeypatch.setattr(mlx_backend, "MLX_AVAILABLE", True)
monkeypatch.setattr(mlx_backend, "_mx", core)
monkeypatch.setattr(mlx_backend, "_mlx_nn", nn)
monkeypatch.setattr(mlx_backend, "_mlx_lm", lm)
return SimpleNamespace(core=core, lm=lm)
@pytest.mark.parametrize(
"call",
[
lambda: mlx_backend.load_model("model"),
lambda: mlx_backend.generate(SimpleNamespace(), "prompt"),
lambda: mlx_backend.get_activations(SimpleNamespace(), [], []),
lambda: mlx_backend.get_weight(SimpleNamespace(), 0, "weight"),
lambda: mlx_backend.modify_weights(SimpleNamespace(), 0, "weight", lambda value: value),
lambda: mlx_backend.project_out_direction(None, None),
lambda: mlx_backend.save_model(SimpleNamespace(), "out"),
lambda: mlx_backend.torch_tensor_to_mlx(torch.ones(1)),
],
)
def test_missing_mlx_fails_with_install_guidance(monkeypatch, call):
monkeypatch.setattr(mlx_backend, "MLX_AVAILABLE", False)
with pytest.raises(RuntimeError, match="pip install mlx>=0.22 mlx-lm>=0.20"):
call()
def test_load_generate_and_handle_config(fake_mlx):
handle = mlx_backend.load_model("local/model", dtype="bfloat16")
assert handle.model_name == "local/model"
assert handle.config == {"kind": "fixture"}
fake_mlx.lm.load.assert_called_once_with("local/model")
assert mlx_backend.generate(
handle,
"prompt",
max_tokens=4,
temperature=0.2,
top_p=0.8,
repetition_penalty=1.1,
) == "generated"
assert fake_mlx.lm.generate.call_args.kwargs == {
"prompt": "prompt",
"max_tokens": 4,
"temp": 0.2,
"top_p": 0.8,
"repetition_penalty": 1.1,
}
mlx_backend.generate(handle, "prompt", repetition_penalty=None)
assert "repetition_penalty" not in fake_mlx.lm.generate.call_args.kwargs
class _Layer:
def __call__(self, value):
return (value + 1, "attention")
def test_activation_capture_supports_llama_and_truncation(fake_mlx):
inner = SimpleNamespace(
layers=[_Layer(), _Layer()],
embed_tokens=lambda ids: np.repeat(ids[:, :, None], 3, axis=2),
)
handle = mlx_backend.MLXModelHandle(
model=SimpleNamespace(model=inner),
tokenizer=SimpleNamespace(encode=lambda _prompt: [1, 2, 3, 4]),
model_name="fixture",
)
activations = mlx_backend.get_activations(handle, ["a", "b"], [0, 1], max_length=2)
assert np.array_equal(activations[0][0], np.array([3, 3, 3]))
assert np.array_equal(activations[1][1], np.array([4, 4, 4]))
assert fake_mlx.core.eval.call_count == 4
def test_activation_capture_rejects_unknown_layers_and_embeddings(fake_mlx):
handle = mlx_backend.MLXModelHandle(SimpleNamespace(), SimpleNamespace(), "unknown")
with pytest.raises(RuntimeError, match="Cannot locate transformer layers"):
mlx_backend.get_activations(handle, ["x"], [0])
model = SimpleNamespace(model=SimpleNamespace(layers=[_Layer()]))
handle = mlx_backend.MLXModelHandle(model, SimpleNamespace(encode=lambda _prompt: [1]), "unknown")
with pytest.raises(RuntimeError, match="Cannot find embedding layer"):
mlx_backend.get_activations(handle, ["x"], [0])
def test_get_and_modify_weights_support_update_and_assignment(fake_mlx):
update_parent = SimpleNamespace(weight=np.array([1.0]), update=Mock())
assign_parent = SimpleNamespace(weight=np.array([2.0]))
layers = [SimpleNamespace(attention=update_parent), SimpleNamespace(attention=assign_parent)]
handle = mlx_backend.MLXModelHandle(
SimpleNamespace(transformer=SimpleNamespace(h=layers)),
SimpleNamespace(),
"fixture",
)
assert np.array_equal(mlx_backend.get_weight(handle, 0, "attention.weight"), np.array([1.0]))
mlx_backend.modify_weights(handle, 0, "attention.weight", lambda value: value + 1)
update_parent.update.assert_called_once()
mlx_backend.modify_weights(handle, 1, "attention.weight", lambda value: value + 1)
assert np.array_equal(assign_parent.weight, np.array([3.0]))
assert fake_mlx.core.eval.call_count == 2
def test_projection_removes_requested_direction(fake_mlx):
weight = np.array([[2.0, 3.0], [4.0, 5.0]])
direction = np.array([1.0, 0.0])
projected = mlx_backend.project_out_direction(weight, direction)
assert np.array_equal(projected, np.array([[0.0, 3.0], [0.0, 5.0]]))
def test_save_model_native_fallback_and_upload(fake_mlx, monkeypatch, tmp_path):
tokenizer = SimpleNamespace(save_pretrained=Mock())
model = SimpleNamespace(parameters=lambda: {"block": {"weight": np.array([1.0])}}.items())
handle = mlx_backend.MLXModelHandle(model, tokenizer, "fixture")
fake_mlx.lm.save_model = Mock()
out = mlx_backend.save_model(handle, tmp_path / "native", upload_repo="org/model")
assert out.is_dir()
fake_mlx.lm.save_model.assert_called_once()
fake_mlx.lm.upload_to_hub.assert_called_once_with(str(out), "org/model")
del fake_mlx.lm.save_model
out = mlx_backend.save_model(handle, tmp_path / "fallback")
fake_mlx.core.save_safetensors.assert_called_once()
tokenizer.save_pretrained.assert_called_once_with(str(out))
def test_tensor_conversions_and_internal_helpers(fake_mlx):
source = torch.tensor([1.0, 2.0], requires_grad=True)
converted = mlx_backend.torch_tensor_to_mlx(source)
assert np.array_equal(converted, np.array([1.0, 2.0], dtype=np.float32))
restored = mlx_backend.mlx_to_torch_tensor(np.array([3.0]), device="cpu")
assert torch.equal(restored, torch.tensor([3.0], dtype=torch.float64))
layers = [object()]
assert mlx_backend._get_layers(SimpleNamespace(gpt_neox=SimpleNamespace(layers=layers))) is layers
with pytest.raises(RuntimeError, match="Cannot locate transformer layers"):
mlx_backend._get_layers(SimpleNamespace())
flattened = {}
mlx_backend._flatten_dict({"a": {"b": 1}, "c": 2}, "", flattened)
assert flattened == {"a.b": 1, "c": 2}