"""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 个 unit,unit 正确性走双向 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 个混合 pair(original 对、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]