Files
Video-Tree-TRM5/research-wiki/plans/2026-07-14-maintenance-pool.md
iomgaa 453cf62088 docs: add maintenance pool implementation plan
Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
2026-07-14 10:22:52 -04:00

20 KiB
Raw Permalink Blame History

Maintenance Pool 自动补入实现计划

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

Goal:PerCategoryPoolStrategy.build() 的池划分前,自动从 benchmark 补入正确题以满足 batch_correct_ratio,并在正确率过高时发出警告。

Architecture:PoolConfig 新增 batch_correct_ratio 字段,PerCategoryPoolStrategy.build() 新增 db_path 参数。在 Phase 2(分组)和 Phase 3(划分)之间插入 Phase 2.5:计算每个 task_type 的正确率缺口 → 从 test_questions_dir 加载 benchmark 同类题 → 查 DB 历史记录筛选正确题 → 补入并标记 family="VME_MAINTENANCE" → 正确率过高时 warning。

Tech Stack: Python 3.11, pytest, sqlite3

关联设计: research-wiki/designs/2026-07-14-maintenance-pool-design.md


Task 1: PoolConfig 新增 batch_correct_ratio 字段

Files:

  • Modify: core/types.py:66-96

  • Modify: app/harness/pools.py:332-364

  • Test: tests/integration/test_pool_strategy.py

  • Step 1: 写失败测试 — PoolConfig 接受 batch_correct_ratio

tests/integration/test_pool_strategy.py 文件顶部的 _make_question 之后追加:

class TestPoolConfigBatchRatio:
    """PoolConfig 新增 batch_correct_ratio 字段。"""

    def test_default_none(self):
        """不传 batch_correct_ratio 时默认 None。"""
        config = PoolConfig(
            task_types=None,
            seed=42,
            baseline_run_id="bl",
            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=0.667,
            test_questions_dir=None,
        )
        assert config.batch_correct_ratio is None

    def test_explicit_value(self):
        """显式传入 batch_correct_ratio。"""
        config = PoolConfig(
            task_types=None,
            seed=42,
            baseline_run_id="bl",
            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=0.667,
            test_questions_dir=None,
            batch_correct_ratio=0.5,
        )
        assert config.batch_correct_ratio == 0.5
  • Step 2: 运行测试验证失败
conda run -n Video-Tree-TRM pytest tests/integration/test_pool_strategy.py::TestPoolConfigBatchRatio -v

预期:TypeError — PoolConfig 不接受 batch_correct_ratio。

  • Step 3: 在 PoolConfig 新增字段

core/types.pyPoolConfig dataclass 中,在 test_questions_dir 之后追加:

    batch_correct_ratio: float | None = None
  • Step 4: 在 _to_pool_config 中映射

app/harness/pools.py_to_pool_config 函数中,return PoolConfig(...) 调用里追加:

        batch_correct_ratio=config.batch_correct_ratio,
  • Step 5: 运行测试验证通过
conda run -n Video-Tree-TRM pytest tests/integration/test_pool_strategy.py::TestPoolConfigBatchRatio tests/integration/test_pool_strategy.py::TestPerCategoryE2E -v

预期:全部 PASS(新测试通过,旧测试不受影响因为字段有默认值 None)。

  • Step 6: 提交
git add core/types.py app/harness/pools.py tests/integration/test_pool_strategy.py
git commit -m "feat(pools): add batch_correct_ratio field to PoolConfig"

Task 2: PerCategoryPoolStrategy.build 新增 db_path 参数 + Phase 2.5 maintenance 补入

Files:

  • Modify: app/ports.py:158-163

  • Modify: app/harness/pools.py:546-624

  • Modify: app/harness/pools.py:526-536 (build_or_load_pools 调用点)

  • Test: tests/integration/test_pool_strategy.py

  • Step 1: 写失败测试 — maintenance 补入核心逻辑

tests/integration/test_pool_strategy.py 末尾追加:

