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:
@@ -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"}
|
||||
|
||||
Reference in New Issue
Block a user