feat: add video-split config knobs and reproducible script

This commit is contained in:
2026-07-15 12:54:25 -04:00
parent 43d7346526
commit 6a21d80313
6 changed files with 351 additions and 0 deletions
+22
View File
@@ -115,6 +115,14 @@ def build_pools(
_VIDEO_ASSIGNMENT_LABELS = ("trainval", "test")
class InsufficientValSignal(Exception): # noqa: N818 领域名「验证信号不足」,非通用错误后缀更贴切
"""validation 池错题数不足以支撑可靠验证信号(如 McNemar 检验功效)时抛出。
fail loud(P5):不静默兜底、不放宽阈值,直接暴露 val 池错题数与所需下限,
由调用方决定放大 val_ratio / 换 trainval 归属或调低 val_wrong_min。
"""
def split_by_video_assignment(
questions: list[GeneratedQuestion],
assignment: dict[str, str],
@@ -122,6 +130,7 @@ def split_by_video_assignment(
val_ratio: float,
seed: int,
baseline_run_id: str = "",
val_wrong_min: int = 0,
) -> Pools:
"""按视频归属做原子切分:同一视频所有题绝不跨 trainval/test 池。
@@ -136,6 +145,8 @@ def split_by_video_assignment(
val_ratio: validation 占 trainval 视频组总数的比例,[0.0, 1.0]。
seed: 随机种子,保证视频组 shuffle 可复现。
baseline_run_id: 基线 run 标识;离线切分阶段可留空,由调用方回填。
val_wrong_min: validation 池最少错题数(默认 0 = 不检查,保持既有调用契约)。
> 0 时切出 val 后统计其错题数,不足即 fail loud(见 InsufficientValSignal)。
返回:
冻结的三池 Poolsdiagnosis/validation 仍是逐题 GeneratedQuestion 列表
@@ -146,6 +157,8 @@ def split_by_video_assignment(
ValueError: assignment 缺失某题 video_idfail-fast 不静默丢题)、
assignment 取值非法、correctness 缺失任一参与 Pools 的题(trainval 或
test)、或 val_ratio 越界。
InsufficientValSignal: val_wrong_min > 0 且 validation 池错题数 < val_wrong_min
(P5,验证信号不足以支撑可靠比较,直接报错而非静默放行)。
关键实现细节:
视频组 correctness 取组内全部题的 AND(组内均答对才记为 correct 组),
@@ -163,6 +176,15 @@ def split_by_video_assignment(
trainval_qs, correctness, val_ratio, random.Random(seed)
)
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(
+25
View File
@@ -7,11 +7,36 @@ build_video_records / select_split。
from __future__ import annotations
import hashlib
import random
from dataclasses import dataclass, field
from loguru import logger
def diag_fingerprint(prompt_version: str, model: str, code_version: str) -> str:
"""由 (诊断 prompt 版本, 模型名, 代码 git 短 SHA) 合成诊断口径指纹。
诊断信号以 (question_id, baseline_run_id, diag_fingerprint) 为主键持久化,
指纹隔离不同诊断配置的信号——换 prompt 版本 / 换模型 / 换代码实现都会得到
新指纹,从而 `--force` 用新指纹重跑诊断时**不覆盖旧记录**(旧指纹行仍在),
保证不同口径的诊断结果可并存、可回溯、可比对。
参数:
prompt_version: 诊断 prompt 的版本标识(如 prompts/diagnose_*.md 的版本)。
model: 执行诊断的模型名(如 "deepseek-v4")。
code_version: 诊断代码的版本(约定为 git 短 SHA)。
返回:
16 位十六进制指纹(sha256 截断),对三分量任一变化敏感、对相同三元组确定。
实现细节:
三分量用 "|" 分隔后 sha256,取前 16 位;分隔符防止 ("ab","c") 与 ("a","bc")
碰撞成同一指纹。纯函数,相同输入永远同输出,可安全用于主键。
"""
return hashlib.sha256("|".join([prompt_version, model, code_version]).encode()).hexdigest()[:16]
_EVOLUTION_TARGET = {
"extraction_failure": "tool",
"search_failure": "skill",