Files
Video-Tree-TRM5/research-wiki/plans/2026-07-12-per-category-pool-strategy.md
T

50 KiB
Raw Blame History

Per-Category Pool Strategy 实现计划

For agentic workers: REQUIRED SUB-SKILL: Use subagent-driven-development to implement this plan task-by-task. Steps use checkbox (- [ ]) syntax for tracking.

Goal: 实现 PoolStrategy Protocol + PerCategoryPoolStrategy,支持按类增量 baseline inference、per-category correctness 分层划分(20 train / 10 val)、task_types 作用域训练,以及 train 模式 CLI 接线。

Architecture:app/ports.py 新增 PoolStrategy Protocol,在 app/harness/pools.py 实现 GlobalPoolStrategy(封装现有逻辑)和 PerCategoryPoolStrategyper-category 2:1 分层)。RunConfig 新增 4 字段,main.py 完成 train 模式接线和 strategy 组装。

Tech Stack: Python 3.11, dataclasses, Protocol, pytest, SQLite


Task 1: core/types.py — 新增 PoolConfig

Files:

  • Modify: core/types.py:57 (在 GeneratedQuestion 之后追加)

  • Test: tests/unit/test_core_types.py

  • Step 1: 写失败测试

tests/unit/test_core_types.py 末尾追加:

from core.types import PoolConfig
from pathlib import Path


class TestPoolConfig:
    """PoolConfig frozen dataclass 基本行为。"""

    def test_pool_config_frozen(self) -> None:
        """PoolConfig 创建后不可变。"""
        cfg = PoolConfig(
            task_types=("Action Reasoning",),
            seed=42,
            baseline_run_id="baseline_v2",
            diag_size=200,
            diag_correct_ratio=0.5,
            val_size=30,
            val_correct_ratio=0.5,
            test_size=60,
            eval_min_per_class=2,
            train_ratio=0.667,
            test_questions_dir=Path("store/questions/benchmarks/Video-MME"),
        )
        assert cfg.task_types == ("Action Reasoning",)
        assert cfg.baseline_run_id == "baseline_v2"
        assert cfg.train_ratio == 0.667

    def test_pool_config_task_types_none(self) -> None:
        """task_types=None 表示全部类别。"""
        cfg = PoolConfig(
            task_types=None,
            seed=0,
            baseline_run_id="run_1",
            diag_size=200,
            diag_correct_ratio=0.5,
            val_size=30,
            val_correct_ratio=0.5,
            test_size=60,
            eval_min_per_class=2,
            train_ratio=0.667,
            test_questions_dir=None,
        )
        assert cfg.task_types is None
        assert cfg.test_questions_dir is None
  • Step 2: 运行测试确认失败

Run: conda activate Video-Tree-TRM && pytest tests/unit/test_core_types.py::TestPoolConfig -v Expected: FAIL with ImportError: cannot import name 'PoolConfig'

  • Step 3: 实现 PoolConfig

core/types.py 末尾追加:

from pathlib import Path as _Path  # 避免与运行时 TYPE_CHECKING 冲突


@dataclass(frozen=True)
class PoolConfig:
    """池构建策略的统一配置。

    两组字段由两个具体策略各取所需,未使用的字段被忽略。

    属性:
        task_types: 限定题型元组;None 表示全部类别。
        seed: 随机种子,保证可复现。
        baseline_run_id: 基线 run 标识(用于读 correctness、写入 pools 指纹)。
        diag_size: 诊断池大小(GlobalStrategy 用)。
        diag_correct_ratio: 诊断池中对题占比(GlobalStrategy 用)。
        val_size: 验证池大小(GlobalStrategy 用)。
        val_correct_ratio: 验证池中对题占比(GlobalStrategy 用)。
        test_size: held-out 测试池大小(GlobalStrategy 用)。
        eval_min_per_class: 验证池中每类保底样本数(GlobalStrategy 用)。
        train_ratio: train/(train+val) 比例(PerCategoryStrategy 用),默认 2/3。
        test_questions_dir: 外部 test 题源路径(PerCategoryStrategy 用)。
    """

    task_types: tuple[str, ...] | None
    seed: int
    baseline_run_id: str
    diag_size: int
    diag_correct_ratio: float
    val_size: int
    val_correct_ratio: float
    test_size: int
    eval_min_per_class: int
    train_ratio: float
    test_questions_dir: _Path | None
  • Step 4: 运行测试确认通过

Run: conda activate Video-Tree-TRM && pytest tests/unit/test_core_types.py::TestPoolConfig -v Expected: PASS

  • Step 5: 提交
git add core/types.py tests/unit/test_core_types.py

Commit message: feat(core): add PoolConfig dataclass for pool strategy configuration


Task 2: app/ports.py — 新增 PoolStrategy Protocol

Files:

  • Modify: app/ports.py:147 (文件末尾追加)

  • Test: tests/unit/test_core_protocols.py

  • Step 1: 写失败测试

tests/unit/test_core_protocols.py 末尾追加:

from app.ports import PoolStrategy


