50 KiB
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(封装现有逻辑)和 PerCategoryPoolStrategy(per-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
注意:PoolStrategy 的 build 方法返回 Pools,但 Pools 在 app/harness/pools.py 中定义。因为 app/ports.py 和 app/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:
"""某类全部 correct(0 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)。
返回:
Pools(diagnosis=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 math 和 import random(random 已有,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: FAIL(save_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.py 的 RunConfig 有默认值区(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_mode 和 train_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.py 的 base_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.yaml 的 harness: 段 early_stop_patience: 8 和 use_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 -> tuple(RunConfig 要求 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 保持原有逻辑不变。保真校验不适用。