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
+202 -16
View File
@@ -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 标识;离线切分阶段可留空,由调用方回填。
返回:
冻结的三池 Poolsdiagnosis/validation 仍是逐题 GeneratedQuestion 列表
(元素粒度不变,仅改变"哪些视频进哪个池"),test 为全部 test 题。
baseline_val_accuracy = validation 池正确率。
异常:
ValueError: assignment 缺失某题 video_idfail-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 -> 退化为非分层随机划分