fix: extend correctness fail-fast to test-side pool questions (P5)

This commit is contained in:
2026-07-15 12:22:11 -04:00
parent 53989078a0
commit fd907aab46
2 changed files with 33 additions and 7 deletions
+12 -7
View File
@@ -143,7 +143,8 @@ def split_by_video_assignment(
异常: 异常:
ValueError: assignment 缺失某题 video_idfail-fast 不静默丢题)、 ValueError: assignment 缺失某题 video_idfail-fast 不静默丢题)、
assignment 取值非法、correctness 缺失 trainval 题、或 val_ratio 越界。 assignment 取值非法、correctness 缺失任一参与 Pools 的题(trainval 或
test)、或 val_ratio 越界。
关键实现细节: 关键实现细节:
视频组 correctness 取组内全部题的 AND(组内均答对才记为 correct 组), 视频组 correctness 取组内全部题的 AND(组内均答对才记为 correct 组),
@@ -170,8 +171,7 @@ def split_by_video_assignment(
baseline_run_id=baseline_run_id, baseline_run_id=baseline_run_id,
baseline_val_accuracy=baseline_val_accuracy, baseline_val_accuracy=baseline_val_accuracy,
correctness={ correctness={
q.question_id: correctness.get(q.question_id, False) q.question_id: correctness[q.question_id] for q in test_qs + validation + diagnosis
for q in test_qs + validation + diagnosis
}, },
) )
@@ -186,24 +186,29 @@ def _partition_by_video_assignment(
参数: 参数:
questions: 题目全集。 questions: 题目全集。
assignment: video_id -> "trainval" | "test" 归属字典。 assignment: video_id -> "trainval" | "test" 归属字典。
correctness: question_id -> 基线是否答对;仅对 trainval 题强制完整。 correctness: question_id -> 基线是否答对;对全部参与 Pools 的题(trainval
与 test 双侧)强制完整。
返回: 返回:
(trainval_qs, test_qs) 逐题列表元组,划分依据每题的 video_id 归属。 (trainval_qs, test_qs) 逐题列表元组,划分依据每题的 video_id 归属。
异常: 异常:
ValueError: assignment 取值非法、缺失某题 video_id、或 correctness 缺失 ValueError: assignment 取值非法、缺失某题 video_id、或 correctness 缺失
trainval 题(fail-fast,不静默丢题)。 任一参与 Pools 的题(trainval 或 testfail-fast,不静默丢题)。
""" """
_assert_valid_assignment(questions, assignment) _assert_valid_assignment(questions, assignment)
trainval_qs = [q for q in questions if assignment[q.video_id] == "trainval"] 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"] 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] # Pools.correctness 会为 test + validation + diagnosis 全体写入基线对错,
# 故 test 侧同样必须有 correctness,缺失即报错而非静默兜底 False(P5)。
missing_correctness = [
q.question_id for q in trainval_qs + test_qs if q.question_id not in correctness
]
if missing_correctness: if missing_correctness:
raise ValueError( raise ValueError(
f"correctness 缺失 {len(missing_correctness)} trainval 题: {missing_correctness[:5]}" f"correctness 缺失 {len(missing_correctness)} 道题: {missing_correctness[:5]}"
) )
return trainval_qs, test_qs return trainval_qs, test_qs
+21
View File
@@ -71,6 +71,27 @@ def test_invalid_assignment_value_fails_fast():
) )
def test_val_ratio_out_of_range_fails_fast():
"""val_ratio 越界(<0 或 >1)时 fail-fast。"""
qs = [_q("v1-1", "v1")]
assignment = {"v1": "trainval"}
correctness = {"v1-1": True}
for bad in (-0.1, 1.5):
with pytest.raises(ValueError, match="val_ratio"):
split_by_video_assignment(
qs, assignment, correctness=correctness, val_ratio=bad, seed=0
)
def test_missing_correctness_on_test_side_fails_fast():
"""test 侧某题缺 correctness 时 fail-fast,不静默兜底 False 污染评估口径。"""
qs = [_q("v1-1", "v1"), _q("v2-1", "v2")]
assignment = {"v1": "trainval", "v2": "test"}
correctness = {"v1-1": True} # 缺 test 侧 v2-1
with pytest.raises(ValueError, match="correctness"):
split_by_video_assignment(qs, assignment, correctness=correctness, val_ratio=0.0, seed=0)
def test_video_group_atomic_in_trainval_split(): def test_video_group_atomic_in_trainval_split():
"""trainval 内 val 切分以视频组为原子:同 video 的题整组同落 train 或 val。""" """trainval 内 val 切分以视频组为原子:同 video 的题整组同落 train 或 val。"""
qs = [ qs = [