fix: count units not questions in trainability pre-flight; fail-fast when all types filtered

This commit is contained in:
2026-07-16 06:45:16 -04:00
parent ef244c52bd
commit e1f08dcd3b
2 changed files with 86 additions and 15 deletions
+67 -7
View File
@@ -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 的题型被剔除。"""
# 题型 Cval=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 的题型被剔除(保留另一可训题型)"""
# 可训题型 Adiag=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)]
# 题型 Cval=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 按单元折叠计数:题目数达标但单元数不足的题型仍被剔除。"""
# 可训题型 Adiag=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)]
# 题型 Pdiag 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-fastraise 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 池不参与过滤(继续报告全题型准确率)。"""