refactor: add video-atomic pool split (algo #5 gate input preserved)
This commit is contained in:
+202
-16
@@ -17,7 +17,7 @@ from typing import TYPE_CHECKING
|
|||||||
|
|
||||||
from loguru import logger
|
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 app.question_gen import stratified_sample
|
||||||
from core.types import GeneratedQuestion, PoolConfig
|
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:
|
class GlobalPoolStrategy:
|
||||||
"""全局三分策略:test -> val -> diag progressive exclusion。
|
"""全局三分策略: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(
|
def _assert_correctness_complete(
|
||||||
units: list[QuestionUnit],
|
units: list[QuestionUnit],
|
||||||
correctness: dict[str, bool],
|
correctness: dict[str, bool],
|
||||||
@@ -849,8 +1035,8 @@ class PerCategoryPoolStrategy:
|
|||||||
|
|
||||||
_assert_correctness_complete(units, correctness)
|
_assert_correctness_complete(units, correctness)
|
||||||
|
|
||||||
correct_units = [u for u in units if _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_correct(u, correctness)]
|
wrong_units = [u for u in units if not unit_correctness(u, correctness, strict=False)]
|
||||||
n_correct = len(correct_units)
|
n_correct = len(correct_units)
|
||||||
|
|
||||||
# 全 correct 或全 wrong -> 退化为非分层随机划分
|
# 全 correct 或全 wrong -> 退化为非分层随机划分
|
||||||
|
|||||||
@@ -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 == []
|
||||||
Reference in New Issue
Block a user