feat: pre-flight filter of untrainable task types before gate

This commit is contained in:
2026-07-16 06:16:02 -04:00
parent c44f6010eb
commit caea9150b3
10 changed files with 148 additions and 7 deletions
+54
View File
@@ -20,6 +20,7 @@ from app.harness.runner import (
_build_comparison_pairs,
_compute_total_steps,
_fallback_summary,
_filter_untrainable_types,
_format_applied_edits,
_guard_infra_failures,
_outcome_to_quadrant_pairs,
@@ -401,6 +402,58 @@ class TestShouldEarlyStop:
assert state.epochs_since_best_improved == 2
class TestFilterUntrainableTypes:
"""_filter_untrainable_types 可训练性预检纯函数。"""
def test_untrainable_types_filtered_before_gate(self) -> None:
"""val<eval_min_per_class 或 非test单元<trainable_min_units 的题型被剔除。"""
# 题型 Adiag=5 + val=5 → units=10、val=5,可训
diag = [_FakeQuestion(question_id=f"A-d{i}", task_type="A") for i in range(5)]
val = [_FakeQuestion(question_id=f"A-v{i}", task_type="A") for i in range(5)]
# 题型 Bval=0<eval_min_per_class)不可训
diag += [_FakeQuestion(question_id=f"B-d{i}", task_type="B") for i in range(2)]
pools = _FakePools(diagnosis=diag, validation=val, test=[])
new_pools, new_types = _filter_untrainable_types(
pools,
task_types=["A", "B"],
eval_min_per_class=2,
trainable_min_units=8,
)
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_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
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 == []
def test_test_pool_untouched(self) -> None:
"""test 池不参与过滤(继续报告全题型准确率)。"""
val = [_FakeQuestion(question_id=f"A-v{i}", task_type="A") for i in range(8)]
diag = [_FakeQuestion(question_id=f"A-d{i}", task_type="A") for i in range(8)]
test = [_FakeQuestion(question_id="B-t0", task_type="B")]
pools = _FakePools(diagnosis=diag, validation=val, test=test)
new_pools, _ = _filter_untrainable_types(
pools, task_types=None, eval_min_per_class=2, trainable_min_units=8
)
assert new_pools.test == test
# =========================================================================
# 13c: Probation 数据结构测试
# =========================================================================
@@ -690,6 +743,7 @@ class TestRunnerFactoryInjection:
"batch_size": 5,
"min_class_per_batch": 2,
"eval_min_per_class": 2,
"trainable_min_units": 8,
"early_stop_patience": 3,
"test_size": 10,
"use_slow_momentum": False,