diff --git a/app/harness/pools.py b/app/harness/pools.py index 66e5769..9c0cda5 100644 --- a/app/harness/pools.py +++ b/app/harness/pools.py @@ -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 diff --git a/tests/unit/test_harness_pools.py b/tests/unit/test_harness_pools.py index c148a40..3a2da7f 100644 --- a/tests/unit/test_harness_pools.py +++ b/tests/unit/test_harness_pools.py @@ -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"}