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:
+247
-1
@@ -9,11 +9,16 @@ test -> validation -> diagnosis 的顺序 progressive exclusion,
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import math
|
||||
import random
|
||||
from collections import defaultdict
|
||||
from dataclasses import dataclass, field
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from loguru import logger
|
||||
|
||||
from app.question_gen import stratified_sample
|
||||
from core.types import GeneratedQuestion
|
||||
from core.types import GeneratedQuestion, PoolConfig
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from pathlib import Path
|
||||
@@ -99,6 +104,66 @@ def build_pools(
|
||||
)
|
||||
|
||||
|
||||
class GlobalPoolStrategy:
|
||||
"""全局三分策略:test -> val -> diag progressive exclusion。
|
||||
|
||||
封装现有 build_pools 逻辑为 PoolStrategy 接口。
|
||||
"""
|
||||
|
||||
def build(
|
||||
self,
|
||||
questions: list[GeneratedQuestion],
|
||||
correctness: dict[str, bool],
|
||||
config: PoolConfig,
|
||||
) -> Pools:
|
||||
"""委托给现有 build_pools 函数。
|
||||
|
||||
参数:
|
||||
questions: 题目全集。
|
||||
correctness: question_id -> 基线是否答对。
|
||||
config: 池构建统一配置。
|
||||
|
||||
返回:
|
||||
冻结的三池 Pools。
|
||||
"""
|
||||
return build_pools(
|
||||
questions,
|
||||
correctness,
|
||||
diag_cfg={
|
||||
"size": config.diag_size,
|
||||
"correct_ratio": config.diag_correct_ratio,
|
||||
"task_types": list(config.task_types) if config.task_types else None,
|
||||
"seed": config.seed,
|
||||
"min_per_class": None,
|
||||
},
|
||||
val_cfg={
|
||||
"size": config.val_size,
|
||||
"correct_ratio": config.val_correct_ratio,
|
||||
"task_types": list(config.task_types) if config.task_types else None,
|
||||
"seed": config.seed,
|
||||
"min_per_class": config.eval_min_per_class,
|
||||
},
|
||||
test_cfg={"size": config.test_size, "seed": config.seed},
|
||||
baseline_run_id=config.baseline_run_id,
|
||||
)
|
||||
|
||||
def build_incremental(
|
||||
self,
|
||||
new_task_types: list[str],
|
||||
questions: list[GeneratedQuestion],
|
||||
correctness: dict[str, bool],
|
||||
config: PoolConfig,
|
||||
) -> dict[str, dict[str, list[str]]]:
|
||||
"""全局策略不支持增量。
|
||||
|
||||
异常:
|
||||
NotImplementedError: 始终抛出。
|
||||
"""
|
||||
raise NotImplementedError(
|
||||
"GlobalPoolStrategy 不支持增量构建,请使用 PerCategoryPoolStrategy。"
|
||||
)
|
||||
|
||||
|
||||
def _sample_excluding(
|
||||
questions: list[GeneratedQuestion],
|
||||
exclude_ids: set[str],
|
||||
@@ -285,3 +350,184 @@ def build_or_load_pools(
|
||||
)
|
||||
save_pools(pools, pools_path)
|
||||
return pools
|
||||
|
||||
|
||||
class PerCategoryPoolStrategy:
|
||||
"""Per-category 分层池构建策略。
|
||||
|
||||
按题型分组,每个题型内部按 correctness 分层,以 train_ratio 比例
|
||||
划分 train(映射到 diagnosis 池)和 val(映射到 validation 池)。
|
||||
与 GlobalPoolStrategy 的全局 progressive exclusion 不同,本策略
|
||||
保证每个类别内部的 train/val 比例精确对齐。
|
||||
"""
|
||||
|
||||
def build(
|
||||
self,
|
||||
questions: list[GeneratedQuestion],
|
||||
correctness: dict[str, bool],
|
||||
config: PoolConfig,
|
||||
) -> Pools:
|
||||
"""按题型分组后,每组做 correctness 分层的 train/val 划分。
|
||||
|
||||
参数:
|
||||
questions: 题目全集。
|
||||
correctness: question_id -> 基线是否答对。
|
||||
config: 池构建配置(使用 train_ratio, task_types, seed,
|
||||
baseline_run_id, test_questions_dir)。
|
||||
|
||||
返回:
|
||||
冻结的 Pools(diagnosis=train, validation=val,
|
||||
test 从 test_questions_dir 加载或为空列表)。
|
||||
"""
|
||||
# Phase 1: 按 task_types 过滤
|
||||
if config.task_types is not None:
|
||||
allowed = set(config.task_types)
|
||||
filtered = [q for q in questions if q.task_type in allowed]
|
||||
else:
|
||||
filtered = list(questions)
|
||||
|
||||
# Phase 2: 按 task_type 分组
|
||||
groups: dict[str, list[GeneratedQuestion]] = defaultdict(list)
|
||||
for q in filtered:
|
||||
groups[q.task_type].append(q)
|
||||
|
||||
# Phase 3: 每组分层划分
|
||||
all_train: list[GeneratedQuestion] = []
|
||||
all_val: list[GeneratedQuestion] = []
|
||||
rng = random.Random(config.seed)
|
||||
|
||||
for task_type in sorted(groups.keys()):
|
||||
train, val = self._split_one_category(
|
||||
groups[task_type], correctness, config.train_ratio, rng,
|
||||
)
|
||||
all_train.extend(train)
|
||||
all_val.extend(val)
|
||||
|
||||
# Phase 4: test 池(从外部目录加载,无则空)
|
||||
test: list[GeneratedQuestion] = []
|
||||
if config.test_questions_dir is not None:
|
||||
from app.question_gen import load_benchmark
|
||||
|
||||
test = load_benchmark(config.test_questions_dir)
|
||||
|
||||
# Phase 5: 计算 baseline_val_accuracy
|
||||
val_correct = sum(
|
||||
1 for q in all_val if correctness.get(q.question_id, False)
|
||||
)
|
||||
baseline_val_accuracy = val_correct / len(all_val) if all_val else 0.0
|
||||
|
||||
return Pools(
|
||||
diagnosis=all_train,
|
||||
validation=all_val,
|
||||
test=test,
|
||||
baseline_run_id=config.baseline_run_id,
|
||||
baseline_val_accuracy=baseline_val_accuracy,
|
||||
correctness={
|
||||
q.question_id: correctness.get(q.question_id, False)
|
||||
for q in all_train + all_val + test
|
||||
},
|
||||
)
|
||||
|
||||
def _split_one_category(
|
||||
self,
|
||||
questions: list[GeneratedQuestion],
|
||||
correctness: dict[str, bool],
|
||||
train_ratio: float,
|
||||
rng: random.Random,
|
||||
) -> tuple[list[GeneratedQuestion], list[GeneratedQuestion]]:
|
||||
"""单类别 correctness 分层划分。
|
||||
|
||||
参数:
|
||||
questions: 单类别全部题目。
|
||||
correctness: question_id -> 基线是否答对。
|
||||
train_ratio: train 占总量的比例。
|
||||
rng: 随机数生成器(保证跨类别可复现)。
|
||||
|
||||
返回:
|
||||
(train, val) 题目列表元组,两池互斥且总量 == len(questions)。
|
||||
|
||||
异常:
|
||||
ValueError: correctness 中缺少某些 question_id。
|
||||
"""
|
||||
n_total = len(questions)
|
||||
n_train = round(n_total * train_ratio)
|
||||
n_val = n_total - n_train
|
||||
|
||||
# 校验 correctness 完整性
|
||||
missing = [q.question_id for q in questions if q.question_id not in correctness]
|
||||
if missing:
|
||||
raise ValueError(
|
||||
f"correctness 缺失 {len(missing)} 题: {missing[:5]}"
|
||||
)
|
||||
|
||||
correct_qs = [q for q in questions if correctness[q.question_id]]
|
||||
wrong_qs = [q for q in questions if not correctness[q.question_id]]
|
||||
n_correct = len(correct_qs)
|
||||
|
||||
# 全 correct 或全 wrong -> 退化为非分层随机划分
|
||||
if n_correct == 0 or n_correct == n_total:
|
||||
label = "全部正确" if n_correct == n_total else "全部错误"
|
||||
logger.warning(
|
||||
"类别 {} {} ({} 题),退化为非分层随机划分",
|
||||
questions[0].task_type, label, n_total,
|
||||
)
|
||||
shuffled = list(questions)
|
||||
rng.shuffle(shuffled)
|
||||
return shuffled[:n_train], shuffled[n_train:]
|
||||
|
||||
# 分层: 按 correctness 比例分配到 train
|
||||
train_correct = math.floor(n_correct * n_train / n_total)
|
||||
train_wrong = n_train - train_correct
|
||||
|
||||
rng.shuffle(correct_qs)
|
||||
rng.shuffle(wrong_qs)
|
||||
|
||||
train = correct_qs[:train_correct] + wrong_qs[:train_wrong]
|
||||
val = correct_qs[train_correct:] + wrong_qs[train_wrong:]
|
||||
|
||||
assert len(train) == n_train, (
|
||||
f"train 数量不匹配: {len(train)} != {n_train}"
|
||||
)
|
||||
assert len(val) == n_val, (
|
||||
f"val 数量不匹配: {len(val)} != {n_val}"
|
||||
)
|
||||
|
||||
return train, val
|
||||
|
||||
def build_incremental(
|
||||
self,
|
||||
new_task_types: list[str],
|
||||
questions: list[GeneratedQuestion],
|
||||
correctness: dict[str, bool],
|
||||
config: PoolConfig,
|
||||
) -> dict[str, dict[str, list[str]]]:
|
||||
"""增量划分:仅处理 new_task_types 中的类别。
|
||||
|
||||
参数:
|
||||
new_task_types: 需要增量划分的类别列表。
|
||||
questions: 题目全集(从中筛选指定类别)。
|
||||
correctness: question_id -> 基线是否答对。
|
||||
config: 池构建配置(使用 train_ratio, seed)。
|
||||
|
||||
返回:
|
||||
{task_type: {"train": [qid, ...], "val": [qid, ...]}}。
|
||||
"""
|
||||
target_types = set(new_task_types)
|
||||
groups: dict[str, list[GeneratedQuestion]] = defaultdict(list)
|
||||
for q in questions:
|
||||
if q.task_type in target_types:
|
||||
groups[q.task_type].append(q)
|
||||
|
||||
result: dict[str, dict[str, list[str]]] = {}
|
||||
rng = random.Random(config.seed)
|
||||
|
||||
for task_type in sorted(groups.keys()):
|
||||
train, val = self._split_one_category(
|
||||
groups[task_type], correctness, config.train_ratio, rng,
|
||||
)
|
||||
result[task_type] = {
|
||||
"train": [q.question_id for q in train],
|
||||
"val": [q.question_id for q in val],
|
||||
}
|
||||
|
||||
return result
|
||||
|
||||
@@ -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