195 lines
7.2 KiB
Python
195 lines
7.2 KiB
Python
"""视频原子切分 split_by_video_assignment 单元测试。
|
||
|
||
覆盖:同一视频绝不跨池、缺失归属 fail-fast、trainval 内部视频组整组同落
|
||
train/val(correctness 分层)、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 必留 diag;T2=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"
|