From 3dbe90d9a94e04d8738dc398aa425ee42c069725 Mon Sep 17 00:00:00 2001 From: faber Date: Wed, 19 Aug 2026 21:34:36 -0400 Subject: [PATCH] tests: 96% coverage --- tests/test_capability_check.py | 59 ++++++++++++++++++++++++++++++++++ 1 file changed, 59 insertions(+) diff --git a/tests/test_capability_check.py b/tests/test_capability_check.py index 124a521..0cf3b04 100644 --- a/tests/test_capability_check.py +++ b/tests/test_capability_check.py @@ -103,3 +103,62 @@ class TestCapabilityCheck: summary = json.loads(summary_path.read_text()) assert "delta_pp" in summary assert summary["method"] == "lm-eval-harness 0-shot log-likelihood" + + @patch("obliteratus.capability_check._run_lm_eval") + def test_custom_subjects(self, mock_lm_eval, tmp_path): + from obliteratus.capability_check import capability_check + + mock_lm_eval.side_effect = [ + {"results": {"mmlu_physics": {"acc,none": 0.75}}}, + {"results": {"mmlu_physics": {"acc,none": 0.80}}}, + ] + + result = capability_check( + "fake/abliterated", "fake/stock", + device="cpu", subjects=["mmlu_physics"], + output_dir=str(tmp_path), + ) + + assert result["tasks"] == "mmlu_physics" + assert result["abliterated_acc"] == 0.75 + + @patch("obliteratus.capability_check._run_lm_eval") + def test_individual_subjects_mean(self, mock_lm_eval, tmp_path): + """Test mean computation when no aggregate mmlu key exists.""" + from obliteratus.capability_check import capability_check + + # Results don't have "mmlu" aggregate key or match tasks.split(",")[0] + mock_lm_eval.side_effect = [ + {"results": {"sub_a": {"acc,none": 0.6}, "sub_b": {"acc,none": 0.8}}}, + {"results": {"sub_a": {"acc,none": 0.7}, "sub_b": {"acc,none": 0.9}}}, + ] + + result = capability_check( + "fake/abliterated", "fake/stock", + device="cpu", subjects=["mmlu_a", "mmlu_b"], + output_dir=str(tmp_path), + ) + + assert abs(result["abliterated_acc"] - 0.7) < 0.01 + assert abs(result["stock_acc"] - 0.8) < 0.01 + + +class TestMainCLI: + @patch("obliteratus.capability_check.capability_check") + def test_main_runs(self, mock_check): + from obliteratus.capability_check import main + + mock_check.return_value = { + "stock_acc": 0.87, "abliterated_acc": 0.81, + "delta_pp": -6.0, "tasks": "mmlu", "method": "test", + } + + import sys + old_argv = sys.argv + sys.argv = ["prog", "--abliterated", "fake/abl", "--stock", "fake/stock", "--device", "cpu"] + try: + main() + finally: + sys.argv = old_argv + + mock_check.assert_called_once()