"""pair 原子性回归测试:三池切分不得把孪生对劈到不同池。 覆盖 pools.py 的两条切分路径: - GlobalPoolStrategy 路径(build_pools 的 test->val->diag progressive exclusion); - PerCategoryPoolStrategy 路径(_split_one_category 的 train/val 分层划分)。 核心不变量:同一 pair_id 的两题必落在同一个池(或同一 split),绝不被拆散。 """ from __future__ import annotations from app.harness.pools import PerCategoryPoolStrategy, build_pools from core.types import GeneratedQuestion, PoolConfig def _pair(pid: str) -> list[GeneratedQuestion]: """构造一个合法孪生对(original + mirror),共享 pair_id / unit_id / flip_axis。 参数: pid: 该孪生对的共享标识。 返回: 含 original 与 mirror 两条题目的列表。 """ base = { "video_id": "v", "task_type": "Action Reasoning", "question": "?", "options": ("A. a", "B. b", "C. c", "D. d"), "answer": "A", "source_nodes": ("n",), "difficulty": "hard", "pair_id": pid, "flip_axis": "before_after", } return [ GeneratedQuestion(question_id=f"{pid}_o", question_role="pair_original", **base), GeneratedQuestion(question_id=f"{pid}_m", question_role="pair_mirror", **base), ] def test_pair_never_split_across_pools() -> None: """build_pools(Global 路径)三池切分后,任一 pair_id 只出现在一个池。""" qs = [q for pid in [f"p{i:02d}" for i in range(12)] for q in _pair(pid)] cfg = { "size": 4, "correct_ratio": None, "task_types": None, "seed": 7, "min_per_class": None, } pools = build_pools( qs, correctness={}, diag_cfg=cfg, val_cfg=cfg, test_cfg={"size": 4, "seed": 7}, baseline_run_id="b", ) by_name = { "diagnosis": pools.diagnosis, "validation": pools.validation, "test": pools.test, } loc: dict[str, set[str]] = {} for name, pool in by_name.items(): for q in pool: loc.setdefault(q.pair_id, set()).add(name) split = {pid: names for pid, names in loc.items() if len(names) > 1} assert not split, f"pair 被劈到多个池: {split}" # 每个被选中的 pair 必须两题齐全(同池内成对),不得只落单题 per_pool_pair_count: dict[tuple[str, str], int] = {} for name, pool in by_name.items(): for q in pool: key = (name, q.pair_id) per_pool_pair_count[key] = per_pool_pair_count.get(key, 0) + 1 assert all(c == 2 for c in per_pool_pair_count.values()), ( f"存在池内落单的 pair 成员: {per_pool_pair_count}" ) def test_pair_never_split_across_train_val() -> None: """PerCategoryPoolStrategy 路径:孪生对不得被 train/val 划分劈开。""" qs = [q for pid in [f"p{i:02d}" for i in range(15)] for q in _pair(pid)] # 让部分 pair 正确、部分错误,触发分层划分(非退化随机) correctness: dict[str, bool] = {} for i, pid in enumerate(f"p{i:02d}" for i in range(15)): val = i < 9 correctness[f"{pid}_o"] = val correctness[f"{pid}_m"] = val config = PoolConfig( task_types=None, seed=42, baseline_run_id="b", diag_size=0, diag_correct_ratio=0.0, val_size=0, val_correct_ratio=0.0, test_size=0, eval_min_per_class=0, train_ratio=2 / 3, test_questions_dir=None, ) pools = PerCategoryPoolStrategy().build(qs, correctness, config) loc: dict[str, set[str]] = {} for name, pool in (("diagnosis", pools.diagnosis), ("validation", pools.validation)): for q in pool: loc.setdefault(q.pair_id, set()).add(name) split = {pid: names for pid, names in loc.items() if len(names) > 1} assert not split, f"pair 被 train/val 划分劈开: {split}"