fix: extend correctness fail-fast to test-side pool questions (P5)
This commit is contained in:
+12
-7
@@ -143,7 +143,8 @@ def split_by_video_assignment(
|
||||
|
||||
异常:
|
||||
ValueError: assignment 缺失某题 video_id(fail-fast 不静默丢题)、
|
||||
assignment 取值非法、correctness 缺失 trainval 题、或 val_ratio 越界。
|
||||
assignment 取值非法、correctness 缺失任一参与 Pools 的题(trainval 或
|
||||
test)、或 val_ratio 越界。
|
||||
|
||||
关键实现细节:
|
||||
视频组 correctness 取组内全部题的 AND(组内均答对才记为 correct 组),
|
||||
@@ -170,8 +171,7 @@ def split_by_video_assignment(
|
||||
baseline_run_id=baseline_run_id,
|
||||
baseline_val_accuracy=baseline_val_accuracy,
|
||||
correctness={
|
||||
q.question_id: correctness.get(q.question_id, False)
|
||||
for q in test_qs + validation + diagnosis
|
||||
q.question_id: correctness[q.question_id] for q in test_qs + validation + diagnosis
|
||||
},
|
||||
)
|
||||
|
||||
@@ -186,24 +186,29 @@ def _partition_by_video_assignment(
|
||||
参数:
|
||||
questions: 题目全集。
|
||||
assignment: video_id -> "trainval" | "test" 归属字典。
|
||||
correctness: question_id -> 基线是否答对;仅对 trainval 题强制完整。
|
||||
correctness: question_id -> 基线是否答对;对全部参与 Pools 的题(trainval
|
||||
与 test 双侧)强制完整。
|
||||
|
||||
返回:
|
||||
(trainval_qs, test_qs) 逐题列表元组,划分依据每题的 video_id 归属。
|
||||
|
||||
异常:
|
||||
ValueError: assignment 取值非法、缺失某题 video_id、或 correctness 缺失
|
||||
trainval 题(fail-fast,不静默丢题)。
|
||||
任一参与 Pools 的题(trainval 或 test,fail-fast,不静默丢题)。
|
||||
"""
|
||||
_assert_valid_assignment(questions, assignment)
|
||||
|
||||
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"]
|
||||
|
||||
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:
|
||||
raise ValueError(
|
||||
f"correctness 缺失 {len(missing_correctness)} 道 trainval 题: {missing_correctness[:5]}"
|
||||
f"correctness 缺失 {len(missing_correctness)} 道题: {missing_correctness[:5]}"
|
||||
)
|
||||
|
||||
return trainval_qs, test_qs
|
||||
|
||||
@@ -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():
|
||||
"""trainval 内 val 切分以视频组为原子:同 video 的题整组同落 train 或 val。"""
|
||||
qs = [
|
||||
|
||||
Reference in New Issue
Block a user