docs: add maintenance pool implementation plan
Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
@@ -175,6 +175,11 @@
|
||||
"id": "design:maintenance-pool",
|
||||
"label": "训练池 Maintenance 正确题自动补入机制",
|
||||
"type": "design"
|
||||
},
|
||||
{
|
||||
"id": "plan:maintenance-pool",
|
||||
"label": "Maintenance Pool 自动补入实现计划",
|
||||
"type": "plan"
|
||||
}
|
||||
],
|
||||
"links": [
|
||||
@@ -324,6 +329,13 @@
|
||||
"relation": "extends",
|
||||
"evidence": "maintenance pool supplements the per-category pool strategy to support SubPattern-targeted training",
|
||||
"added": "2026-07-14T14:11:13.527027+00:00"
|
||||
},
|
||||
{
|
||||
"source": "plan:maintenance-pool",
|
||||
"target": "design:maintenance-pool",
|
||||
"relation": "implements",
|
||||
"evidence": "Implements maintenance pool auto-supplement in PerCategoryPoolStrategy",
|
||||
"added": "2026-07-14T14:22:44.443012+00:00"
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -1,6 +1,6 @@
|
||||
# Research Wiki 索引
|
||||
|
||||
> 自动生成,更新时间:2026-07-14 14:11 UTC
|
||||
> 自动生成,更新时间:2026-07-14 14:22 UTC
|
||||
|
||||
## design (27)
|
||||
- [2026-07-06-core-agent-adapters-llm-design](designs/2026-07-06-core-agent-adapters-llm-design.md) `design:2026-07-06-core-agent-adapters-llm-design`
|
||||
@@ -38,7 +38,7 @@
|
||||
- [Harness 评估: Spec-1 修复验证 (infer_spec1check)](findings/eval-spec1check.md) `finding:eval-spec1check`
|
||||
- [Harness 评估: Spec-2 批量并行建树](findings/eval-spec2-batch-tree-build.md) `finding:eval-spec2-batch-tree-build`
|
||||
|
||||
## plan (30)
|
||||
## plan (32)
|
||||
- [2026-07-06-core-agent-adapters-llm](plans/2026-07-06-core-agent-adapters-llm.md) `plan:2026-07-06-core-agent-adapters-llm`
|
||||
- [2026-07-07-app-harness](plans/2026-07-07-app-harness.md) `plan:2026-07-07-app-harness`
|
||||
- [2026-07-07-core-evolution](plans/2026-07-07-core-evolution.md) `plan:2026-07-07-core-evolution`
|
||||
@@ -52,6 +52,7 @@
|
||||
- [2026-07-12-per-category-pool-strategy](plans/2026-07-12-per-category-pool-strategy.md) `plan:2026-07-12-per-category-pool-strategy`
|
||||
- [2026-07-14-action-recognition-strategy](plans/2026-07-14-action-recognition-strategy.md) `plan:2026-07-14-action-recognition-strategy`
|
||||
- [2026-07-14-action-recognition-training](plans/2026-07-14-action-recognition-training.md) `plan:2026-07-14-action-recognition-training`
|
||||
- [2026-07-14-maintenance-pool](plans/2026-07-14-maintenance-pool.md) `plan:2026-07-14-maintenance-pool`
|
||||
- [2026-07-14-task-type-strategy-framework](plans/2026-07-14-task-type-strategy-framework.md) `plan:2026-07-14-task-type-strategy-framework`
|
||||
- [Action Recognition 单题型首次训练实验计划](plans/action-recognition-training.md) `plan:action-recognition-training`
|
||||
- [ActionRecognitionStrategy 特化实现计划 (Plan B)](plans/action-recognition-strategy.md) `plan:action-recognition-strategy`
|
||||
@@ -59,6 +60,7 @@
|
||||
- [app/search/ 搜索 Agent 装配层实现计划](plans/2026-07-07-search-module.md) `plan:2026-07-07-search-module`
|
||||
- [core/agent/ + adapters/llm 基础设施实现计划](plans/core-agent-adapters-llm.md) `plan:core-agent-adapters-llm`
|
||||
- [main.py 推理入口 + 初始 Prompt 集实现计划](plans/main-inference-entry.md) `plan:main-inference-entry`
|
||||
- [Maintenance Pool 自动补入实现计划](plans/maintenance-pool.md) `plan:maintenance-pool`
|
||||
- [Per-Category Pool Strategy 实现计划](plans/per-category-pool-strategy.md) `plan:per-category-pool-strategy`
|
||||
- [question_gen 模块实现计划](plans/question-gen.md) `plan:question-gen`
|
||||
- [Spec-1 Agent 执行环境修复实现计划](plans/agent-runtime-fixes-plan.md) `plan:agent-runtime-fixes-plan`
|
||||
|
||||
@@ -82,3 +82,6 @@
|
||||
- [2026-07-14 14:11 UTC] 新增 design: 训练池 Maintenance 正确题自动补入机制 (design:maintenance-pool)
|
||||
- [2026-07-14 14:11 UTC] 新增边: design:maintenance-pool --extends--> design:task-type-strategy
|
||||
- [2026-07-14 14:11 UTC] 重建索引: 68 篇页面
|
||||
- [2026-07-14 14:22 UTC] 新增 plan: Maintenance Pool 自动补入实现计划 (plan:maintenance-pool)
|
||||
- [2026-07-14 14:22 UTC] 新增边: plan:maintenance-pool --implements--> design:maintenance-pool
|
||||
- [2026-07-14 14:22 UTC] 重建索引: 70 篇页面
|
||||
|
||||
@@ -0,0 +1,589 @@
|
||||
# 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` 之后追加:
|
||||
|
||||
```python
|
||||
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: 运行测试验证失败**
|
||||
|
||||
```bash
|
||||
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.py` 的 `PoolConfig` dataclass 中,在 `test_questions_dir` 之后追加:
|
||||
|
||||
```python
|
||||
batch_correct_ratio: float | None = None
|
||||
```
|
||||
|
||||
- [ ] **Step 4: 在 _to_pool_config 中映射**
|
||||
|
||||
在 `app/harness/pools.py` 的 `_to_pool_config` 函数中,`return PoolConfig(...)` 调用里追加:
|
||||
|
||||
```python
|
||||
batch_correct_ratio=config.batch_correct_ratio,
|
||||
```
|
||||
|
||||
- [ ] **Step 5: 运行测试验证通过**
|
||||
|
||||
```bash
|
||||
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: 提交**
|
||||
|
||||
```bash
|
||||
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` 末尾追加:
|
||||
|
||||
```python
|
||||
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: 运行测试验证失败**
|
||||
|
||||
```bash
|
||||
conda run -n Video-Tree-TRM pytest tests/integration/test_pool_strategy.py::TestMaintenanceSupplementation -v
|
||||
```
|
||||
|
||||
预期:`test_supplements_when_correct_ratio_too_low` FAIL(build 不接受 db_path)。
|
||||
|
||||
- [ ] **Step 3: 修改 PoolStrategy Protocol 签名**
|
||||
|
||||
在 `app/ports.py` 中,修改 `PoolStrategy.build` 的签名,新增 `db_path` 关键字参数:
|
||||
|
||||
```python
|
||||
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.py` 的 `GlobalPoolStrategy.build` 中,新增 `db_path` 参数但忽略:
|
||||
|
||||
```python
|
||||
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.py` 的 `PerCategoryPoolStrategy.build` 中:
|
||||
|
||||
**5a.** 修改签名新增 `db_path`:
|
||||
|
||||
```python
|
||||
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:
|
||||
|
||||
```python
|
||||
# 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` 方法:
|
||||
|
||||
```python
|
||||
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.py` 的 `build_or_load_pools` 函数中(约 536 行),将:
|
||||
|
||||
```python
|
||||
pools = strategy.build(questions, correctness, pool_config)
|
||||
```
|
||||
|
||||
改为:
|
||||
|
||||
```python
|
||||
pools = strategy.build(questions, correctness, pool_config, db_path=db_path)
|
||||
```
|
||||
|
||||
- [ ] **Step 7: 运行测试验证通过**
|
||||
|
||||
```bash
|
||||
conda run -n Video-Tree-TRM pytest tests/integration/test_pool_strategy.py -v --tb=short
|
||||
```
|
||||
|
||||
预期:全部 PASS。
|
||||
|
||||
- [ ] **Step 8: 提交**
|
||||
|
||||
```bash
|
||||
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**
|
||||
|
||||
```bash
|
||||
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: 全量测试**
|
||||
|
||||
```bash
|
||||
conda run -n Video-Tree-TRM pytest tests/unit/ tests/integration/ -q --tb=short
|
||||
```
|
||||
|
||||
预期:1208+ 全部 PASS
|
||||
|
||||
- [ ] **Step 3: 提交(如有 lint 修复)**
|
||||
|
||||
```bash
|
||||
git add -A && git commit -m "chore: lint and format maintenance pool changes"
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 核心算法保真校验
|
||||
|
||||
本计划不涉及核心算法迁移,保真校验不适用。
|
||||
|
||||
修改仅限于 `PerCategoryPoolStrategy.build()` 内部新增 Phase 2.5(补入逻辑),不改变 `_split_one_category`、`build_pools`、`GlobalPoolStrategy` 或 batching 算法的任何行为。
|
||||
Reference in New Issue
Block a user