feat: add video-split config knobs and reproducible script
This commit is contained in:
@@ -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)。
|
||||
|
||||
返回:
|
||||
冻结的三池 Pools:diagnosis/validation 仍是逐题 GeneratedQuestion 列表
|
||||
@@ -146,6 +157,8 @@ def split_by_video_assignment(
|
||||
ValueError: assignment 缺失某题 video_id(fail-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(
|
||||
|
||||
@@ -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",
|
||||
|
||||
Reference in New Issue
Block a user