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

138 lines
4.8 KiB
Python
Raw 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 == []