Files
2026-07-14 10:41:46 -04:00

387 lines
13 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""PerCategoryPoolStrategy 端到端集成测试。
验证从构造题目 → 伪造 baseline → 池构建 → 冻结 → 加载的完整流程。
"""
from __future__ import annotations
import json
import sqlite3
from collections import Counter
from typing import TYPE_CHECKING
from loguru import logger
from app.harness.pools import (
PerCategoryPoolStrategy,
load_pools,
save_pools,
)
from core.types import GeneratedQuestion, PoolConfig
if TYPE_CHECKING:
from pathlib import Path
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 Recognition",
"Action Reasoning",
"Attribute Perception",
"Counting Problem",
"Information Synopsis",
"Object Recognition",
"Object Reasoning",
"OCR Problems",
"Spatial Perception",
"Spatial Reasoning",
"Temporal Perception",
"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
# ── helpers ──
def _make_benchmark_dir(tmp_path: Path, task_type: str, n: int) -> Path:
"""在 tmp_path 下创建含 n 道题的 benchmark JSON 目录。"""
bench_dir = tmp_path / "benchmark"
bench_dir.mkdir(exist_ok=True)
questions = []
for i in range(n):
questions.append(
{
"question_id": f"bench_{task_type}_{i:03d}",
"video_id": "v_bench",
"task_type": task_type,
"question": f"Benchmark Q{i}?",
"options": ["A. a", "B. b", "C. c", "D. d"],
"answer": "A",
"source_nodes": ["n1"],
"difficulty": "medium",
}
)
(bench_dir / "bench.json").write_text(
json.dumps(questions, ensure_ascii=False),
encoding="utf-8",
)
return bench_dir
def _make_db_with_correctness(
tmp_path: Path,
qids_correct: list[str],
qids_wrong: list[str],
) -> Path:
"""创建含 predictions 表的 SQLite DB,模拟 benchmark 历史推理记录。"""
db_path = tmp_path / "harness.db"
conn = sqlite3.connect(str(db_path))
conn.execute(
"CREATE TABLE predictions ("
" question_id TEXT, prediction TEXT, answer TEXT, timestamp TEXT"
")"
)
for qid in qids_correct:
conn.execute(
"INSERT INTO predictions VALUES (?, ?, ?, ?)",
(qid, "A", "A", "2026-01-01T00:00:00"),
)
for qid in qids_wrong:
conn.execute(
"INSERT INTO predictions VALUES (?, ?, ?, ?)",
(qid, "B", "A", "2026-01-01T00:00:00"),
)
conn.commit()
conn.close()
return db_path
class TestMaintenanceSupplementation:
"""maintenance 自动补入集成测试。"""
def test_supplements_when_correct_ratio_too_low(self, tmp_path: Path) -> None:
"""正确率过低时从 benchmark 补入正确题。
30 题(5 correct, 25 wrong),benchmark 40 题(30 correct),
batch_correct_ratio=0.5 → 需补入 k=ceil((0.5*25 - 0.5*5)/0.5)=20
总量变为 50。
"""
tt = "Object Recognition"
questions = [_make_question(f"q_{i:03d}", tt) for i in range(30)]
correctness: dict[str, bool] = {}
for i, q in enumerate(questions):
correctness[q.question_id] = i < 5 # 前 5 正确
bench_dir = _make_benchmark_dir(tmp_path, tt, 40)
# benchmark 前 30 题在 DB 中为正确
bench_qids_correct = [f"bench_{tt}_{i:03d}" for i in range(30)]
bench_qids_wrong = [f"bench_{tt}_{i:03d}" for i in range(30, 40)]
db_path = _make_db_with_correctness(tmp_path, bench_qids_correct, bench_qids_wrong)
config = PoolConfig(
task_types=(tt,),
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.7,
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 == 50, f"补入后总量应为 50,实际 {total}"
# 验证补入题的 family 标记
all_qs = pools.diagnosis + pools.validation
maint = [q for q in all_qs if q.family == "VME_MAINTENANCE"]
assert len(maint) == 20
def test_no_supplement_when_ratio_satisfied(self, tmp_path: Path) -> None:
"""正确率达标时无需补入。
30 题(20 correct),batch_correct_ratio=0.5,无 test_questions_dir。
"""
tt = "Object Recognition"
questions = [_make_question(f"q_{i:03d}", tt) for i in range(30)]
correctness = {q.question_id: (i < 20) for i, q in enumerate(questions)}
config = PoolConfig(
task_types=(tt,),
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.7,
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, f"无补入时总量应为 30,实际 {total}"
def test_no_supplement_when_ratio_not_configured(self) -> None:
"""batch_correct_ratio=None 时不触发补入逻辑。"""
tt = "Object Recognition"
questions = [_make_question(f"q_{i:03d}", tt) for i in range(30)]
correctness = {q.question_id: (i < 5) for i, q in enumerate(questions)}
config = PoolConfig(
task_types=(tt,),
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.7,
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:
"""benchmark 候选不足时以实际可用数量为上限。
30 题全部 wrongbenchmark 仅 10 题 correct → 补入 10 题,总量 40。
"""
tt = "Object Recognition"
questions = [_make_question(f"q_{i:03d}", tt) for i in range(30)]
correctness = {q.question_id: False for q in questions}
bench_dir = _make_benchmark_dir(tmp_path, tt, 15)
bench_qids_correct = [f"bench_{tt}_{i:03d}" for i in range(10)]
bench_qids_wrong = [f"bench_{tt}_{i:03d}" for i in range(10, 15)]
db_path = _make_db_with_correctness(tmp_path, bench_qids_correct, bench_qids_wrong)
config = PoolConfig(
task_types=(tt,),
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.7,
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, f"候选不足时总量应为 40,实际 {total}"
def test_high_correct_ratio_warning(self) -> None:
"""正确率过高时发出警告。
30 题(28 correct),batch_correct_ratio=0.5test_questions_dir=None。
阈值 1-0.5=0.5,实际 28/30=93.3% > 50% → 触发警告。
"""
tt = "Object Recognition"
questions = [_make_question(f"q_{i:03d}", tt) for i in range(30)]
correctness = {q.question_id: (i < 28) for i, q in enumerate(questions)}
config = PoolConfig(
task_types=(tt,),
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.7,
test_questions_dir=None,
batch_correct_ratio=0.5,
)
strategy = PerCategoryPoolStrategy()
captured: list[str] = []
sink_id = logger.add(lambda msg: captured.append(str(msg)), level="WARNING")
try:
strategy.build(questions, correctness, config)
finally:
logger.remove(sink_id)
assert any("出题可能太简单" in m for m in captured), (
f"未捕获到正确率过高警告,captured: {captured}"
)