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 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 -> 退化为非分层随机划分
|
||||
|
||||
Reference in New Issue
Block a user