Files
Video-Tree-TRM5/tests/unit/test_pools_video_atomic.py

195 lines
7.2 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.
"""视频原子切分 split_by_video_assignment 单元测试。
覆盖:同一视频绝不跨池、缺失归属 fail-fast、trainval 内部视频组整组同落
train/valcorrectness 分层)、baseline_val_accuracy 计算正确。
"""
from __future__ import annotations
import pytest
from app.harness.pools import split_by_video_assignment
from core.types import GeneratedQuestion
def _q(qid: str, vid: str, tt: str = "Counting Problem") -> GeneratedQuestion:
"""构造最小可用题目(补齐 GeneratedQuestion 的必填 source_nodes/difficulty)。"""
return GeneratedQuestion(
question_id=qid,
video_id=vid,
task_type=tt,
question="",
options=("A", "B", "C", "D"),
answer="A",
source_nodes=(),
difficulty="medium",
)
def test_video_never_split_across_pools():
"""同一视频的所有题绝不跨 trainval/test 池。"""
qs = [_q("v1-1", "v1"), _q("v1-2", "v1"), _q("v1-3", "v1"), _q("v2-1", "v2")]
assignment = {"v1": "trainval", "v2": "test"}
pools = split_by_video_assignment(
qs,
assignment,
correctness={q.question_id: True for q in qs},
val_ratio=0.0,
seed=0,
)
test_vids = {q.video_id for q in pools.test}
train_vids = {q.video_id for q in pools.diagnosis + pools.validation}
assert test_vids & train_vids == set() # 视频不跨池
assert test_vids == {"v2"} and train_vids == {"v1"}
def test_missing_assignment_fails_fast():
"""assignment 缺失某 video_id 时 fail-fast,不静默丢题。"""
qs = [_q("v1-1", "v1"), _q("v2-1", "v2")]
assignment = {"v1": "trainval"} # 缺 v2
with pytest.raises(ValueError, match="assignment"):
split_by_video_assignment(
qs,
assignment,
correctness={q.question_id: True for q in qs},
val_ratio=0.0,
seed=0,
)
def test_invalid_assignment_value_fails_fast():
"""assignment 取值非法(非 trainval/test)时 fail-fast。"""
qs = [_q("v1-1", "v1")]
assignment = {"v1": "holdout"}
with pytest.raises(ValueError, match="归属"):
split_by_video_assignment(
qs,
assignment,
correctness={q.question_id: True for q in qs},
val_ratio=0.0,
seed=0,
)
def test_val_ratio_out_of_range_fails_fast():
"""val_ratio 越界(<0 或 >1)时 fail-fast。"""
qs = [_q("v1-1", "v1")]
assignment = {"v1": "trainval"}
correctness = {"v1-1": True}
for bad in (-0.1, 1.5):
with pytest.raises(ValueError, match="val_ratio"):
split_by_video_assignment(
qs, assignment, correctness=correctness, val_ratio=bad, seed=0
)
def test_missing_correctness_on_test_side_fails_fast():
"""test 侧某题缺 correctness 时 fail-fast,不静默兜底 False 污染评估口径。"""
qs = [_q("v1-1", "v1"), _q("v2-1", "v2")]
assignment = {"v1": "trainval", "v2": "test"}
correctness = {"v1-1": True} # 缺 test 侧 v2-1
with pytest.raises(ValueError, match="correctness"):
split_by_video_assignment(qs, assignment, correctness=correctness, val_ratio=0.0, seed=0)
def test_video_group_atomic_in_trainval_split():
"""trainval 内 val 切分以视频组为原子:同 video 的题整组同落 train 或 val。"""
qs = [
_q("v1-1", "v1"),
_q("v1-2", "v1"),
_q("v2-1", "v2"),
_q("v2-2", "v2"),
_q("v3-1", "v3"),
_q("v3-2", "v3"),
_q("v4-1", "v4"),
]
assignment = {"v1": "trainval", "v2": "trainval", "v3": "trainval", "v4": "trainval"}
pools = split_by_video_assignment(
qs,
assignment,
correctness={q.question_id: True for q in qs},
val_ratio=0.5,
seed=0,
)
diag_vids = {q.video_id for q in pools.diagnosis}
val_vids = {q.video_id for q in pools.validation}
# 视频组不跨 train/val
assert diag_vids & val_vids == set()
# 无题丢失
assert len(pools.diagnosis) + len(pools.validation) == len(qs)
# 同一 video 的所有题落在同侧
for vid in {q.video_id for q in qs}:
vid_pools = {"diag" if q in pools.diagnosis else "val" for q in qs if q.video_id == vid}
assert len(vid_pools) <= 1
def test_baseline_val_accuracy_reflects_validation():
"""baseline_val_accuracy = validation 池正确率。"""
qs = [_q("v1-1", "v1"), _q("v2-1", "v2")]
assignment = {"v1": "trainval", "v2": "trainval"}
correctness = {"v1-1": True, "v2-1": False}
pools = split_by_video_assignment(
qs, assignment, correctness=correctness, val_ratio=1.0, seed=0
)
# val_ratio=1.0 → 全部进 validation
assert len(pools.validation) == 2
assert pools.baseline_val_accuracy == pytest.approx(0.5)
assert pools.diagnosis == []
def test_tier_aware_keeps_high_t2_in_diag():
"""错题视频组按 T2 含量升序进 val:T2 高的组保留在 diagnosis。"""
from app.harness.pools import split_by_video_assignment
from core.types import GeneratedQuestion
def _q(qid, vid):
return GeneratedQuestion(
question_id=qid, video_id=vid, task_type="X", question="q",
options=["A", "B"], answer="A", source_nodes=[], difficulty="easy",
)
# 4 个错题视频(每视频 1 题),T2 数分别 2/1/0/0
questions = [_q(f"{v}-1", v) for v in ("vA", "vB", "vC", "vD")]
assignment = {v: "trainval" for v in ("vA", "vB", "vC", "vD")}
correctness = {f"{v}-1": False for v in ("vA", "vB", "vC", "vD")}
wrong_tier = {"vA": 2, "vB": 1, "vC": 0, "vD": 0}
pools = split_by_video_assignment(
questions, assignment, correctness, val_ratio=0.5, seed=7,
wrong_tier_by_video=wrong_tier,
)
diag_vids = {q.video_id for q in pools.diagnosis}
# T2 最高的 vA 必留 diagT2=0 的组优先进 val
assert "vA" in diag_vids
assert "vB" in diag_vids
def test_val_wrong_min_repair_pulls_from_diag():
"""val 错题不足 val_wrong_min 时从 diag 换入低 T2 错题组补足。"""
from app.harness.pools import split_by_video_assignment
from core.types import GeneratedQuestion
def _q(qid, vid, correct):
return GeneratedQuestion(
question_id=qid, video_id=vid, task_type="X", question="q",
options=["A", "B"], answer="A", source_nodes=[], difficulty="easy",
)
# 8 错题视频 + 2 正确视频;val_ratio 小使初分 val 错题不足,触发修复
vids_wrong = [f"w{i}" for i in range(8)]
vids_correct = ["c0", "c1"]
questions = [_q(f"{v}-1", v, False) for v in vids_wrong] + [
_q(f"{v}-1", v, True) for v in vids_correct
]
assignment = {v: "trainval" for v in vids_wrong + vids_correct}
correctness = {f"{v}-1": False for v in vids_wrong}
correctness.update({f"{v}-1": True for v in vids_correct})
wrong_tier = {v: i for i, v in enumerate(vids_wrong)} # 递增 T2
pools = split_by_video_assignment(
questions, assignment, correctness, val_ratio=0.1, seed=7,
wrong_tier_by_video=wrong_tier, val_wrong_min=4,
)
val_wrong = sum(1 for q in pools.validation if not correctness[q.question_id])
assert val_wrong >= 4, f"功效修复后 val 错题 {val_wrong} < 4"