Files
Video-Tree-TRM5/tests/unit/test_pools_pair_atomic.py
iomgaa ddb9a44f75 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 原子性回归。
2026-07-15 06:09:14 -04:00

117 lines
3.9 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""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}"