From fd907aab465a7c402c3df1a5cce42642ba13c2a6 Mon Sep 17 00:00:00 2001 From: iomgaa Date: Wed, 15 Jul 2026 12:22:11 -0400 Subject: [PATCH] fix: extend correctness fail-fast to test-side pool questions (P5) --- app/harness/pools.py | 19 ++++++++++++------- tests/unit/test_pools_video_atomic.py | 21 +++++++++++++++++++++ 2 files changed, 33 insertions(+), 7 deletions(-) diff --git a/app/harness/pools.py b/app/harness/pools.py index d91c6d6..91f04a4 100644 --- a/app/harness/pools.py +++ b/app/harness/pools.py @@ -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 diff --git a/tests/unit/test_pools_video_atomic.py b/tests/unit/test_pools_video_atomic.py index aacc71c..6ac5664 100644 --- a/tests/unit/test_pools_video_atomic.py +++ b/tests/unit/test_pools_video_atomic.py @@ -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 = [