class TestMaintenanceSupplementation:
    """Phase 2.5: maintenance 正确题自动补入。"""

    def _make_benchmark_dir(self, tmp_path: Path, task_type: str, n: int) -> Path:
        """创建 benchmark 题目目录(模拟 VME)。"""
        import json

        bench_dir = tmp_path / "benchmark"
        bench_dir.mkdir()
        questions = []
        for i in range(n):
            questions.append({
                "question_id": f"vme_{task_type}_{i:03d}",
                "video_id": f"vid_{i:03d}",
                "task_type": task_type,
                "question": f"VME Q{i}?",
                "options": ["A. a", "B. b", "C. c", "D. d"],
                "answer": "A",
            })
        (bench_dir / "benchmark.json").write_text(
            json.dumps(questions, ensure_ascii=False),
            encoding="utf-8",
        )
        return bench_dir

    def _make_db_with_correctness(
        self, tmp_path: Path, qids_correct: list[str], qids_wrong: list[str],
    ) -> "Path":
        """创建带历史推理记录的 harness.db。"""
        import sqlite3

        db_path = tmp_path / "harness.db"
        conn = sqlite3.connect(str(db_path))
        conn.execute(
            "CREATE TABLE IF NOT EXISTS predictions ("
            "run_id TEXT, timestamp TEXT, video_id TEXT, question_id TEXT, "
            "task_type TEXT, prediction TEXT, answer TEXT, evidence TEXT, "
            "reasoning TEXT, steps_used INTEGER, prompt_tokens INTEGER, "
            "completion_tokens INTEGER, stop_reason TEXT, steps_json JSON)"
        )
        for qid in qids_correct:
            conn.execute(
                "INSERT INTO predictions (run_id, question_id, prediction, answer) "
                "VALUES (?, ?, ?, ?)",
                ("infer_adhoc", qid, "A", "A"),
            )
        for qid in qids_wrong:
            conn.execute(
                "INSERT INTO predictions (run_id, question_id, prediction, answer) "
                "VALUES (?, ?, ?, ?)",
                ("infer_adhoc", qid, "B", "A"),
            )
        conn.commit()
        conn.close()
        return db_path

    def test_supplements_when_correct_ratio_too_low(self, tmp_path: Path) -> None:
        """正确率低于 batch_correct_ratio 时自动补入。"""
        # 30 道训练题:5 correct, 25 wrong
        questions = [_make_question(f"q_{i:03d}", "Action Recognition") for i in range(30)]
        correctness = {q.question_id: (i < 5) for i, q in enumerate(questions)}

        # 40 道 benchmark 题(模拟 VME),其中 30 道在 DB 中标记为正确
        bench_dir = self._make_benchmark_dir(tmp_path, "Action Recognition", 40)
        correct_vme_ids = [f"vme_Action Recognition_{i:03d}" for i in range(30)]
        wrong_vme_ids = [f"vme_Action Recognition_{i:03d}" for i in range(30, 40)]
        db_path = self._make_db_with_correctness(tmp_path, correct_vme_ids, wrong_vme_ids)

        config = PoolConfig(
            task_types=("Action Recognition",),
            seed=42,
            baseline_run_id="bl",
            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=0.667,
            test_questions_dir=bench_dir,
            batch_correct_ratio=0.5,
        )

        strategy = PerCategoryPoolStrategy()
        pools = strategy.build(questions, correctness, config, db_path=db_path)

        # 需补入 k = (0.5*25 - 0.5*5) / 0.5 = 20 道
        total = len(pools.diagnosis) + len(pools.validation)
        assert total == 50  # 30 原始 + 20 补入

        # 验证补入题标记
        all_qs = pools.diagnosis + pools.validation
        maintenance_qs = [q for q in all_qs if q.family == "VME_MAINTENANCE"]
        assert len(maintenance_qs) == 20

        # 验证 correctness 中补入题为 True
        for q in maintenance_qs:
            assert pools.correctness[q.question_id] is True

    def test_no_supplement_when_ratio_satisfied(self, tmp_path: Path) -> None:
        """正确率已满足时不补入。"""
        questions = [_make_question(f"q_{i:03d}", "Action Recognition") for i in range(30)]
        correctness = {q.question_id: (i < 20) for i, q in enumerate(questions)}

        config = PoolConfig(
            task_types=("Action Recognition",),
            seed=42,
            baseline_run_id="bl",
            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=0.667,
            test_questions_dir=None,
            batch_correct_ratio=0.5,
        )

        strategy = PerCategoryPoolStrategy()
        pools = strategy.build(questions, correctness, config)

        total = len(pools.diagnosis) + len(pools.validation)
        assert total == 30  # 无补入

    def test_no_supplement_when_ratio_not_configured(self, tmp_path: Path) -> None:
        """batch_correct_ratio 未配置时不补入。"""
        questions = [_make_question(f"q_{i:03d}", "Action Recognition") for i in range(30)]
        correctness = {q.question_id: (i < 2) for i, q in enumerate(questions)}

        config = PoolConfig(
            task_types=("Action Recognition",),
            seed=42,
            baseline_run_id="bl",
            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=0.667,
            test_questions_dir=None,
            batch_correct_ratio=None,
        )

        strategy = PerCategoryPoolStrategy()
        pools = strategy.build(questions, correctness, config)

        total = len(pools.diagnosis) + len(pools.validation)
        assert total == 30

    def test_caps_at_available_candidates(self, tmp_path: Path) -> None:
        """候选不足时补入全部可用,接受不完美比例。"""
        # 30 题全错
        questions = [_make_question(f"q_{i:03d}", "Action Recognition") for i in range(30)]
        correctness = {q.question_id: False for q in questions}

        # 只有 10 道 benchmark 正确题
        bench_dir = self._make_benchmark_dir(tmp_path, "Action Recognition", 10)
        correct_ids = [f"vme_Action Recognition_{i:03d}" for i in range(10)]
        db_path = self._make_db_with_correctness(tmp_path, correct_ids, [])

        config = PoolConfig(
            task_types=("Action Recognition",),
            seed=42,
            baseline_run_id="bl",
            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=0.667,
            test_questions_dir=bench_dir,
            batch_correct_ratio=0.5,
        )

        strategy = PerCategoryPoolStrategy()
        pools = strategy.build(questions, correctness, config, db_path=db_path)

        total = len(pools.diagnosis) + len(pools.validation)
        assert total == 40  # 30 + 10(全部可用)

    def test_high_correct_ratio_warning(self, tmp_path: Path) -> None:
        """正确率过高时发出 warning。"""
        from loguru import logger

        questions = [_make_question(f"q_{i:03d}", "Action Recognition") for i in range(30)]
        correctness = {q.question_id: (i < 28) for i, q in enumerate(questions)}

        config = PoolConfig(
            task_types=("Action Recognition",),
            seed=42,
            baseline_run_id="bl",
            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=0.667,
            test_questions_dir=None,
            batch_correct_ratio=0.5,
        )

        captured: list[str] = []
        sink_id = logger.add(lambda msg: captured.append(str(msg)), level="WARNING")
        try:
            strategy = PerCategoryPoolStrategy()
            strategy.build(questions, correctness, config)
        finally:
            logger.remove(sink_id)

        assert any("出题可能太简单" in m for m in captured)
  • Step 2: 运行测试验证失败