class TestPoolStrategyProtocol:
    """PoolStrategy Protocol runtime_checkable 验证。"""

    def test_pool_strategy_is_runtime_checkable(self) -> None:
        """PoolStrategy 支持 isinstance 检查。"""
        from app.harness.pools import Pools
        from core.types import GeneratedQuestion, PoolConfig

        class FakeStrategy:
            def build(self, questions, correctness, config):
                return Pools(
                    diagnosis=[], validation=[], test=[],
                    baseline_run_id="", baseline_val_accuracy=0.0,
                )

            def build_incremental(self, new_task_types, questions, correctness, config):
                return {}

        assert isinstance(FakeStrategy(), PoolStrategy)
  • Step 2: 运行测试确认失败

Run: conda activate Video-Tree-TRM && pytest tests/unit/test_core_protocols.py::TestPoolStrategyProtocol -v Expected: FAIL with ImportError: cannot import name 'PoolStrategy'

  • Step 3: 实现 PoolStrategy Protocol

app/ports.py 末尾追加:

@runtime_checkable
class PoolStrategy(Protocol):
    """池构建策略端口。

    应用层端口(非 core 层),因为返回类型 Pools 定义在 app/harness/pools.py。
    两个具体策略(GlobalPoolStrategy / PerCategoryPoolStrategy)实现此接口。
    """

    def build(
        self,
        questions: list[GeneratedQuestion],
        correctness: dict[str, bool],
        config: PoolConfig,
    ) -> Pools: ...

    def build_incremental(
        self,
        new_task_types: list[str],
        questions: list[GeneratedQuestion],
        correctness: dict[str, bool],
        config: PoolConfig,
    ) -> dict[str, dict[str, list[str]]]: ...

同时在文件头部的 TYPE_CHECKING 块中添加:

from app.harness.pools import Pools
from core.types import PoolConfig

注意:PoolStrategybuild 方法返回 Pools,但 Poolsapp/harness/pools.py 中定义。因为 app/ports.pyapp/harness/pools.py 同属 app 层,不违反依赖方向。使用 TYPE_CHECKING 保护导入以避免循环引用。

  • Step 4: 运行测试确认通过

Run: conda activate Video-Tree-TRM && pytest tests/unit/test_core_protocols.py::TestPoolStrategyProtocol -v Expected: PASS

  • Step 5: 提交
git add app/ports.py tests/unit/test_core_protocols.py

Commit message: feat(app): add PoolStrategy Protocol to application ports


Task 3: app/harness/pools.py — GlobalPoolStrategy 封装

Files:

  • Modify: app/harness/pools.py

  • Test: tests/unit/test_harness_pools.py

  • Step 1: 写失败测试

tests/unit/test_harness_pools.py 追加:

from app.harness.pools import GlobalPoolStrategy
from core.types import PoolConfig
from pathlib import Path


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 不支持增量,调用 build_incremental 应报错。"""
        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)
  • Step 2: 运行测试确认失败

Run: conda activate Video-Tree-TRM && pytest tests/unit/test_harness_pools.py::TestGlobalPoolStrategy -v Expected: FAIL with ImportError: cannot import name 'GlobalPoolStrategy'

  • Step 3: 实现 GlobalPoolStrategy

app/harness/pools.py 中,现有 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 函数。"""
        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]]]:
        """全局策略不支持增量。"""
        raise NotImplementedError("GlobalPoolStrategy 不支持增量构建,请使用 PerCategoryPoolStrategy。")

在文件头部添加导入:

from core.types import PoolConfig
  • Step 4: 运行测试确认通过

Run: conda activate Video-Tree-TRM && pytest tests/unit/test_harness_pools.py::TestGlobalPoolStrategy -v Expected: PASS

  • Step 5: 运行全部现有池测试确认无回归

Run: conda activate Video-Tree-TRM && pytest tests/unit/test_harness_pools.py -v Expected: 全部 PASS

  • Step 6: 提交
git add app/harness/pools.py tests/unit/test_harness_pools.py

Commit message: refactor(harness): wrap existing build_pools in GlobalPoolStrategy


Task 4: app/harness/pools.py — PerCategoryPoolStrategy

Files:

  • Modify: app/harness/pools.py

  • Test: tests/unit/test_harness_pools.py

  • Step 1: 写失败测试

tests/unit/test_harness_pools.py 追加:

from app.harness.pools import PerCategoryPoolStrategy


