c412698cff
TC003 将 Path 导入移入 TYPE_CHECKING 块(仅注解使用), C408 将 base = dict(...) 改为字典字面量。行为不变。
301 lines
11 KiB
Python
301 lines
11 KiB
Python
"""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 typing import TYPE_CHECKING
|
||
|
||
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
|
||
|
||
if TYPE_CHECKING:
|
||
from pathlib import Path
|
||
|
||
|
||
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 = {
|
||
"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]
|