feat: tier-aware diag/val split with val-power repair (design 5.1)
This commit is contained in:
@@ -13,7 +13,7 @@ from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import sqlite3
|
||||
from collections import Counter
|
||||
from collections import Counter, defaultdict
|
||||
from dataclasses import asdict, dataclass, field
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
@@ -59,6 +59,8 @@ class SplitBuildConfig:
|
||||
select_seed: 贪心选择器预洗牌种子(打破等增益平局)。
|
||||
val_ratio: validation 占 trainval 视频组总数的比例。
|
||||
split_seed: 视频组题级切分的洗牌种子。
|
||||
val_wrong_min: validation 池最少错题数,切分时保证功效(不足则从 diag 换入
|
||||
低 T2 错题组补足,耗尽 fail-loud)。
|
||||
"""
|
||||
|
||||
n_trainval: int
|
||||
@@ -68,6 +70,7 @@ class SplitBuildConfig:
|
||||
select_seed: int
|
||||
val_ratio: float
|
||||
split_seed: int
|
||||
val_wrong_min: int
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
@@ -119,10 +122,10 @@ def build_split(
|
||||
test → 加载题库并以视频归属切三池 → 原子冻结 pools.json → 写溯源 manifest →
|
||||
六条防御断言 fail-fast 校验。
|
||||
|
||||
契约(Task 11,非疏漏):build_split 有意保持 val_wrong_min-agnostic——内部调
|
||||
split_by_video_assignment 时不传 val_wrong_min(默认 0,不校验 validation 错题
|
||||
数)。McNemar 功效护栏是切分**冻结后**的独立校验,由 CLI 的 check_mcnemar_power
|
||||
在 build_split 返回后执行;切分构造本身不因功效阈失败,二者关注点分离。
|
||||
契约(Task 11):val_wrong_min 前置到切分内保证功效——build_split 计算
|
||||
wrong_tier_by_video 并连同 config.val_wrong_min 传入 split_by_video_assignment,
|
||||
切分时若 val 错题不足即从 diag 换入低 T2 错题组补足(耗尽 fail-loud)。CLI 的
|
||||
check_mcnemar_power 作切分冻结后的冗余最终确认。
|
||||
|
||||
参数:
|
||||
db_path: harness.db 路径(只读读取 predictions,不改动)。
|
||||
@@ -181,6 +184,11 @@ def build_split(
|
||||
# Phase 3: 加载题库 + 视频归属切三池 + 原子冻结。
|
||||
questions = load_benchmark(questions_dir)
|
||||
correctness = {pred["question_id"]: pred["correct"] for pred in preds}
|
||||
tier_by_q = {row["question_id"]: row["tier"] for row in signal_rows}
|
||||
wrong_tier_by_video: dict[str, int] = defaultdict(int)
|
||||
for pred in preds:
|
||||
if not pred["correct"] and tier_by_q.get(pred["question_id"]) == "T2":
|
||||
wrong_tier_by_video[pred["video_id"]] += 1
|
||||
pools = split_by_video_assignment(
|
||||
questions,
|
||||
assignment,
|
||||
@@ -188,6 +196,8 @@ def build_split(
|
||||
config.val_ratio,
|
||||
config.split_seed,
|
||||
baseline_run_id=baseline_run_id,
|
||||
val_wrong_min=config.val_wrong_min,
|
||||
wrong_tier_by_video=dict(wrong_tier_by_video),
|
||||
)
|
||||
save_pools(pools, out_path)
|
||||
|
||||
|
||||
+48
-13
@@ -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,16 +178,9 @@ 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)
|
||||
)
|
||||
|
||||
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 归属或调低阈值。"
|
||||
trainval_qs, correctness, val_ratio, random.Random(seed),
|
||||
wrong_tier_by_video=wrong_tier_by_video,
|
||||
val_wrong_min=val_wrong_min,
|
||||
)
|
||||
|
||||
val_correct = sum(1 for q in validation if correctness.get(q.question_id))
|
||||
@@ -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)
|
||||
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
|
||||
|
||||
@@ -477,8 +477,9 @@ def load_questions_by_id(questions_dir: Path) -> dict[str, GeneratedQuestion]:
|
||||
def check_mcnemar_power(pools: Pools, val_wrong_min: int) -> int:
|
||||
"""校验 validation 池错题数达 McNemar 功效阈,不足即 fail loud。
|
||||
|
||||
build_split 按契约用 split_by_video_assignment(val_wrong_min=0)(不破契约),
|
||||
故功效护栏在 capstone 层单独核验:val 错题数 < 阈 → 验证信号不足以支撑可靠比较。
|
||||
val_wrong_min 已前置到 build_split 内的切分保证功效(不足即从 diag 换入低 T2
|
||||
错题组补足,耗尽 fail-loud);本函数作切分冻结后的冗余最终确认:val 错题数 < 阈
|
||||
→ 验证信号不足以支撑可靠比较。
|
||||
|
||||
参数:
|
||||
pools: 冻结三池(含 validation 与 correctness)。
|
||||
@@ -571,6 +572,7 @@ async def run_pipeline(
|
||||
select_seed=config.seed,
|
||||
val_ratio=config.val_ratio,
|
||||
split_seed=config.seed,
|
||||
val_wrong_min=config.val_wrong_min,
|
||||
),
|
||||
out_path=out_dir / "pools.json",
|
||||
manifest_path=out_dir / "split_manifest.json",
|
||||
|
||||
@@ -101,6 +101,7 @@ def test_end_to_end_freezes_valid_pools(tmp_path: Path) -> None:
|
||||
select_seed=7,
|
||||
val_ratio=0.3,
|
||||
split_seed=7,
|
||||
val_wrong_min=0,
|
||||
)
|
||||
|
||||
result = build_split(
|
||||
|
||||
@@ -135,3 +135,60 @@ def test_baseline_val_accuracy_reflects_validation():
|
||||
assert len(pools.validation) == 2
|
||||
assert pools.baseline_val_accuracy == pytest.approx(0.5)
|
||||
assert pools.diagnosis == []
|
||||
|
||||
|
||||
def test_tier_aware_keeps_high_t2_in_diag():
|
||||
"""错题视频组按 T2 含量升序进 val:T2 高的组保留在 diagnosis。"""
|
||||
from app.harness.pools import split_by_video_assignment
|
||||
from core.types import GeneratedQuestion
|
||||
|
||||
def _q(qid, vid):
|
||||
return GeneratedQuestion(
|
||||
question_id=qid, video_id=vid, task_type="X", question="q",
|
||||
options=["A", "B"], answer="A", source_nodes=[], difficulty="easy",
|
||||
)
|
||||
|
||||
# 4 个错题视频(每视频 1 题),T2 数分别 2/1/0/0
|
||||
questions = [_q(f"{v}-1", v) for v in ("vA", "vB", "vC", "vD")]
|
||||
assignment = {v: "trainval" for v in ("vA", "vB", "vC", "vD")}
|
||||
correctness = {f"{v}-1": False for v in ("vA", "vB", "vC", "vD")}
|
||||
wrong_tier = {"vA": 2, "vB": 1, "vC": 0, "vD": 0}
|
||||
|
||||
pools = split_by_video_assignment(
|
||||
questions, assignment, correctness, val_ratio=0.5, seed=7,
|
||||
wrong_tier_by_video=wrong_tier,
|
||||
)
|
||||
diag_vids = {q.video_id for q in pools.diagnosis}
|
||||
# T2 最高的 vA 必留 diag;T2=0 的组优先进 val
|
||||
assert "vA" in diag_vids
|
||||
assert "vB" in diag_vids
|
||||
|
||||
|
||||
def test_val_wrong_min_repair_pulls_from_diag():
|
||||
"""val 错题不足 val_wrong_min 时从 diag 换入低 T2 错题组补足。"""
|
||||
from app.harness.pools import split_by_video_assignment
|
||||
from core.types import GeneratedQuestion
|
||||
|
||||
def _q(qid, vid, correct):
|
||||
return GeneratedQuestion(
|
||||
question_id=qid, video_id=vid, task_type="X", question="q",
|
||||
options=["A", "B"], answer="A", source_nodes=[], difficulty="easy",
|
||||
)
|
||||
|
||||
# 8 错题视频 + 2 正确视频;val_ratio 小使初分 val 错题不足,触发修复
|
||||
vids_wrong = [f"w{i}" for i in range(8)]
|
||||
vids_correct = ["c0", "c1"]
|
||||
questions = [_q(f"{v}-1", v, False) for v in vids_wrong] + [
|
||||
_q(f"{v}-1", v, True) for v in vids_correct
|
||||
]
|
||||
assignment = {v: "trainval" for v in vids_wrong + vids_correct}
|
||||
correctness = {f"{v}-1": False for v in vids_wrong}
|
||||
correctness.update({f"{v}-1": True for v in vids_correct})
|
||||
wrong_tier = {v: i for i, v in enumerate(vids_wrong)} # 递增 T2
|
||||
|
||||
pools = split_by_video_assignment(
|
||||
questions, assignment, correctness, val_ratio=0.1, seed=7,
|
||||
wrong_tier_by_video=wrong_tier, val_wrong_min=4,
|
||||
)
|
||||
val_wrong = sum(1 for q in pools.validation if not correctness[q.question_id])
|
||||
assert val_wrong >= 4, f"功效修复后 val 错题 {val_wrong} < 4"
|
||||
|
||||
Reference in New Issue
Block a user