feat: tier-aware diag/val split with val-power repair (design 5.1)
This commit is contained in:
+49
-14
@@ -131,6 +131,7 @@ def split_by_video_assignment(
|
||||
seed: int,
|
||||
baseline_run_id: str = "",
|
||||
val_wrong_min: int = 0,
|
||||
wrong_tier_by_video: dict[str, int] | None = None,
|
||||
) -> Pools:
|
||||
"""按视频归属做原子切分:同一视频所有题绝不跨 trainval/test 池。
|
||||
|
||||
@@ -146,7 +147,11 @@ def split_by_video_assignment(
|
||||
seed: 随机种子,保证视频组 shuffle 可复现。
|
||||
baseline_run_id: 基线 run 标识;离线切分阶段可留空,由调用方回填。
|
||||
val_wrong_min: validation 池最少错题数(默认 0 = 不检查,保持既有调用契约)。
|
||||
> 0 时切出 val 后统计其错题数,不足即 fail loud(见 InsufficientValSignal)。
|
||||
> 0 时切分时保证(不足则从 diag 换入低 T2 错题组补足,耗尽 fail-loud,
|
||||
见 InsufficientValSignal)。
|
||||
wrong_tier_by_video: video_id -> 该视频错题中 T2(defect) 数量;透传给
|
||||
_split_trainval_by_video_group 做 tier 感知 diag/val 分配,None 时退化为
|
||||
原随机 shuffle。
|
||||
|
||||
返回:
|
||||
冻结的三池 Pools:diagnosis/validation 仍是逐题 GeneratedQuestion 列表
|
||||
@@ -173,18 +178,11 @@ def split_by_video_assignment(
|
||||
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)
|
||||
trainval_qs, correctness, val_ratio, random.Random(seed),
|
||||
wrong_tier_by_video=wrong_tier_by_video,
|
||||
val_wrong_min=val_wrong_min,
|
||||
)
|
||||
|
||||
if val_wrong_min > 0:
|
||||
val_wrong = sum(1 for q in validation if not correctness[q.question_id])
|
||||
if val_wrong < val_wrong_min:
|
||||
raise InsufficientValSignal(
|
||||
f"validation 池错题数 {val_wrong} < val_wrong_min={val_wrong_min},"
|
||||
"验证信号不足以支撑可靠比较(如 McNemar 检验功效),"
|
||||
"请放大 val_ratio / 调整 trainval 归属或调低阈值。"
|
||||
)
|
||||
|
||||
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(
|
||||
@@ -291,6 +289,8 @@ def _split_trainval_by_video_group(
|
||||
correctness: dict[str, bool],
|
||||
val_ratio: float,
|
||||
rng: random.Random,
|
||||
wrong_tier_by_video: dict[str, int] | None = None,
|
||||
val_wrong_min: int = 0,
|
||||
) -> tuple[list[GeneratedQuestion], list[GeneratedQuestion]]:
|
||||
"""以视频组为原子对 trainval 题集做 correctness 分层,切出 (diagnosis, validation)。
|
||||
|
||||
@@ -299,6 +299,12 @@ def _split_trainval_by_video_group(
|
||||
correctness: question_id -> 基线是否答对;视频组正确性取组内全部题 AND。
|
||||
val_ratio: validation 占视频组总数的比例。
|
||||
rng: 随机数生成器,保证视频组 shuffle 可复现。
|
||||
wrong_tier_by_video: video_id -> 该视频错题中 T2(defect) 的数量。提供时错题
|
||||
视频组按 T2 含量升序进 val(T2 高的组保留在 diagnosis,把高价值缺陷信号
|
||||
留给诊断),确定性排序取代随机 shuffle;None 时退化为原随机 shuffle。
|
||||
val_wrong_min: validation 池最少错题数(切分时保证功效)。> 0 且初分 val 错题
|
||||
不足时,从 diag 侧的错题组按 T2 升序换入 val 直到满足(每组至多移动一次),
|
||||
耗尽仍不足则抛 InsufficientValSignal(fail loud,P5)。
|
||||
|
||||
返回:
|
||||
(diagnosis, validation) 逐题列表元组;同一 video 的全部题整组落在同一侧,
|
||||
@@ -306,7 +312,8 @@ def _split_trainval_by_video_group(
|
||||
|
||||
关键实现细节:
|
||||
与 _split_one_category 同构:先按视频组 correctness 分正确组/错误组,按比例
|
||||
把 n_val 个组分层落入 validation(全正确或全错误时退化为非分层随机划分),
|
||||
把 n_val 个组分层落入 validation(全正确退化为非分层随机划分;全错误时若有
|
||||
wrong_tier_by_video 仍按 T2 升序分配,否则随机划分),
|
||||
再把选中组内所有题展开。视频组按 video_id 排序后再 shuffle,保证确定性。
|
||||
"""
|
||||
groups: dict[str, list[GeneratedQuestion]] = defaultdict(list)
|
||||
@@ -320,7 +327,11 @@ def _split_trainval_by_video_group(
|
||||
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:
|
||||
if n_correct == 0 and wrong_tier_by_video is not None:
|
||||
# 全部错误 + 有 tier 信号:按 T2 升序,低 T2 组优先进 val(保留高 T2 在 diag)
|
||||
wrong_vids.sort(key=lambda v: (wrong_tier_by_video.get(v, 0), v))
|
||||
val_vids = set(wrong_vids[:n_val])
|
||||
elif 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)
|
||||
@@ -330,9 +341,33 @@ def _split_trainval_by_video_group(
|
||||
val_correct = math.floor(n_correct * n_val / n_total)
|
||||
val_wrong = n_val - val_correct
|
||||
rng.shuffle(correct_vids)
|
||||
rng.shuffle(wrong_vids)
|
||||
if wrong_tier_by_video is None:
|
||||
rng.shuffle(wrong_vids)
|
||||
else:
|
||||
# T2 少的错题组优先进 val(保留 T2 高的组在 diag),确定性排序
|
||||
wrong_vids.sort(key=lambda v: (wrong_tier_by_video.get(v, 0), v))
|
||||
val_vids = set(correct_vids[:val_correct] + wrong_vids[:val_wrong])
|
||||
|
||||
if val_wrong_min > 0:
|
||||
val_wrong_now = sum(
|
||||
1 for v in val_vids for q in groups[v] if not correctness[q.question_id]
|
||||
)
|
||||
# diag 侧仍在的错题组,按 T2 升序(低价值优先移交 val)
|
||||
diag_wrong_pool = sorted(
|
||||
(v for v in wrong_vids if v not in val_vids),
|
||||
key=lambda v: ((wrong_tier_by_video or {}).get(v, 0), v),
|
||||
)
|
||||
for v in diag_wrong_pool:
|
||||
if val_wrong_now >= val_wrong_min:
|
||||
break
|
||||
val_vids.add(v)
|
||||
val_wrong_now += sum(1 for q in groups[v] if not correctness[q.question_id])
|
||||
if val_wrong_now < val_wrong_min:
|
||||
raise InsufficientValSignal(
|
||||
f"trainval 错题不足以让 val 达到 val_wrong_min={val_wrong_min}"
|
||||
f"(修复后仅 {val_wrong_now}),请放大 val_ratio 或调整 trainval 归属。"
|
||||
)
|
||||
|
||||
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
|
||||
|
||||
Reference in New Issue
Block a user