refactor: add video-atomic pool split (algo #5 gate input preserved)

This commit is contained in:
2026-07-15 12:14:36 -04:00
parent 9d19328cc9
commit 20eea98cdd
2 changed files with 318 additions and 16 deletions
+116
View File
@@ -0,0 +1,116 @@
"""视频原子切分 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_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 == []