feat(harness): add PerCategoryPoolStrategy with correctness-stratified 2:1 split

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
2026-07-12 22:41:48 -04:00
parent ec4cbbdd44
commit 21c6a53aed
2 changed files with 420 additions and 2 deletions
+173 -1
View File
@@ -16,11 +16,13 @@ from typing import TYPE_CHECKING
import pytest
from app.harness.pools import (
GlobalPoolStrategy,
PerCategoryPoolStrategy,
build_pools,
load_pools,
save_pools,
)
from core.types import GeneratedQuestion
from core.types import GeneratedQuestion, PoolConfig
if TYPE_CHECKING:
from pathlib import Path
@@ -288,3 +290,173 @@ class TestBuildOrLoadPoolsFrozen:
orig_ids = [q.question_id for q in getattr(frozen, pool_name)]
load_ids = [q.question_id for q in getattr(loaded, pool_name)]
assert orig_ids == load_ids, f"{pool_name} 冻结后 ID 顺序不一致"
class TestGlobalPoolStrategy:
"""GlobalPoolStrategy 封装现有全局三分逻辑。"""
def test_global_strategy_builds_three_pools(self) -> None:
"""GlobalPoolStrategy.build 产出三个互斥池。"""
questions = _make_question_set(200)
correctness = _make_correctness(questions, 0.5)
config = PoolConfig(
task_types=None,
seed=42,
baseline_run_id="run_baseline",
diag_size=30,
diag_correct_ratio=0.5,
val_size=30,
val_correct_ratio=0.5,
test_size=30,
eval_min_per_class=1,
train_ratio=0.667,
test_questions_dir=None,
)
strategy = GlobalPoolStrategy()
pools = strategy.build(questions, correctness, config)
diag_ids = {q.question_id for q in pools.diagnosis}
val_ids = {q.question_id for q in pools.validation}
test_ids = {q.question_id for q in pools.test}
assert diag_ids & val_ids == set()
assert diag_ids & test_ids == set()
assert val_ids & test_ids == set()
assert len(pools.diagnosis) == 30
assert len(pools.validation) == 30
assert len(pools.test) == 30
def test_global_strategy_build_incremental_raises(self) -> None:
"""GlobalPoolStrategy 不支持增量。"""
strategy = GlobalPoolStrategy()
config = PoolConfig(
task_types=None,
seed=0,
baseline_run_id="r",
diag_size=10,
diag_correct_ratio=0.5,
val_size=10,
val_correct_ratio=0.5,
test_size=10,
eval_min_per_class=1,
train_ratio=0.667,
test_questions_dir=None,
)
with pytest.raises(NotImplementedError):
strategy.build_incremental(["Action Reasoning"], [], {}, config)
def _make_per_category_questions():
"""构造 12 类各 30 题,共 360 题。"""
task_types = [
"Action Prediction", "Action Reasoning", "Action Recognition",
"Action Sequence", "Causal Reasoning", "Event Reasoning",
"Object Interaction", "Object Reasoning", "Object Recognition",
"Scene Understanding", "Spatial Reasoning", "Temporal Reasoning",
]
questions = []
for tt in task_types:
for i in range(30):
questions.append(_make_question(f"{tt}_{i:03d}", tt))
return questions
class TestPerCategoryPoolStrategy:
"""PerCategoryPoolStrategy per-category 2:1 分层划分。"""
def test_per_category_split_20_10(self):
"""每类 30 题按 correctness 2:1 分层 -> 20 train + 10 val。"""
questions = _make_per_category_questions()
correctness = {}
for q in questions:
idx = int(q.question_id.split("_")[-1])
correctness[q.question_id] = idx < 18
config = PoolConfig(
task_types=None, seed=42, baseline_run_id="baseline_v2",
diag_size=0, diag_correct_ratio=0.0, val_size=0,
val_correct_ratio=0.0, test_size=0, eval_min_per_class=0,
train_ratio=20 / 30, test_questions_dir=None,
)
strategy = PerCategoryPoolStrategy()
pools = strategy.build(questions, correctness, config)
assert len(pools.diagnosis) == 240
assert len(pools.validation) == 120
from collections import Counter
diag_counts = Counter(q.task_type for q in pools.diagnosis)
val_counts = Counter(q.task_type for q in pools.validation)
for tt in diag_counts:
assert diag_counts[tt] == 20
assert val_counts[tt] == 10
diag_ids = {q.question_id for q in pools.diagnosis}
val_ids = {q.question_id for q in pools.validation}
assert diag_ids & val_ids == set()
def test_per_category_correctness_ratio_aligned(self):
"""train 和 val 的 correctness 比例应对齐。"""
questions = _make_per_category_questions()
correctness = {}
for q in questions:
idx = int(q.question_id.split("_")[-1])
correctness[q.question_id] = idx < 18
config = PoolConfig(
task_types=("Action Reasoning",), seed=42, baseline_run_id="baseline_v2",
diag_size=0, diag_correct_ratio=0.0, val_size=0,
val_correct_ratio=0.0, test_size=0, eval_min_per_class=0,
train_ratio=20 / 30, test_questions_dir=None,
)
strategy = PerCategoryPoolStrategy()
pools = strategy.build(questions, correctness, config)
assert len(pools.diagnosis) == 20
assert len(pools.validation) == 10
diag_correct = sum(1 for q in pools.diagnosis if correctness[q.question_id])
val_correct = sum(1 for q in pools.validation if correctness[q.question_id])
assert diag_correct == 12
assert val_correct == 6
def test_per_category_all_correct_degrades(self):
"""某类全部 correct -> 退化为非分层 random 20/10。"""
questions = [_make_question(f"q_{i:03d}", "Object Recognition") for i in range(30)]
correctness = {q.question_id: True for q in questions}
config = PoolConfig(
task_types=None, seed=42, baseline_run_id="r",
diag_size=0, diag_correct_ratio=0.0, val_size=0,
val_correct_ratio=0.0, test_size=0, eval_min_per_class=0,
train_ratio=20 / 30, test_questions_dir=None,
)
strategy = PerCategoryPoolStrategy()
pools = strategy.build(questions, correctness, config)
assert len(pools.diagnosis) == 20
assert len(pools.validation) == 10
def test_per_category_missing_correctness_fails(self):
"""correctness 不完整时 fail-fast。"""
questions = [_make_question(f"q_{i:03d}", "Object Recognition") for i in range(30)]
correctness = {q.question_id: True for q in questions[:25]}
config = PoolConfig(
task_types=None, seed=42, baseline_run_id="r",
diag_size=0, diag_correct_ratio=0.0, val_size=0,
val_correct_ratio=0.0, test_size=0, eval_min_per_class=0,
train_ratio=20 / 30, test_questions_dir=None,
)
strategy = PerCategoryPoolStrategy()
with pytest.raises(ValueError, match="correctness 缺失"):
strategy.build(questions, correctness, config)
def test_per_category_task_types_filter(self):
"""task_types 过滤只处理指定类别。"""
questions = _make_per_category_questions()
correctness = {q.question_id: True for q in questions}
config = PoolConfig(
task_types=("Action Reasoning", "Scene Understanding"), seed=42,
baseline_run_id="r", diag_size=0, diag_correct_ratio=0.0,
val_size=0, val_correct_ratio=0.0, test_size=0,
eval_min_per_class=0, train_ratio=20 / 30, test_questions_dir=None,
)
strategy = PerCategoryPoolStrategy()
pools = strategy.build(questions, correctness, config)
assert len(pools.diagnosis) == 40
assert len(pools.validation) == 20
types_in_diag = {q.task_type for q in pools.diagnosis}
assert types_in_diag == {"Action Reasoning", "Scene Understanding"}