diff --git a/app/harness/pools.py b/app/harness/pools.py index 32e3fb6..d91c6d6 100644 --- a/app/harness/pools.py +++ b/app/harness/pools.py @@ -17,7 +17,7 @@ from typing import TYPE_CHECKING from loguru import logger -from app.harness.question_units import build_units, flatten_units +from app.harness.question_units import build_units, flatten_units, unit_correctness from app.question_gen import stratified_sample from core.types import GeneratedQuestion, PoolConfig @@ -111,6 +111,205 @@ def build_pools( ) +_VIDEO_ASSIGNMENT_LABELS = ("trainval", "test") + + +def split_by_video_assignment( + questions: list[GeneratedQuestion], + assignment: dict[str, str], + correctness: dict[str, bool], + val_ratio: float, + seed: int, + baseline_run_id: str = "", +) -> Pools: + """按视频归属做原子切分:同一视频所有题绝不跨 trainval/test 池。 + + 切分原子从 unit 提升为 **视频组**(同 video 的全部题同进同出),彻底杜绝 + 同视频多题散落不同池造成的内容泄漏。trainval 题集内部再以视频组为原子做 + correctness 分层,切出 validation(占 val_ratio)与 diagnosis(其余)。 + + 参数: + questions: 题目全集。 + assignment: video_id -> "trainval" | "test" 归属字典(由选择器上游产出)。 + correctness: question_id -> 基线是否答对;trainval 分层与验证池准确率均依赖它。 + val_ratio: validation 占 trainval 视频组总数的比例,[0.0, 1.0]。 + seed: 随机种子,保证视频组 shuffle 可复现。 + baseline_run_id: 基线 run 标识;离线切分阶段可留空,由调用方回填。 + + 返回: + 冻结的三池 Pools:diagnosis/validation 仍是逐题 GeneratedQuestion 列表 + (元素粒度不变,仅改变"哪些视频进哪个池"),test 为全部 test 题。 + baseline_val_accuracy = validation 池正确率。 + + 异常: + ValueError: assignment 缺失某题 video_id(fail-fast 不静默丢题)、 + assignment 取值非法、correctness 缺失 trainval 题、或 val_ratio 越界。 + + 关键实现细节: + 视频组 correctness 取组内全部题的 AND(组内均答对才记为 correct 组), + 据此在 trainval 内做与 _split_one_category 同构的比例分层,但原子是视频组。 + val_ratio 决定 validation 组数:val_correct = floor(n_correct * n_val / n_total), + 余额补 wrong 组,全 correct / 全 wrong 时退化为非分层随机划分。下游 + gate_ladder(信息阶梯冷启动 2:1,核心算法保真 #5)消费的 unit 结构不变。 + """ + if not 0.0 <= val_ratio <= 1.0: + raise ValueError(f"val_ratio 必须在 [0.0, 1.0],实际 {val_ratio}") + + trainval_qs, test_qs = _partition_by_video_assignment(questions, assignment, correctness) + + diagnosis, validation = _split_trainval_by_video_group( + trainval_qs, correctness, val_ratio, random.Random(seed) + ) + + val_correct = sum(1 for q in validation if correctness.get(q.question_id)) + baseline_val_accuracy = val_correct / len(validation) if validation else 0.0 + return Pools( + diagnosis=diagnosis, + validation=validation, + test=test_qs, + baseline_run_id=baseline_run_id, + baseline_val_accuracy=baseline_val_accuracy, + correctness={ + q.question_id: correctness.get(q.question_id, False) + for q in test_qs + validation + diagnosis + }, + ) + + +def _partition_by_video_assignment( + questions: list[GeneratedQuestion], + assignment: dict[str, str], + correctness: dict[str, bool], +) -> tuple[list[GeneratedQuestion], list[GeneratedQuestion]]: + """校验归属字典并按 video 归属把题划成 (trainval_qs, test_qs)。 + + 参数: + questions: 题目全集。 + assignment: video_id -> "trainval" | "test" 归属字典。 + correctness: question_id -> 基线是否答对;仅对 trainval 题强制完整。 + + 返回: + (trainval_qs, test_qs) 逐题列表元组,划分依据每题的 video_id 归属。 + + 异常: + ValueError: assignment 取值非法、缺失某题 video_id、或 correctness 缺失 + trainval 题(fail-fast,不静默丢题)。 + """ + _assert_valid_assignment(questions, assignment) + + trainval_qs = [q for q in questions if assignment[q.video_id] == "trainval"] + test_qs = [q for q in questions if assignment[q.video_id] == "test"] + + missing_correctness = [q.question_id for q in trainval_qs if q.question_id not in correctness] + if missing_correctness: + raise ValueError( + f"correctness 缺失 {len(missing_correctness)} 道 trainval 题: {missing_correctness[:5]}" + ) + + return trainval_qs, test_qs + + +def _assert_valid_assignment( + questions: list[GeneratedQuestion], + assignment: dict[str, str], +) -> None: + """校验归属字典取值合法且覆盖全部题的 video_id,否则 fail-fast。 + + 参数: + questions: 题目全集。 + assignment: video_id -> "trainval" | "test" 归属字典。 + + 异常: + ValueError: assignment 含非法取值,或缺失某题的 video_id。 + """ + bad_labels = {v for v in assignment.values() if v not in _VIDEO_ASSIGNMENT_LABELS} + if bad_labels: + raise ValueError( + f"assignment 含非法归属值 {sorted(bad_labels)},仅允许 {_VIDEO_ASSIGNMENT_LABELS}" + ) + + missing_videos = sorted({q.video_id for q in questions if q.video_id not in assignment}) + if missing_videos: + raise ValueError(f"assignment 缺失 {len(missing_videos)} 个 video_id: {missing_videos[:5]}") + + +def _partition_video_groups_by_correctness( + groups: dict[str, list[GeneratedQuestion]], + correctness: dict[str, bool], +) -> tuple[list[str], list[str]]: + """按视频组正确性把 video_id 分成 (correct_vids, wrong_vids)。 + + 组正确性取组内全部题的 AND(组内均答对才记为 correct 组),排序保证确定性。 + + 参数: + groups: video_id -> 该视频全部题列表。 + correctness: question_id -> 基线是否答对(调用方已校验完整)。 + + 返回: + (correct_vids, wrong_vids) 两个 video_id 列表,按 video_id 升序。 + """ + correct_vids: list[str] = [] + wrong_vids: list[str] = [] + for vid in sorted(groups.keys()): + if all(correctness[q.question_id] for q in groups[vid]): + correct_vids.append(vid) + else: + wrong_vids.append(vid) + return correct_vids, wrong_vids + + +def _split_trainval_by_video_group( + trainval_qs: list[GeneratedQuestion], + correctness: dict[str, bool], + val_ratio: float, + rng: random.Random, +) -> tuple[list[GeneratedQuestion], list[GeneratedQuestion]]: + """以视频组为原子对 trainval 题集做 correctness 分层,切出 (diagnosis, validation)。 + + 参数: + trainval_qs: trainval 归属的全部题(correctness 已在调用方校验完整)。 + correctness: question_id -> 基线是否答对;视频组正确性取组内全部题 AND。 + val_ratio: validation 占视频组总数的比例。 + rng: 随机数生成器,保证视频组 shuffle 可复现。 + + 返回: + (diagnosis, validation) 逐题列表元组;同一 video 的全部题整组落在同一侧, + 两侧互斥且并集 == trainval_qs。 + + 关键实现细节: + 与 _split_one_category 同构:先按视频组 correctness 分正确组/错误组,按比例 + 把 n_val 个组分层落入 validation(全正确或全错误时退化为非分层随机划分), + 再把选中组内所有题展开。视频组按 video_id 排序后再 shuffle,保证确定性。 + """ + groups: dict[str, list[GeneratedQuestion]] = defaultdict(list) + for q in trainval_qs: + groups[q.video_id].append(q) + + video_ids = sorted(groups.keys()) + n_total = len(video_ids) + n_val = round(n_total * val_ratio) + + correct_vids, wrong_vids = _partition_video_groups_by_correctness(groups, correctness) + n_correct = len(correct_vids) + + if n_correct == 0 or n_correct == n_total: + label = "全部正确" if n_correct == n_total else "全部错误" + logger.warning("trainval 视频组 {} ({} 组),退化为非分层随机划分", label, n_total) + shuffled = list(video_ids) + rng.shuffle(shuffled) + val_vids = set(shuffled[:n_val]) + else: + val_correct = math.floor(n_correct * n_val / n_total) + val_wrong = n_val - val_correct + rng.shuffle(correct_vids) + rng.shuffle(wrong_vids) + val_vids = set(correct_vids[:val_correct] + wrong_vids[:val_wrong]) + + diagnosis = [q for q in trainval_qs if q.video_id not in val_vids] + validation = [q for q in trainval_qs if q.video_id in val_vids] + return diagnosis, validation + + class GlobalPoolStrategy: """全局三分策略:test -> val -> diag progressive exclusion。 @@ -173,19 +372,6 @@ class GlobalPoolStrategy: ) -def _unit_correct(unit: QuestionUnit, correctness: dict[str, bool]) -> bool: - """单元级正确性:成员全部答对才算对(缺失按 False,宽松口径)。 - - 参数: - unit: 目标单元(single 1 题,pair 2 题)。 - correctness: question_id -> 基线是否答对。 - - 返回: - pair 走双向 AND、single 即单题正确性;任一成员缺失或答错即 False。 - """ - return all(correctness.get(q.question_id, False) for q in unit.questions) - - def _assert_correctness_complete( units: list[QuestionUnit], correctness: dict[str, bool], @@ -849,8 +1035,8 @@ class PerCategoryPoolStrategy: _assert_correctness_complete(units, correctness) - correct_units = [u for u in units if _unit_correct(u, correctness)] - wrong_units = [u for u in units if not _unit_correct(u, correctness)] + correct_units = [u for u in units if unit_correctness(u, correctness, strict=False)] + wrong_units = [u for u in units if not unit_correctness(u, correctness, strict=False)] n_correct = len(correct_units) # 全 correct 或全 wrong -> 退化为非分层随机划分 diff --git a/tests/unit/test_pools_video_atomic.py b/tests/unit/test_pools_video_atomic.py new file mode 100644 index 0000000..aacc71c --- /dev/null +++ b/tests/unit/test_pools_video_atomic.py @@ -0,0 +1,116 @@ +"""视频原子切分 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 == []