conda run -n Video-Tree-TRM pytest tests/integration/test_pool_strategy.py::TestMaintenanceSupplementation -v

预期:test_supplements_when_correct_ratio_too_low FAILbuild 不接受 db_path)。

  • Step 3: 修改 PoolStrategy Protocol 签名

app/ports.py 中,修改 PoolStrategy.build 的签名,新增 db_path 关键字参数:

    def build(
        self,
        questions: list[GeneratedQuestion],
        correctness: dict[str, bool],
        config: PoolConfig,
        *,
        db_path: _Path | None = None,
    ) -> Pools: ...

需在文件顶部 TYPE_CHECKING 块中添加 from pathlib import Path as _Path(若不存在)。

  • Step 4: 修改 GlobalPoolStrategy.build 签名兼容

app/harness/pools.pyGlobalPoolStrategy.build 中,新增 db_path 参数但忽略:

    def build(
        self,
        questions: list[GeneratedQuestion],
        correctness: dict[str, bool],
        config: PoolConfig,
        *,
        db_path: Path | None = None,
    ) -> Pools:

函数体不变。

  • Step 5: 实现 PerCategoryPoolStrategy.build 的 Phase 2.5

app/harness/pools.pyPerCategoryPoolStrategy.build 中:

5a. 修改签名新增 db_path

    def build(
        self,
        questions: list[GeneratedQuestion],
        correctness: dict[str, bool],
        config: PoolConfig,
        *,
        db_path: Path | None = None,
    ) -> Pools:

5b. 在 Phase 2(分组循环)和 Phase 3(分层划分循环)之间插入 Phase 2.5:

        # Phase 2.5: 正确率检查 + maintenance 补入
        if config.batch_correct_ratio is not None:
            self._check_and_supplement_maintenance(
                groups, correctness, config, db_path,
            )

