From e1f08dcd3b6cda52625c39d554232490870c8b54 Mon Sep 17 00:00:00 2001 From: iomgaa Date: Thu, 16 Jul 2026 06:45:16 -0400 Subject: [PATCH] fix: count units not questions in trainability pre-flight; fail-fast when all types filtered --- app/harness/runner.py | 27 +++++++---- tests/unit/test_harness_runner.py | 74 ++++++++++++++++++++++++++++--- 2 files changed, 86 insertions(+), 15 deletions(-) diff --git a/app/harness/runner.py b/app/harness/runner.py index 3c7f682..ae35031 100644 --- a/app/harness/runner.py +++ b/app/harness/runner.py @@ -298,37 +298,48 @@ def _filter_untrainable_types( eval_min_per_class: int, trainable_min_units: int, ) -> tuple[Pools, list[str] | None]: - """剔除不可训练题型(val={eval_min_per_class} 且 units>={trainable_min_units}:" + f"{detail}。请调整池切分或降低阈值。" + ) new_pools = replace( pools, diagnosis=[q for q in pools.diagnosis if q.task_type in keep], diff --git a/tests/unit/test_harness_runner.py b/tests/unit/test_harness_runner.py index f328f04..98dad2f 100644 --- a/tests/unit/test_harness_runner.py +++ b/tests/unit/test_harness_runner.py @@ -75,6 +75,26 @@ class _FakeQuestion: object.__setattr__(self, "unit_id", self.pair_id or self.question_id) +def _fake_pair(pair_id: str, task_type: str, video_id: str = "v1") -> list[_FakeQuestion]: + """构造合法孪生对(original + mirror),共享 pair_id/video_id/task_type/flip_axis。""" + return [ + _FakeQuestion( + question_id=f"{pair_id}-o", + video_id=video_id, + task_type=task_type, + pair_id=pair_id, + question_role="pair_original", + ), + _FakeQuestion( + question_id=f"{pair_id}-m", + video_id=video_id, + task_type=task_type, + pair_id=pair_id, + question_role="pair_mirror", + ), + ] + + @dataclass class _FakePools: """Pools 替身。""" @@ -426,19 +446,59 @@ class TestFilterUntrainableTypes: assert new_types == ["A"] def test_units_below_threshold_filtered(self) -> None: - """val 达标但 diag+val 单元数 < trainable_min_units 的题型被剔除。""" - # 题型 C:val=2(>=2)但 units=2+0=... 补 diag 使总数不足 - val = [_FakeQuestion(question_id=f"C-v{i}", task_type="C") for i in range(2)] - diag = [_FakeQuestion(question_id="C-d0", task_type="C")] # units=3 < 8 + """val 达标但 diag+val 单元数 < trainable_min_units 的题型被剔除(保留另一可训题型)。""" + # 可训题型 A:diag=8 + val=2 → units=10、val=2,保留 + keep_diag = [_FakeQuestion(question_id=f"A-d{i}", task_type="A") for i in range(8)] + keep_val = [_FakeQuestion(question_id=f"A-v{i}", task_type="A") for i in range(2)] + # 题型 C:val=2(>=2)但 units=2+1=3 < 8 → 剔除 + val = keep_val + [_FakeQuestion(question_id=f"C-v{i}", task_type="C") for i in range(2)] + diag = keep_diag + [_FakeQuestion(question_id="C-d0", task_type="C")] pools = _FakePools(diagnosis=diag, validation=val, test=[]) new_pools, new_types = _filter_untrainable_types( pools, task_types=None, eval_min_per_class=2, trainable_min_units=8 ) - assert new_pools.diagnosis == [] - assert new_pools.validation == [] - assert new_types == [] + assert {q.task_type for q in new_pools.diagnosis} == {"A"} + assert {q.task_type for q in new_pools.validation} == {"A"} + assert new_types == ["A"] + + def test_ar_pair_counted_as_units_not_questions(self) -> None: + """AR pair 按单元折叠计数:题目数达标但单元数不足的题型仍被剔除。""" + # 可训题型 A:diag=8 + val=2 single → units=10,保留 + keep_diag = [_FakeQuestion(question_id=f"A-d{i}", task_type="A") for i in range(8)] + keep_val = [_FakeQuestion(question_id=f"A-v{i}", task_type="A") for i in range(2)] + # 题型 P:diag 3 对(6 题=3 单元)+ val 2 对(4 题=2 单元)→ 单元数=5<8, + # 但题目数=10>=8。按单元计数须剔除(按题目计数会误通过)。 + pair_diag: list[_FakeQuestion] = [] + for i in range(3): + pair_diag.extend(_fake_pair(f"P-d{i}", "P")) + pair_val: list[_FakeQuestion] = [] + for i in range(2): + pair_val.extend(_fake_pair(f"P-v{i}", "P")) + pools = _FakePools( + diagnosis=keep_diag + pair_diag, + validation=keep_val + pair_val, + test=[], + ) + + new_pools, new_types = _filter_untrainable_types( + pools, task_types=None, eval_min_per_class=2, trainable_min_units=8 + ) + + assert "P" not in {q.task_type for q in new_pools.diagnosis} + assert "P" not in new_types + assert "A" in new_types + + def test_all_filtered_raises(self) -> None: + """所有题型都被剔除时 fail-fast:raise RuntimeError 并列出剔除原因。""" + diag = [_FakeQuestion(question_id="B-d0", task_type="B")] # val=0 < 2 + pools = _FakePools(diagnosis=diag, validation=[], test=[]) + + with pytest.raises(RuntimeError, match="可训练"): + _filter_untrainable_types( + pools, task_types=None, eval_min_per_class=2, trainable_min_units=8 + ) def test_test_pool_untouched(self) -> None: """test 池不参与过滤(继续报告全题型准确率)。"""