feat(pools): 三池切分以 unit 为原子,孪生对同池不被拆散

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 原子性回归。
This commit is contained in:
2026-07-15 06:09:14 -04:00
parent 5ef5f2b8b7
commit ddb9a44f75
2 changed files with 214 additions and 48 deletions
+116
View File
@@ -0,0 +1,116 @@
"""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_poolsGlobal 路径)三池切分后,任一 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}"