5c.PerCategoryPoolStrategy 类中新增 _supplement_maintenance 方法:

    def _check_and_supplement_maintenance(
        self,
        groups: dict[str, list[GeneratedQuestion]],
        correctness: dict[str, bool],
        config: PoolConfig,
        db_path: Path | None,
    ) -> None:
        """按 task_type 检查正确率,过高警告,过低则从 benchmark 补入正确题。

        修改 groups 和 correctness(原地更新)。

        参数:
            groups: task_type → 题目列表映射(原地追加补入题)。
            correctness: question_id → 是否正确映射(原地追加补入题标记)。
            config: 含 batch_correct_ratio 和 test_questions_dir。
            db_path: harness.db 路径,用于查询 benchmark 历史推理记录。
        """
        import sqlite3

        r = config.batch_correct_ratio

        # Phase 2.5a: 正确率检查(不依赖 test_questions_dir
        for task_type, group in groups.items():
            c = sum(1 for q in group if correctness.get(q.question_id, False))
            n = len(group)
            ratio = c / n if n > 0 else 0.0

            if ratio > 1 - r:
                logger.warning(
                    "类别 {} 正确率 {:.1%} 过高(阈值 {:.1%}),出题可能太简单",
                    task_type,
                    ratio,
                    1 - r,
                )

        # Phase 2.5b: maintenance 补入(需要 test_questions_dir
        if config.test_questions_dir is None:
            return

        from app.question_gen import load_benchmark

        bench_questions = load_benchmark(config.test_questions_dir)

        # 查询 DB 中 benchmark 题的历史正确性
        bench_correctness: dict[str, bool] = {}
        if db_path is not None and db_path.exists():
            conn = sqlite3.connect(str(db_path))
            bench_qids = [q.question_id for q in bench_questions]
            if bench_qids:
                placeholders = ",".join("?" for _ in bench_qids)
                rows = conn.execute(
                    f"SELECT question_id, prediction, answer FROM predictions "
                    f"WHERE question_id IN ({placeholders}) "
                    f"ORDER BY timestamp DESC",
                    bench_qids,
                ).fetchall()
                for qid, pred, ans in rows:
                    if qid not in bench_correctness:
                        bench_correctness[qid] = (pred == ans)
                conn.close()

        # 按 task_type 索引 benchmark 题
        bench_by_type: dict[str, list[GeneratedQuestion]] = defaultdict(list)
        for q in bench_questions:
            bench_by_type[q.task_type].append(q)

        for task_type, group in groups.items():
            c = sum(1 for q in group if correctness.get(q.question_id, False))
            w = len(group) - c
            n = len(group)
            ratio = c / n if n > 0 else 0.0

            if ratio >= r:
                continue

            # 计算需补入数
            k = math.ceil((r * w - (1 - r) * c) / (1 - r))

            # 筛选候选:同 task_type + DB 历史正确 + 不在当前组中
            existing_ids = {q.question_id for q in group}
            candidates = [
                q for q in bench_by_type.get(task_type, [])
                if bench_correctness.get(q.question_id, False)
                and q.question_id not in existing_ids
            ]

            if not candidates:
                logger.warning(
                    "类别 {} 需补入 {} 道正确题,但 benchmark 中无可用候选",
                    task_type,
                    k,
                )
                continue

            actual = min(k, len(candidates))
            for q in candidates[:actual]:
                supplemented = GeneratedQuestion(
                    question_id=q.question_id,
                    video_id=q.video_id,
                    task_type=q.task_type,
                    question=q.question,
                    options=q.options,
                    answer=q.answer,
                    source_nodes=q.source_nodes,
                    difficulty=q.difficulty,
                    family="VME_MAINTENANCE",
                    skill_target=q.skill_target,
                    difficulty_steps=q.difficulty_steps,
                )
                group.append(supplemented)
                correctness[supplemented.question_id] = True

            logger.info(
                "类别 {} 正确率 {:.1%} < {:.1%},从 benchmark 补入 {} 道 maintenance 正确题",
                task_type,
                ratio,
                r,
                actual,
            )
  • Step 6: 修改 build_or_load_pools 透传 db_path

app/harness/pools.pybuild_or_load_pools 函数中(约 536 行),将:

    pools = strategy.build(questions, correctness, pool_config)

改为:

    pools = strategy.build(questions, correctness, pool_config, db_path=db_path)
  • Step 7: 运行测试验证通过
conda run -n Video-Tree-TRM pytest tests/integration/test_pool_strategy.py -v --tb=short

预期:全部 PASS。

  • Step 8: 提交
git add app/ports.py app/harness/pools.py tests/integration/test_pool_strategy.py
git commit -m "feat(pools): auto-supplement maintenance correct questions in PerCategoryPoolStrategy"

Task 3: 全量回归测试 + lint

  • Step 1: lint
conda run -n Video-Tree-TRM ruff format app/harness/pools.py app/ports.py core/types.py tests/integration/test_pool_strategy.py
conda run -n Video-Tree-TRM ruff check app/harness/pools.py app/ports.py core/types.py
  • Step 2: 全量测试
conda run -n Video-Tree-TRM pytest tests/unit/ tests/integration/ -q --tb=short

预期:1208+ 全部 PASS

  • Step 3: 提交(如有 lint 修复)
git add -A && git commit -m "chore: lint and format maintenance pool changes"

核心算法保真校验

本计划不涉及核心算法迁移,保真校验不适用。

修改仅限于 PerCategoryPoolStrategy.build() 内部新增 Phase 2.5(补入逻辑),不改变 _split_one_categorybuild_poolsGlobalPoolStrategy 或 batching 算法的任何行为。