def _make_per_category_questions() -> list[GeneratedQuestion]:
    """构造 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) -> None:
        """每类 30 题按 correctness 2:1 分层 → 20 train + 10 val。"""
        questions = _make_per_category_questions()
        # 每类前 18 题 correct,后 12 题 wrong
        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)

        # 12 类 × 20 = 240 train, 12 类 × 10 = 120 val
        assert len(pools.diagnosis) == 240
        assert len(pools.validation) == 120

        # 逐类验证 train=20, val=10
        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, f"{tt} train 应为 20,实际 {diag_counts[tt]}"
            assert val_counts[tt] == 10, f"{tt} val 应为 10,实际 {val_counts[tt]}"

        # train 和 val 互斥
        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) -> None:
        """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  # 18/30 = 60% correct

        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])
        # 18 correct: floor(18 * 20/30) = 12 train, 6 val
        assert diag_correct == 12
        assert val_correct == 6

    def test_per_category_all_correct_degrades(self) -> None:
        """某类全部 correct0 wrong)时退化为非分层 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) -> None:
        """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]}  # 缺 5 题

        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) -> None:
        """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  # 2 类 × 20
        assert len(pools.validation) == 20  # 2 类 × 10
        types_in_diag = {q.task_type for q in pools.diagnosis}
        assert types_in_diag == {"Action Reasoning", "Scene Understanding"}
  • Step 2: 运行测试确认失败

Run: conda activate Video-Tree-TRM && pytest tests/unit/test_harness_pools.py::TestPerCategoryPoolStrategy -v Expected: FAIL with ImportError: cannot import name 'PerCategoryPoolStrategy'

  • Step 3: 实现 PerCategoryPoolStrategy

app/harness/pools.py 中追加:

import math
from loguru import logger


class PerCategoryPoolStrategy:
    """Per-category correctness 分层策略。

    每个 task_type 内部按 correct/wrong 分层,
    各自按 train_ratio 比例分配到 train(diagnosis) 和 val 池。
    test 从外部 benchmark 加载。
    """

    def build(
        self,
        questions: list[GeneratedQuestion],
        correctness: dict[str, bool],
        config: PoolConfig,
    ) -> Pools:
        """按类别分层构建 train/val 池。

        参数:
            questions: 题目全集(可含多类别)。
            correctness: question_id -> 基线是否答对。
            config: 池配置(使用 task_types, seed, train_ratio, test_questions_dir)。

        返回:
            Poolsdiagnosis=train 合并, validation=val 合并, test=外部 benchmark)。
        """
        filtered = questions
        if config.task_types is not None:
            allowed = set(config.task_types)
            filtered = [q for q in questions if q.task_type in allowed]

        # 按 task_type 分组
        by_type: dict[str, list[GeneratedQuestion]] = {}
        for q in filtered:
            by_type.setdefault(q.task_type, []).append(q)

        rng = random.Random(config.seed)
        all_train: list[GeneratedQuestion] = []
        all_val: list[GeneratedQuestion] = []

        for task_type in sorted(by_type):
            type_qs = by_type[task_type]
            train_qs, val_qs = self._split_one_category(
                type_qs, correctness, config.train_ratio, rng, task_type,
            )
            all_train.extend(train_qs)
            all_val.extend(val_qs)

        # test 从外部 benchmark 加载
        test_qs: list[GeneratedQuestion] = []
        if config.test_questions_dir is not None:
            from app.question_gen import load_benchmark
            all_test = load_benchmark(config.test_questions_dir)
            if config.task_types is not None:
                allowed = set(config.task_types)
                test_qs = [q for q in all_test if q.task_type in allowed]
            else:
                test_qs = all_test
            if not test_qs:
                logger.warning("test 池为空:test_questions_dir 中无匹配的 task_type")

        # baseline_val_accuracy
        val_correct = sum(1 for q in all_val if correctness.get(q.question_id))
        baseline_val_acc = val_correct / len(all_val) if all_val else 0.0

        # 汇总 correctness
        all_correctness = {
            q.question_id: correctness.get(q.question_id, False)
            for q in all_train + all_val + test_qs
        }

        return Pools(
            diagnosis=all_train,
            validation=all_val,
            test=test_qs,
            baseline_run_id=config.baseline_run_id,
            baseline_val_accuracy=baseline_val_acc,
            correctness=all_correctness,
        )

    def _split_one_category(
        self,
        questions: list[GeneratedQuestion],
        correctness: dict[str, bool],
        train_ratio: float,
        rng: random.Random,
        task_type: str,
    ) -> tuple[list[GeneratedQuestion], list[GeneratedQuestion]]:
        """单类别按 correctness 分层划分。

        参数:
            questions: 该类别全部题目。
            correctness: question_id -> 基线是否答对。
            train_ratio: train/(train+val) 比例。
            rng: 随机数发生器。
            task_type: 类别名(用于日志)。

        返回:
            (train 题目列表, val 题目列表)。
        """
        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"类别 {task_type!r} 的 correctness 缺失 {len(missing)} 题: "
                + ", ".join(missing[:5])
                + ("..." if len(missing) > 5 else "")
            )

        correct = [q for q in questions if correctness[q.question_id]]
        wrong = [q for q in questions if not correctness[q.question_id]]

        # 边界:全 correct 或全 wrong → 退化为非分层
        if not wrong or not correct:
            logger.warning(
                "类别 {!r} 全部 {} ({} 题),退化为非分层 random {}/{}",
                task_type,
                "correct" if not wrong else "wrong",
                n_total, n_train, n_val,
            )
            shuffled = list(questions)
            rng.shuffle(shuffled)
            return shuffled[:n_train], shuffled[n_train:]

        # 按比例分配 correct/wrong 到 train
        n_correct = len(correct)
        train_correct = math.floor(n_correct * n_train / n_total)
        train_wrong = n_train - train_correct
        val_correct = n_correct - train_correct
        val_wrong = len(wrong) - train_wrong

        assert train_correct + train_wrong == n_train
        assert val_correct + val_wrong == n_val

        rng.shuffle(correct)
        rng.shuffle(wrong)

        train_qs = correct[:train_correct] + wrong[:train_wrong]
        val_qs = correct[train_correct:] + wrong[train_wrong:]
        return train_qs, val_qs

    def build_incremental(
        self,
        new_task_types: list[str],
        questions: list[GeneratedQuestion],
        correctness: dict[str, bool],
        config: PoolConfig,
    ) -> dict[str, dict[str, list[str]]]:
        """增量构建新类别的 train/val 划分。

        参数:
            new_task_types: 待新增的类别列表。
            questions: 题目全集。
            correctness: question_id -> 基线是否答对。
            config: 池配置。

        返回:
            {task_type: {"train": [qid, ...], "val": [qid, ...]}}。
        """
        rng = random.Random(config.seed)
        result: dict[str, dict[str, list[str]]] = {}

        by_type: dict[str, list[GeneratedQuestion]] = {}
        for q in questions:
            if q.task_type in new_task_types:
                by_type.setdefault(q.task_type, []).append(q)

        for task_type in sorted(by_type):
            train_qs, val_qs = self._split_one_category(
                by_type[task_type], correctness, config.train_ratio, rng, task_type,
            )
            result[task_type] = {
                "train": [q.question_id for q in train_qs],
                "val": [q.question_id for q in val_qs],
            }
        return result

在文件头部添加 import mathimport randomrandom 已有,math 需新增)。

  • Step 4: 运行测试确认通过

Run: conda activate Video-Tree-TRM && pytest tests/unit/test_harness_pools.py::TestPerCategoryPoolStrategy -v Expected: 全部 PASS

  • Step 5: 运行全部池测试确认无回归

Run: conda activate Video-Tree-TRM && pytest tests/unit/test_harness_pools.py -v Expected: 全部 PASS

  • Step 6: 提交
git add app/harness/pools.py tests/unit/test_harness_pools.py

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


Task 5: app/harness/config.py — RunConfig 新增字段(原 Task 6,提前以消除前向引用)

Files:

  • Modify: app/harness/pools.py:170-288

  • Test: tests/unit/test_harness_pools.py

  • Step 1: 写失败测试

tests/unit/test_harness_pools.py 追加:

class TestPerCategorySaveLoad:
    """per_category 格式的 pools.json 冻结/加载。"""

    def test_save_load_per_category_roundtrip(self, tmp_path: Path) -> None:
        """per_category 模式 save → load 往返一致。"""
        questions = _make_per_category_questions()
        correctness = {q.question_id: True for q in questions}
        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)

        pools_path = tmp_path / "pools.json"
        save_pools(pools, pools_path, split_mode="per_category", config=config)
        loaded = load_pools(pools_path)

        assert loaded.baseline_run_id == pools.baseline_run_id
        assert len(loaded.diagnosis) == len(pools.diagnosis)
        assert len(loaded.validation) == len(pools.validation)

    def test_load_per_category_rejects_mismatched_config(self, tmp_path: Path) -> None:
        """加载时 seed/train_ratio 不匹配 → 报错。"""
        data = {
            "split_mode": "per_category",
            "train_ratio": 0.667,
            "seed": 42,
            "baseline_run_id": "r1",
            "baseline_val_accuracy": 0.5,
            "correctness": {},
            "categories": {},
            "diagnosis": [],
            "validation": [],
            "test": [],
        }
        pools_path = tmp_path / "pools.json"
        pools_path.write_text(json.dumps(data), encoding="utf-8")

        # 正常加载应通过
        loaded = load_pools(pools_path)
        assert loaded.baseline_run_id == "r1"
  • Step 2: 运行测试确认失败

Run: conda activate Video-Tree-TRM && pytest tests/unit/test_harness_pools.py::TestPerCategorySaveLoad -v Expected: FAILsave_pools 不接受 split_mode 参数)

  • Step 3: 扩展 save_pools 和 load_pools

修改 save_pools 签名和实现,支持 per_category 格式:

def save_pools(
    pools: Pools,
    path: Path,
    *,
    split_mode: str = "global",
    config: PoolConfig | None = None,
) -> None:
    """将三池及基线指标冻结为 JSON。

    参数:
        pools: 待冻结的三池。
        path: 目标 JSON 文件路径。
        split_mode: 分割模式标签("global" / "per_category")。
        config: 池配置(per_category 模式下需要记录 seed/train_ratio)。
    """
    data: dict = {
        "split_mode": split_mode,
        "baseline_run_id": pools.baseline_run_id,
        "baseline_val_accuracy": pools.baseline_val_accuracy,
        "correctness": pools.correctness,
        "diagnosis": [_q_to_dict(q) for q in pools.diagnosis],
        "validation": [_q_to_dict(q) for q in pools.validation],
        "test": [_q_to_dict(q) for q in pools.test],
    }
    if split_mode == "per_category" and config is not None:
        # 按类别记录 train/val qid 分配
        categories: dict[str, dict[str, list[str]]] = {}
        for q in pools.diagnosis:
            categories.setdefault(q.task_type, {"train": [], "val": []})["train"].append(
                q.question_id
            )
        for q in pools.validation:
            categories.setdefault(q.task_type, {"train": [], "val": []})["val"].append(
                q.question_id
            )
        data["categories"] = categories
        data["seed"] = config.seed
        data["train_ratio"] = config.train_ratio
        if config.test_questions_dir is not None:
            data["test_source"] = str(config.test_questions_dir)
    path.write_text(
        json.dumps(data, ensure_ascii=False, indent=2),
        encoding="utf-8",
    )

修改 load_pools 以兼容新旧格式:

def load_pools(path: Path) -> Pools:
    """从 JSON 恢复冻结的三池。兼容 global 和 per_category 格式。"""
    d = json.loads(path.read_text(encoding="utf-8"))
    if "test" not in d:
        raise ValueError(
            f"{path} 为旧格式 pools.json(缺 test 池),"
            "请删除后重新切分。"
        )
    return Pools(
        diagnosis=[_dict_to_q(x) for x in d["diagnosis"]],
        validation=[_dict_to_q(x) for x in d["validation"]],
        test=[_dict_to_q(x) for x in d["test"]],
        baseline_run_id=d["baseline_run_id"],
        baseline_val_accuracy=d["baseline_val_accuracy"],
        correctness=d.get("correctness", {}),
    )
  • Step 4: 重构 build_or_load_pools 接受 strategy
def build_or_load_pools(
    config: RunConfig,
    strategy: PoolStrategy,
    db_path: Path,
) -> Pools:
    """构建或加载三池。

    参数:
        config: 运行配置。
        strategy: 池构建策略实例。
        db_path: harness.db 路径(用于读取 baseline correctness)。

    返回:
        冻结的三池 Pools。
    """
    from app.harness.log import HarnessLog
    from app.harness.workspace import resolve_paths
    from app.question_gen import load_benchmark

    pools_path = config.workspace_dir / "pools.json"

    # baseline_run_id 从 seed 解析(train 的 config.run_id 是训练 ID,非 baseline
    from app.harness.workspace import resolve_paths as _resolve
    _paths = _resolve(config.workspace_dir)
    manifest = json.loads((_paths.workspace_dir / "manifest.json").read_text(encoding="utf-8"))
    baseline_run_id = manifest.get("baseline_run_id", config.run_id or "infer_adhoc")

    if pools_path.exists():
        loaded_data = json.loads(pools_path.read_text(encoding="utf-8"))
        stored_mode = loaded_data.get("split_mode", "global")
        if stored_mode == "per_category":
            # 一致性校验
            if loaded_data.get("seed") != getattr(_to_pool_config(config), "seed", None):
                raise ValueError(
                    f"pools.json 的 seed({loaded_data.get('seed')}) 与当前配置不一致,"
                    "请删除 pools.json 重建。"
                )
            # 检查是否有新类别需要增量
            existing_categories = set(loaded_data.get("categories", {}).keys())
            requested = set(config.task_types) if config.task_types else set()
            new_types = requested - existing_categories
            if new_types:
                paths = resolve_paths(config.workspace_dir)
                questions = load_benchmark(paths.questions_dir)
                with HarnessLog(str(db_path), baseline_run_id) as log:
                    rows = log.query(
                        "SELECT question_id, prediction, answer FROM predictions WHERE run_id=?",
                        (baseline_run_id,),
                    )
                correctness = {r["question_id"]: r["prediction"] == r["answer"] for r in rows}
                pool_config = _to_pool_config(config)
                incremental = strategy.build_incremental(
                    list(new_types), questions, correctness, pool_config,
                )
                loaded_data["categories"].update(incremental)
                # 重建 pools 并保存
                pools_path.write_text(
                    json.dumps(loaded_data, ensure_ascii=False, indent=2),
                    encoding="utf-8",
                )
        return load_pools(pools_path)

    # 首次构建
    paths = resolve_paths(config.workspace_dir)
    questions = load_benchmark(paths.questions_dir)
    with HarnessLog(str(db_path), baseline_run_id) as log:
        rows = log.query(
            "SELECT question_id, prediction, answer FROM predictions WHERE run_id=?",
            (baseline_run_id,),
        )
    correctness = {r["question_id"]: r["prediction"] == r["answer"] for r in rows}
    pool_config = _to_pool_config(config)
    pools = strategy.build(questions, correctness, pool_config)
    save_pools(pools, pools_path, split_mode=config.pool_split_mode, config=pool_config)
    return pools

新增辅助函数 _to_pool_config

def _to_pool_config(config: RunConfig) -> PoolConfig:
    """从 RunConfig 提取 PoolConfig。"""
    return PoolConfig(
        task_types=config.task_types,
        seed=0,
        baseline_run_id=config.run_id if config.run_id else "infer_adhoc",
        diag_size=config.diag_size,
        diag_correct_ratio=config.diag_correct_ratio,
        val_size=config.val_size,
        val_correct_ratio=config.val_correct_ratio,
        test_size=config.test_size,
        eval_min_per_class=config.eval_min_per_class,
        train_ratio=config.train_ratio,
        test_questions_dir=(
            Path(config.store_dir) / "questions" / config.test_questions
            if hasattr(config, "test_questions") and config.test_questions
            else None
        ),
    )
  • Step 5: 运行测试确认通过

Run: conda activate Video-Tree-TRM && pytest tests/unit/test_harness_pools.py -v Expected: 全部 PASS

  • Step 6: 提交
git add app/harness/pools.py tests/unit/test_harness_pools.py

Commit message: feat(harness): refactor build_or_load_pools to accept PoolStrategy + per_category freeze format


Task 6: pools.py — 重构 build_or_load_pools + per_category 冻结格式(原 Task 5,后移以依赖 Task 5 的 RunConfig 字段)

Files:

  • Modify: app/harness/config.py:131-138 (在有默认值字段区追加)

  • Modify: app/harness/config.py:259-298 (校验逻辑调整)

  • Test: tests/unit/test_harness_config.py

  • Step 1: 写失败测试

tests/unit/test_harness_config.py 末尾追加:

class TestRunConfigNewFields:
    """RunConfig 新增字段校验。"""

    def test_pool_split_mode_valid(self, base_config: dict) -> None:
        """pool_split_mode 合法值。"""
        base_config["pool_split_mode"] = "per_category"
        cfg = RunConfig(**base_config)
        assert cfg.pool_split_mode == "per_category"

    def test_pool_split_mode_invalid(self, base_config: dict) -> None:
        """pool_split_mode 非法值应报错。"""
        base_config["pool_split_mode"] = "invalid"
        from app.harness.config import _validate
        cfg = RunConfig(**base_config)
        with pytest.raises(ValueError, match="pool_split_mode"):
            _validate(cfg)

    def test_task_types_tuple(self, base_config: dict) -> None:
        """task_types 接受 tuple。"""
        base_config["task_types"] = ("Action Reasoning", "Scene Understanding")
        cfg = RunConfig(**base_config)
        assert cfg.task_types == ("Action Reasoning", "Scene Understanding")

    def test_task_types_none(self, base_config: dict) -> None:
        """task_types 默认 None。"""
        cfg = RunConfig(**base_config)
        assert cfg.task_types is None

    def test_train_ratio_range(self, base_config: dict) -> None:
        """train_ratio 必须在 (0, 1) 内。"""
        base_config["train_ratio"] = 1.5
        from app.harness.config import _validate
        cfg = RunConfig(**base_config)
        with pytest.raises(ValueError, match="train_ratio"):
            _validate(cfg)

    def test_per_category_skips_val_size_check(self, base_config: dict) -> None:
        """per_category 模式跳过 val_size >= eval_min_per_class * 12 的校验。"""
        base_config["pool_split_mode"] = "per_category"
        base_config["val_size"] = 1  # 远小于 12*2=24,全局模式会报错
        from app.harness.config import _validate
        cfg = RunConfig(**base_config)
        _validate(cfg)  # 不应报错

注意:需要检查已有测试的 base_config fixture,确保它包含新字段的默认值。

  • Step 2: 运行测试确认失败

Run: conda activate Video-Tree-TRM && pytest tests/unit/test_harness_config.py::TestRunConfigNewFields -v Expected: FAIL

  • Step 3: 修改 RunConfig

app/harness/config.pyRunConfig 有默认值区(line 131-138 之后)追加:

    task_types: tuple[str, ...] | None = None
    pool_split_mode: str = "global"
    train_ratio: float = 0.667
    test_questions: str = "benchmarks/Video-MME"

_VALID_POOL_SPLIT_MODES 常量区追加:

_VALID_POOL_SPLIT_MODES = {"global", "per_category"}
  • Step 4: 修改校验逻辑

_validate_basic 中追加 pool_split_modetrain_ratio 校验:

    if config.pool_split_mode not in _VALID_POOL_SPLIT_MODES:
        raise ValueError(
            f"pool_split_mode 必须为 {_VALID_POOL_SPLIT_MODES} 之一,"
            f"实际: {config.pool_split_mode!r}"
        )
    if not (0 < config.train_ratio < 1):
        raise ValueError(
            f"train_ratio 必须在 (0, 1) 内,实际: {config.train_ratio}"
        )

修改 _validate_minibatch 中的 val_size 校验,per_category 模式跳过:

将现有的:

    floor = config.eval_min_per_class * _VIDEO_MME_TASK_TYPE_COUNT
    if config.val_size < floor:

改为:

    if config.pool_split_mode != "per_category":
        floor = config.eval_min_per_class * _VIDEO_MME_TASK_TYPE_COUNT
        if config.val_size < floor:
  • Step 5: 更新 base_config fixture

tests/unit/test_harness_config.pybase_config fixture 中追加新字段默认值:

    "task_types": None,
    "pool_split_mode": "global",
    "train_ratio": 0.667,
    "test_questions": "benchmarks/Video-MME",
  • Step 6: 运行测试确认通过

Run: conda activate Video-Tree-TRM && pytest tests/unit/test_harness_config.py -v Expected: 全部 PASS

  • Step 7: 提交
git add app/harness/config.py tests/unit/test_harness_config.py

Commit message: feat(harness): add task_types, pool_split_mode, train_ratio, test_questions to RunConfig


Task 7: config/default.yaml — 新增配置项

Files:

  • Modify: config/default.yaml

  • Step 1: 在 harness 段末尾追加

config/default.yamlharness:early_stop_patience: 8use_slow_momentum: true 之后追加:

  # 池构建策略
  pool_split_mode: global          # global | per_category
  train_ratio: 0.667               # per_category 模式下 train/(train+val) 比例
  test_questions: "benchmarks/Video-MME"  # test 池的题目来源
  • Step 2: 验证 YAML 可解析

Run: conda activate Video-Tree-TRM && python -c "import yaml; yaml.safe_load(open('config/default.yaml'))" Expected: 无输出(成功)

  • Step 3: 提交
git add config/default.yaml

Commit message: config: add pool_split_mode, train_ratio, test_questions to default.yaml


Task 8: app/harness/log.py — _runs 表 upsert

Files:

  • Modify: app/harness/log.py:72-78

  • Step 1: 写失败测试

tests/unit/ 新建或追加到已有 test 文件:

# tests/unit/test_harness_log.py

from app.harness.log import HarnessLog


class TestHarnessLogUpsert:
    """_runs 表 upsert 行为。"""

    def test_same_run_id_updates_started_at(self, tmp_path) -> None:
        """同 run_id 第二次创建 HarnessLog 应更新 started_at。"""
        db = str(tmp_path / "test.db")
        with HarnessLog(db, "run_1", git_sha="abc") as log1:
            rows = log1.query("SELECT started_at FROM _runs WHERE run_id='run_1'")
            first_time = rows[0]["started_at"]

        import time
        time.sleep(0.01)

        with HarnessLog(db, "run_1", git_sha="abc") as log2:
            rows = log2.query("SELECT started_at FROM _runs WHERE run_id='run_1'")
            second_time = rows[0]["started_at"]

        assert second_time > first_time
        # 应只有一行
        with HarnessLog(db, "run_1") as log3:
            rows = log3.query("SELECT COUNT(*) as cnt FROM _runs WHERE run_id='run_1'")
            assert rows[0]["cnt"] == 1
  • Step 2: 运行测试确认失败

Run: conda activate Video-Tree-TRM && pytest tests/unit/test_harness_log.py::TestHarnessLogUpsert -v Expected: FAIL(当前 INSERT OR IGNORE 不更新 started_at

  • Step 3: 修改 HarnessLog.init

app/harness/log.py line 72-78 的:

        self._conn.execute(
            "INSERT OR IGNORE INTO _runs"
            " (run_id, git_sha, started_at, config, status)"
            " VALUES (?, ?, ?, ?, ?)",
            (run_id, resolved_sha, _now_iso(), config_json, "running"),
        )

改为:

        self._conn.execute(
            "INSERT INTO _runs"
            " (run_id, git_sha, started_at, config, status)"
            " VALUES (?, ?, ?, ?, ?)"
            " ON CONFLICT(run_id) DO UPDATE SET"
            " started_at=excluded.started_at,"
            " config=excluded.config,"
            " status=excluded.status",
            (run_id, resolved_sha, _now_iso(), config_json, "running"),
        )
  • Step 4: 运行测试确认通过

Run: conda activate Video-Tree-TRM && pytest tests/unit/test_harness_log.py -v Expected: PASS

  • Step 5: 提交
git add app/harness/log.py tests/unit/test_harness_log.py

Commit message: fix(harness): change _runs INSERT OR IGNORE to ON CONFLICT DO UPDATE for incremental infer


Task 9: main.py — task_types 纳入 config + train 接线

Files:

  • Modify: main.py:258,293-298

  • Step 1: 将 task_types 纳入 cli_overrides

修改 main.py line 258 的:

    cli_overrides = {k: v for k, v in vars(args).items() if k not in ("config", "task_types")}

改为:

    cli_args = vars(args)
    # task_types: list -> tupleRunConfig 要求 tuple
    if cli_args.get("task_types") is not None:
        cli_args["task_types"] = tuple(cli_args["task_types"])
    cli_overrides = {k: v for k, v in cli_args.items() if k != "config"}
  • Step 2: 添加 --pool-split-mode CLI 参数

_build_parser() 中追加:

    parser.add_argument(
        "--pool-split-mode",
        choices=["global", "per_category"],
        dest="pool_split_mode",
    )
    parser.add_argument("--train-ratio", type=float, dest="train_ratio")
    parser.add_argument("--test-questions", type=str, dest="test_questions")
  • Step 3: 接线 train 模式

修改 main.py line 293-298 的:

    if config.mode == "infer":
        task_types = getattr(args, "task_types", None)
        result = asyncio.run(runner.infer(task_types=task_types))
        _log_result(result)
    else:
        raise SystemExit(f"模式 {config.mode!r} 尚未实现")

改为:

    if config.mode == "infer":
        result = asyncio.run(runner.infer(task_types=config.task_types))
        _log_result(result)
    elif config.mode == "train":
        from app.harness.pools import (
            GlobalPoolStrategy,
            PerCategoryPoolStrategy,
            build_or_load_pools,
        )
        from app.harness.workspace import resolve_paths

        strategy = (
            PerCategoryPoolStrategy()
            if config.pool_split_mode == "per_category"
            else GlobalPoolStrategy()
        )
        paths = resolve_paths(config.workspace_dir)
        pools = build_or_load_pools(config, strategy, paths.db_path)
        asyncio.run(runner.train(pools))
    else:
        raise SystemExit(f"模式 {config.mode!r} 尚未实现")
  • Step 4: 验证 CLI 解析

Run: conda activate Video-Tree-TRM && python main.py harness --help Expected: 输出中包含 --pool-split-mode, --train-ratio, --test-questions, --task-types

  • Step 5: 提交
git add main.py

Commit message: feat(cli): wire train mode with PoolStrategy selection and task_types in RunConfig


Task 10: 集成测试

Files:

  • Create: tests/integration/test_pool_strategy.py

  • Step 1: 编写集成测试

"""PerCategoryPoolStrategy 端到端集成测试。

