ddb9a44f75
build_pools/_sample_excluding 与 PerCategoryPoolStrategy._split_one_category/ build_incremental 两条切分路径均改为以 QuestionUnit 为采样原子:progressive exclusion 互斥集合与 train/val 分层划分都按 unit_id 计数(pair 计 1 个 unit), 命中单元整体展开,AR 孪生对两题永不落入不同池/split。 复用 app.harness.question_units 的 build_units/flatten_units,不重写分组逻辑。 single-only 输入下 unit 与 question 一一对应、rng 消耗量不变,采样与划分结果 与逐题口径完全一致;抽出 _unit_correct/_assert_correctness_complete 两个 helper 将 _split_one_category 复杂度压回基线以下。 新增 tests/unit/test_pools_pair_atomic.py 覆盖两条路径的 pair 原子性回归。
117 lines
3.9 KiB
Python
117 lines
3.9 KiB
Python
"""pair 原子性回归测试:三池切分不得把孪生对劈到不同池。
|
||
|
||
覆盖 pools.py 的两条切分路径:
|
||
- GlobalPoolStrategy 路径(build_pools 的 test->val->diag progressive exclusion);
|
||
- PerCategoryPoolStrategy 路径(_split_one_category 的 train/val 分层划分)。
|
||
|
||
核心不变量:同一 pair_id 的两题必落在同一个池(或同一 split),绝不被拆散。
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
from app.harness.pools import PerCategoryPoolStrategy, build_pools
|
||
from core.types import GeneratedQuestion, PoolConfig
|
||
|
||
|
||
def _pair(pid: str) -> list[GeneratedQuestion]:
|
||
"""构造一个合法孪生对(original + mirror),共享 pair_id / unit_id / flip_axis。
|
||
|
||
参数:
|
||
pid: 该孪生对的共享标识。
|
||
|
||
返回:
|
||
含 original 与 mirror 两条题目的列表。
|
||
"""
|
||
base = {
|
||
"video_id": "v",
|
||
"task_type": "Action Reasoning",
|
||
"question": "?",
|
||
"options": ("A. a", "B. b", "C. c", "D. d"),
|
||
"answer": "A",
|
||
"source_nodes": ("n",),
|
||
"difficulty": "hard",
|
||
"pair_id": pid,
|
||
"flip_axis": "before_after",
|
||
}
|
||
return [
|
||
GeneratedQuestion(question_id=f"{pid}_o", question_role="pair_original", **base),
|
||
GeneratedQuestion(question_id=f"{pid}_m", question_role="pair_mirror", **base),
|
||
]
|
||
|
||
|
||
def test_pair_never_split_across_pools() -> None:
|
||
"""build_pools(Global 路径)三池切分后,任一 pair_id 只出现在一个池。"""
|
||
qs = [q for pid in [f"p{i:02d}" for i in range(12)] for q in _pair(pid)]
|
||
cfg = {
|
||
"size": 4,
|
||
"correct_ratio": None,
|
||
"task_types": None,
|
||
"seed": 7,
|
||
"min_per_class": None,
|
||
}
|
||
pools = build_pools(
|
||
qs,
|
||
correctness={},
|
||
diag_cfg=cfg,
|
||
val_cfg=cfg,
|
||
test_cfg={"size": 4, "seed": 7},
|
||
baseline_run_id="b",
|
||
)
|
||
|
||
by_name = {
|
||
"diagnosis": pools.diagnosis,
|
||
"validation": pools.validation,
|
||
"test": pools.test,
|
||
}
|
||
loc: dict[str, set[str]] = {}
|
||
for name, pool in by_name.items():
|
||
for q in pool:
|
||
loc.setdefault(q.pair_id, set()).add(name)
|
||
|
||
split = {pid: names for pid, names in loc.items() if len(names) > 1}
|
||
assert not split, f"pair 被劈到多个池: {split}"
|
||
|
||
# 每个被选中的 pair 必须两题齐全(同池内成对),不得只落单题
|
||
per_pool_pair_count: dict[tuple[str, str], int] = {}
|
||
for name, pool in by_name.items():
|
||
for q in pool:
|
||
key = (name, q.pair_id)
|
||
per_pool_pair_count[key] = per_pool_pair_count.get(key, 0) + 1
|
||
assert all(c == 2 for c in per_pool_pair_count.values()), (
|
||
f"存在池内落单的 pair 成员: {per_pool_pair_count}"
|
||
)
|
||
|
||
|
||
def test_pair_never_split_across_train_val() -> None:
|
||
"""PerCategoryPoolStrategy 路径:孪生对不得被 train/val 划分劈开。"""
|
||
qs = [q for pid in [f"p{i:02d}" for i in range(15)] for q in _pair(pid)]
|
||
# 让部分 pair 正确、部分错误,触发分层划分(非退化随机)
|
||
correctness: dict[str, bool] = {}
|
||
for i, pid in enumerate(f"p{i:02d}" for i in range(15)):
|
||
val = i < 9
|
||
correctness[f"{pid}_o"] = val
|
||
correctness[f"{pid}_m"] = val
|
||
|
||
config = PoolConfig(
|
||
task_types=None,
|
||
seed=42,
|
||
baseline_run_id="b",
|
||
diag_size=0,
|
||
diag_correct_ratio=0.0,
|
||
val_size=0,
|
||
val_correct_ratio=0.0,
|
||
test_size=0,
|
||
eval_min_per_class=0,
|
||
train_ratio=2 / 3,
|
||
test_questions_dir=None,
|
||
)
|
||
pools = PerCategoryPoolStrategy().build(qs, correctness, config)
|
||
|
||
loc: dict[str, set[str]] = {}
|
||
for name, pool in (("diagnosis", pools.diagnosis), ("validation", pools.validation)):
|
||
for q in pool:
|
||
loc.setdefault(q.pair_id, set()).add(name)
|
||
|
||
split = {pid: names for pid, names in loc.items() if len(names) > 1}
|
||
assert not split, f"pair 被 train/val 划分劈开: {split}"
|