"""视频原子切分 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_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 == []