Files
Video-Tree-TRM5/tests/unit/test_loader_unit_sampling.py
T
iomgaa d6a3107e4e feat(question_gen): loader 按 unit 分层采样 + load_benchmark 读回 pair 字段
stratified_sample 先 build_units 聚合,以 QuestionUnit 为采样原子做
分层/去重/补足/rng.sample,返回前 flatten_units 展开为逐题列表;
size/correct_ratio/min_per_class 均按 unit 计数,单元正确性走成员 AND,
孪生对两题永不被劈开。纯 single 输入下 build_units 1:1 折叠、顺序不变,
rng 消耗与旧逐题实现字节级一致(新增回归测试守护)。

_backfill_per_class candidates 改按 unit 枚举去重;build_units/flatten_units
函数内延迟导入以规避 question_gen<->harness 循环依赖(沿用 adversarial_filter)。

load_benchmark 反序列化补 pair_id/question_role/flip_axis/unit_id 四字段,
用 .get 兼容旧 JSON(缺失退化为 single,unit_id 由 __post_init__ 回填)。

pools._sample_excluding 随之改为透传 flatten_units(candidates) 给已单元化的
stratified_sample(不再用 lone pair-original 代表),行为对 single-only 保持等价。
2026-07-15 06:27:41 -04:00

298 lines
11 KiB
Python
Raw 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.
"""loader.stratified_sample 按 unit 采样 + load_benchmark 读回 pair 字段的单元测试。
覆盖 question-gen v3 Phase 1 Task 4 的四项契约:
(a) pair 采样同进同出——同一 pair_id 两题要么都被选、要么都不被选;
(b) size / correct_ratio 按 **unit 计数**pair 计 1 个 unitunit 正确性走双向 AND);
(c) min_per_class 补足路径不拆 pair
(d) load_benchmark 反序列化读回 pair_id/question_role/flip_axis/unit_id
且旧 JSON 缺这些字段时按 single 默认兜底、不崩。
另加一条纯 single 输入的字节级回归守护:单题场景采样序列必须与旧逐题行为一致。
"""
from __future__ import annotations
import json
import random
from pathlib import Path
import pytest
from app.harness.question_units import build_units
from app.question_gen.loader import load_benchmark, stratified_sample
from core.types import GeneratedQuestion
def _single(qid: str, task_type: str = "Single") -> GeneratedQuestion:
"""构造一道 single 题(无 pair 归属)。"""
return GeneratedQuestion(
question_id=qid,
video_id="v",
task_type=task_type,
question="?",
options=("A", "B", "C", "D"),
answer="A",
source_nodes=(),
difficulty="medium",
)
def _pair(pid: str, task_type: str = "AR") -> list[GeneratedQuestion]:
"""构造一个合法孪生对(original + mirror),共享 pair_id / unit_id / flip_axis。"""
base = dict(
video_id="v",
task_type=task_type,
question="?",
options=("A", "B", "C", "D"),
answer="A",
source_nodes=(),
difficulty="hard",
pair_id=pid,
unit_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 _pair_roles_by_id(questions: list[GeneratedQuestion]) -> dict[str, set[str]]:
"""将采样结果按 pair_id 聚合出现的角色集合(仅统计配对题)。"""
loc: dict[str, set[str]] = {}
for q in questions:
if q.pair_id:
loc.setdefault(q.pair_id, set()).add(q.question_role)
return loc
class TestPairAtomicSampling:
def test_pair_never_split_in_natural_sample(self) -> None:
"""自然分布采样:命中的 pair 必两题齐全,size 按 unit 计数。"""
qs = [q for i in range(8) for q in _pair(f"p{i}")]
result = stratified_sample(
questions=qs,
correctness={},
size=4,
correct_ratio=None,
task_types=None,
seed=7,
min_per_class=None,
)
loc = _pair_roles_by_id(result)
assert len(loc) == 4, "size=4 应命中 4 个 pair 单元"
for pid, roles in loc.items():
assert roles == {"pair_original", "pair_mirror"}, f"pair {pid} 被拆: {roles}"
assert len(result) == 8, "4 个 pair 单元展开应为 8 道题"
def test_size_counts_units_not_questions(self) -> None:
"""混合 single + pair 时 size 仍按 unit 计数(pair 计 1)。"""
qs = [_single(f"s{i}") for i in range(3)]
qs += [q for i in range(3) for q in _pair(f"p{i}")]
result = stratified_sample(
questions=qs,
correctness={},
size=6, # 全部 6 个单元(3 single + 3 pair
correct_ratio=None,
task_types=None,
seed=1,
min_per_class=None,
)
singles = [q for q in result if not q.pair_id]
loc = _pair_roles_by_id(result)
assert len(singles) == 3
assert len(loc) == 3
for roles in loc.values():
assert roles == {"pair_original", "pair_mirror"}
assert len(result) == 3 + 3 * 2
def test_ratio_counts_units(self) -> None:
"""correct_ratio 按 unit 计数:对/错单元按比例各取整数个。"""
correct_pairs = [_pair(f"c{i}") for i in range(4)]
wrong_pairs = [_pair(f"w{i}") for i in range(4)]
qs = [q for p in correct_pairs + wrong_pairs for q in p]
correctness: dict[str, bool] = {}
for p in correct_pairs:
for q in p:
correctness[q.question_id] = True
result = stratified_sample(
questions=qs,
correctness=correctness,
size=4,
correct_ratio=0.5,
task_types=None,
seed=3,
min_per_class=None,
)
loc = _pair_roles_by_id(result)
assert len(loc) == 4, "4 个 unit"
for roles in loc.values():
assert roles == {"pair_original", "pair_mirror"}
correct_units = sum(
1 for pid in loc if all(correctness.get(f"{pid}_{s}", False) for s in ("o", "m"))
)
assert correct_units == 2, "correct_ratio=0.5 * 4 unit = 2 个对单元"
def test_unit_correctness_uses_and_not_any(self) -> None:
"""unit 正确性走双向 AND:混合对(P 对 Q 错)不得计入对单元层。
构造 2 个全对 pair + 1 个混合 pairoriginal 对、mirror 错)。请求 3 个对单元,
若按 AND,对池仅 2 个 → 分层不足报错;若错误地按 any-member,对池 3 个 → 不报错。
以是否抛错来判别语义,规避随机命中导致的假阳性。
"""
good = [_pair("g0"), _pair("g1")]
mixed = _pair("mx")
qs = [q for p in good for q in p] + mixed
correctness: dict[str, bool] = {}
for p in good:
for q in p:
correctness[q.question_id] = True
correctness[mixed[0].question_id] = True
correctness[mixed[1].question_id] = False
with pytest.raises(ValueError, match="分层不足"):
stratified_sample(
questions=qs,
correctness=correctness,
size=3,
correct_ratio=1.0,
task_types=None,
seed=1,
min_per_class=None,
)
class TestBackfillPairAtomic:
def test_backfill_keeps_pairs_atomic(self) -> None:
"""min_per_class 补足稀疏 pair 题型时,补入的 pair 两题齐全不拆。"""
ar = [q for i in range(3) for q in _pair(f"a{i}", task_type="AR")]
main = [_single(f"m{i}", task_type="Main") for i in range(10)]
result = stratified_sample(
questions=ar + main,
correctness={},
size=2,
correct_ratio=None,
task_types=None,
seed=5,
min_per_class=2,
)
loc = _pair_roles_by_id(result)
assert len(loc) >= 2, "AR 至少被补足到 2 个 pair 单元"
for pid, roles in loc.items():
assert roles == {"pair_original", "pair_mirror"}, f"补足拆散了 pair {pid}: {roles}"
def test_backfill_no_duplicate_units(self) -> None:
"""补足不得重复选入同一 pair 单元。"""
ar = [q for i in range(4) for q in _pair(f"a{i}", task_type="AR")]
result = stratified_sample(
questions=ar,
correctness={},
size=1,
correct_ratio=None,
task_types=None,
seed=9,
min_per_class=3,
)
ids = [q.question_id for q in result]
assert len(ids) == len(set(ids)), "存在重复题目"
class TestLoadBenchmarkPairFields:
def test_reads_pair_fields(self, tmp_path: Path) -> None:
"""新格式 JSON 的 pair_id/question_role/flip_axis/unit_id 被正确读回。"""
data = [
{
"question_id": "pp_o",
"video_id": "v",
"task_type": "AR",
"question": "?",
"options": ["A", "B", "C", "D"],
"answer": "A",
"pair_id": "pp",
"question_role": "pair_original",
"flip_axis": "before_after",
"unit_id": "pp",
},
{
"question_id": "pp_m",
"video_id": "v",
"task_type": "AR",
"question": "?",
"options": ["A", "B", "C", "D"],
"answer": "A",
"pair_id": "pp",
"question_role": "pair_mirror",
"flip_axis": "before_after",
"unit_id": "pp",
},
]
(tmp_path / "vid.json").write_text(json.dumps(data), encoding="utf-8")
qs = load_benchmark(tmp_path)
q0 = next(q for q in qs if q.question_id == "pp_o")
assert q0.pair_id == "pp"
assert q0.question_role == "pair_original"
assert q0.flip_axis == "before_after"
assert q0.unit_id == "pp"
# 读回后应能被 build_units 重新聚合成 1 个 pair 单元
units = build_units(qs)
assert len(units) == 1
assert units[0].kind == "pair"
def test_legacy_json_defaults_to_single(self, tmp_path: Path) -> None:
"""旧 JSON 缺 pair 字段时按 single 兜底、unit_id 回填为 question_id,不崩。"""
data = [
{
"question_id": "x1",
"video_id": "v",
"task_type": "T",
"question": "?",
"options": ["A", "B", "C", "D"],
"answer": "A",
}
]
(tmp_path / "legacy.json").write_text(json.dumps(data), encoding="utf-8")
qs = load_benchmark(tmp_path)
assert qs[0].pair_id is None
assert qs[0].question_role == "single"
assert qs[0].flip_axis is None
assert qs[0].unit_id == "x1"
units = build_units(qs)
assert len(units) == 1
assert units[0].kind == "single"
class TestPureSingleByteIdentical:
def test_natural_matches_legacy_rng(self) -> None:
"""纯 single 输入的自然分布采样,与旧逐题 rng.sample 序列字节级一致。"""
qs = [_single(f"s{i}") for i in range(20)]
seed = 123
result = stratified_sample(
questions=qs,
correctness={},
size=8,
correct_ratio=None,
task_types=None,
seed=seed,
min_per_class=None,
)
reference = random.Random(seed).sample(qs, 8)
assert [q.question_id for q in result] == [q.question_id for q in reference]
def test_ratio_matches_legacy_rng(self) -> None:
"""纯 single 输入的比例分层采样,与旧逐题实现的抽样序列一致。"""
qs = [_single(f"s{i}") for i in range(20)]
correctness = {f"s{i}": i < 10 for i in range(20)}
seed = 77
result = stratified_sample(
questions=qs,
correctness=correctness,
size=10,
correct_ratio=0.6,
task_types=None,
seed=seed,
min_per_class=None,
)
rng = random.Random(seed)
correct = [q for q in qs if correctness.get(q.question_id, False)]
wrong = [q for q in qs if not correctness.get(q.question_id, False)]
reference = rng.sample(correct, 6) + rng.sample(wrong, 4)
assert [q.question_id for q in result] == [q.question_id for q in reference]