验证从构造题目 → 伪造 baseline → 池构建 → 冻结 → 加载的完整流程。
"""

from __future__ import annotations

import json
from collections import Counter
from pathlib import Path

import pytest

from app.harness.pools import (
    PerCategoryPoolStrategy,
    load_pools,
    save_pools,
)
from core.types import GeneratedQuestion, PoolConfig


def _make_question(qid: str, task_type: str) -> GeneratedQuestion:
    """构造测试用 GeneratedQuestion。"""
    return GeneratedQuestion(
        question_id=qid, video_id="v1", task_type=task_type,
        question=f"Q {qid}?",
        options=("A. a", "B. b", "C. c", "D. d"),
        answer="A", source_nodes=("n1",), difficulty="medium",
    )


class TestPerCategoryE2E:
    """端到端:构建 → 冻结 → 加载 → 校验。"""

    def test_full_flow(self, tmp_path: Path) -> None:
        """完整流程:12 类各 30 题 → 策略构建 → 冻结 → 加载 → 三池校验。"""
        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))

        # 每类前 18 correct,后 12 wrong
        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

        # 验证逐类均匀
        diag_by_type = Counter(q.task_type for q in pools.diagnosis)
        val_by_type = Counter(q.task_type for q in pools.validation)
        for tt in task_types:
            assert diag_by_type[tt] == 20
            assert val_by_type[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()

        # 验证 correctness 对齐
        for tt in task_types:
            tt_diag = [q for q in pools.diagnosis if q.task_type == tt]
            tt_val = [q for q in pools.validation if q.task_type == tt]
            diag_ratio = sum(1 for q in tt_diag if correctness[q.question_id]) / len(tt_diag)
            val_ratio = sum(1 for q in tt_val if correctness[q.question_id]) / len(tt_val)
            assert abs(diag_ratio - val_ratio) < 0.05, (
                f"{tt}: train ratio {diag_ratio:.2f} vs val ratio {val_ratio:.2f}"
            )

        # 冻结 → 加载
        pools_path = tmp_path / "pools.json"
        save_pools(pools, pools_path, split_mode="per_category", config=config)
        loaded = load_pools(pools_path)
        assert len(loaded.diagnosis) == 240
        assert len(loaded.validation) == 120

        # 验证冻结格式
        data = json.loads(pools_path.read_text())
        assert data["split_mode"] == "per_category"
        assert "categories" in data
        assert len(data["categories"]) == 12
  • Step 2: 运行测试

Run: conda activate Video-Tree-TRM && pytest tests/integration/test_pool_strategy.py -v Expected: PASS

  • Step 3: 运行全量测试确认无回归

Run: conda activate Video-Tree-TRM && pytest tests/unit/test_harness_pools.py tests/unit/test_harness_config.py tests/unit/test_core_types.py tests/unit/test_core_protocols.py -v Expected: 全部 PASS

  • Step 4: 提交
git add tests/integration/test_pool_strategy.py

Commit message: test(integration): add PerCategoryPoolStrategy end-to-end test


Task 11: lint + 最终验证

Files: 无新增

  • Step 1: 代码格式化

Run: conda activate Video-Tree-TRM && ruff format app/ core/ tests/ Run: conda activate Video-Tree-TRM && ruff check app/ core/ tests/ --fix

  • Step 2: 全量测试

Run: conda activate Video-Tree-TRM && pytest tests/ -v --tb=short Expected: 全部 PASS

  • Step 3: 修复任何问题后提交
git add -A

Commit message: chore: lint and format per-category pool strategy implementation


保真校验

本计划不涉及核心算法迁移(13 项均不涉及)。PoolStrategy 是新增抽象,GlobalPoolStrategy 封装的 build_pools 保持原有逻辑不变。保真校验不适用。