37d4519905
Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
153 lines
5.1 KiB
Python
153 lines
5.1 KiB
Python
"""PerCategoryPoolStrategy 端到端集成测试。
|
|
|
|
验证从构造题目 → 伪造 baseline → 池构建 → 冻结 → 加载的完整流程。
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import json
|
|
from collections import Counter
|
|
from typing import TYPE_CHECKING
|
|
|
|
if TYPE_CHECKING:
|
|
from pathlib import Path
|
|
|
|
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
|
|
|
|
def test_single_category_flow(self, tmp_path: Path) -> None:
|
|
"""单类别训练流程:30 题 → 20 train + 10 val。"""
|
|
questions = [_make_question(f"q_{i:03d}", "Object Recognition") for i in range(30)]
|
|
correctness = {q.question_id: (i < 20) for i, q in enumerate(questions)}
|
|
|
|
config = PoolConfig(
|
|
task_types=("Object 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=20 / 30,
|
|
test_questions_dir=None,
|
|
)
|
|
|
|
strategy = PerCategoryPoolStrategy()
|
|
pools = strategy.build(questions, correctness, config)
|
|
assert len(pools.diagnosis) == 20
|
|
assert len(pools.validation) == 10
|
|
|
|
# 冻结 → 加载
|
|
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) == 20
|
|
assert len(loaded.validation) == 10
|