"""app/harness/batching.py 的 unit 粒度切分测试(Task 5)。 覆盖 pair 契约在 mini-batch 构建中的三条铁律: - (a) 同 pair_id 两题整锁进同一 batch(像小类整组不拆); - (b) pair 按 unit correctness(双向 AND)落 correct/error 桶,不因 P 对 Q 错被劈; - (c) 非 AR byte-identical:AR unit 折叠不改变非 AR 的 rng.sample/shuffle 抽样序列—— 纯非 AR 输入下 build_batches 结果与"引入 QuestionUnit 前"的旧算法逐字节一致。 """ from __future__ import annotations import math import random from app.harness.batching import build_batches from core.types import GeneratedQuestion # --------------------------------------------------------------------------- # 辅助构造 # --------------------------------------------------------------------------- def _make_single( qid: str, task_type: str = "default", video_id: str = "v1", ) -> GeneratedQuestion: """构造最小 single GeneratedQuestion。""" return GeneratedQuestion( question_id=qid, video_id=video_id, task_type=task_type, question=f"question_{qid}", options=("A. a", "B. b", "C. c", "D. d"), answer="A", source_nodes=("n1",), difficulty="medium", question_role="single", ) def _make_pair( pair_id: str, task_type: str = "AR", video_id: str = "v1", ) -> tuple[GeneratedQuestion, GeneratedQuestion]: """构造一个孪生对(original + mirror),共享 pair_id / flip_axis。""" common = { "video_id": video_id, "task_type": task_type, "question": "?", "options": ("A. a", "B. b", "C. c", "D. d"), "answer": "A", "source_nodes": ("n1",), "difficulty": "medium", "pair_id": pair_id, "flip_axis": "before_after", } original = GeneratedQuestion( question_id=f"{pair_id}_o", question_role="pair_original", **common ) mirror = GeneratedQuestion(question_id=f"{pair_id}_m", question_role="pair_mirror", **common) return original, mirror def _batch_of(batches: list[list[GeneratedQuestion]], qid: str) -> int: """返回 qid 所在 batch 的下标;不存在返回 -1。""" for i, b in enumerate(batches): if any(q.question_id == qid for q in b): return i return -1 # --------------------------------------------------------------------------- # (a) 同 pair_id 两题落同一 batch(整锁不拆) # --------------------------------------------------------------------------- class TestPairStaysInSameBatch: """孪生对两题必须整锁进同一 batch,无论散布在多少 single 之间。""" def test_single_pair_same_batch(self) -> None: po, pm = _make_pair("pairA", task_type="AR") singles = [_make_single(f"s{i}", task_type="RETRIEVAL") for i in range(10)] items = [*singles, po, pm] correctness = {q.question_id: False for q in items} batches, count = build_batches( items, correctness, batch_size=4, min_class_per_batch=2, seed=3 ) assert count == 12 idx_o = _batch_of(batches, "pairA_o") idx_m = _batch_of(batches, "pairA_m") assert idx_o != -1 and idx_o == idx_m def test_many_pairs_each_intact(self) -> None: items: list[GeneratedQuestion] = [] for k in range(5): po, pm = _make_pair(f"p{k}", task_type="AR") items.extend([po, pm]) items.extend(_make_single(f"s{i}", task_type="SPATIAL") for i in range(6)) correctness = {q.question_id: False for q in items} batches, _ = build_batches(items, correctness, batch_size=6, min_class_per_batch=2, seed=11) for k in range(5): idx_o = _batch_of(batches, f"p{k}_o") idx_m = _batch_of(batches, f"p{k}_m") assert idx_o != -1 and idx_o == idx_m, f"pair p{k} 被拆到不同 batch" # 每个 batch 容量不超限(pair 占 2) for b in batches: assert len(b) <= 6 # --------------------------------------------------------------------------- # (b) pair 按 unit correctness(双向 AND)分桶 # --------------------------------------------------------------------------- class TestPairBucketedByUnitCorrectness: """P 对 Q 错的 pair 是 error 单元,不因单题分歧被劈到两个桶。""" def test_p_correct_q_wrong_pair_is_error_unit(self) -> None: po, pm = _make_pair("pairErr", task_type="AR") # 混一个 single 错题避免空池边界 s0 = _make_single("s0", task_type="AR") items = [po, pm, s0] # original 对、mirror 错 → 单元级 AND = 错 → 应作为 error 单元整体进 batch correctness = {"pairErr_o": True, "pairErr_m": False, "s0": False} batches, count = build_batches( items, correctness, batch_size=6, min_class_per_batch=2, seed=0, correct_ratio=0.0 ) idx_o = _batch_of(batches, "pairErr_o") idx_m = _batch_of(batches, "pairErr_m") # 两题都在(未被"P 对"劈掉)且同 batch assert idx_o != -1 and idx_o == idx_m # 纯错题模式下 pair 单元整体被选入 assert count == 3 def test_both_correct_pair_excluded_in_pure_error_mode(self) -> None: po, pm = _make_pair("pairOk", task_type="AR") s_err = _make_single("s_err", task_type="AR") items = [po, pm, s_err] correctness = {"pairOk_o": True, "pairOk_m": True, "s_err": False} batches, count = build_batches( items, correctness, batch_size=6, min_class_per_batch=2, seed=0, correct_ratio=0.0 ) # both-correct pair 是 correct 单元,纯错题模式下不入池 assert count == 1 assert _batch_of(batches, "pairOk_o") == -1 assert _batch_of(batches, "pairOk_m") == -1 assert _batch_of(batches, "s_err") != -1 def test_both_correct_pair_stays_intact_when_mixed(self) -> None: """correct_ratio>0 时 both-correct pair 可作为整体 correct 单元混入。""" errs = [_make_single(f"e{i}", task_type="AR") for i in range(4)] po, pm = _make_pair("pairOk", task_type="AR") items = [*errs, po, pm] correctness = {f"e{i}": False for i in range(4)} correctness["pairOk_o"] = True correctness["pairOk_m"] = True batches, _ = build_batches( items, correctness, batch_size=10, min_class_per_batch=2, seed=1, correct_ratio=0.5 ) idx_o = _batch_of(batches, "pairOk_o") idx_m = _batch_of(batches, "pairOk_m") # 若被混入,两题必须整体同 batch;若未被采样,两题都不在 if idx_o == -1: assert idx_m == -1 else: assert idx_o == idx_m # --------------------------------------------------------------------------- # (c) 非 AR byte-identical:与"引入 QuestionUnit 前"旧算法逐字节一致 # --------------------------------------------------------------------------- def _reference_select_mixed( items: list[GeneratedQuestion], correctness: dict[str, bool], correct_ratio: float, rng: random.Random, ) -> dict[str, list[GeneratedQuestion]]: """旧版 _select_mixed_by_task_type 的忠实副本(逐题、单一 rng)。""" errors_by_type: dict[str, list[GeneratedQuestion]] = {} correct_by_type: dict[str, list[GeneratedQuestion]] = {} for q in items: qid = q.question_id if correctness.get(qid) is False: errors_by_type.setdefault(q.task_type, []).append(q) elif correctness.get(qid, False): correct_by_type.setdefault(q.task_type, []).append(q) if correct_ratio <= 0: return errors_by_type grouped: dict[str, list[GeneratedQuestion]] = {} for task_type in sorted(errors_by_type): errs = errors_by_type[task_type] n_correct = round(len(errs) * correct_ratio / (1 - correct_ratio)) available = correct_by_type.get(task_type, []) sampled = ( list(available) if len(available) <= n_correct else rng.sample(available, n_correct) ) grouped[task_type] = errs + sampled return grouped def _reference_build_batches( items: list[GeneratedQuestion], correctness: dict[str, bool], batch_size: int, min_class_per_batch: int, seed: int, correct_ratio: float = 0.0, ) -> tuple[list[list[GeneratedQuestion]], int]: """引入 QuestionUnit 前的旧版 build_batches 的忠实副本(逐题、单一 rng)。""" rng = random.Random(seed) grouped = _reference_select_mixed(items, correctness, correct_ratio, rng) total = sum(len(g) for g in grouped.values()) if total == 0: return [], 0 nb = max(1, math.ceil(total / batch_size)) batches: list[list[GeneratedQuestion]] = [[] for _ in range(nb)] small = {t: g for t, g in grouped.items() if len(g) <= min_class_per_batch} large = {t: g for t, g in grouped.items() if len(g) > min_class_per_batch} for t in sorted(small, key=lambda t: (-len(small[t]), t)): group = small[t] placed = False for b in batches: if len(b) + len(group) <= batch_size: b.extend(group) placed = True break if not placed: batches.append(list(group)) nb_live = len(batches) pointer = 0 for task_type in sorted(large): group = list(large[task_type]) rng.shuffle(group) for q in group: for offset in range(nb_live): idx = (pointer + offset) % nb_live if len(batches[idx]) < batch_size: batches[idx].append(q) pointer = (idx + 1) % nb_live break result = [b for b in batches if b] return result, sum(len(b) for b in result) class TestNonARByteIdentical: """纯 single 输入下 build_batches 与旧逐题算法逐字节一致(黄金对照)。""" def _assert_identical( self, items: list[GeneratedQuestion], correctness: dict[str, bool], batch_size: int, min_class_per_batch: int, seed: int, correct_ratio: float, ) -> None: got_batches, got_count = build_batches( items, correctness, batch_size, min_class_per_batch, seed, correct_ratio ) ref_batches, ref_count = _reference_build_batches( items, correctness, batch_size, min_class_per_batch, seed, correct_ratio ) got_ids = [[q.question_id for q in b] for b in got_batches] ref_ids = [[q.question_id for q in b] for b in ref_batches] assert got_ids == ref_ids assert got_count == ref_count def test_pure_errors_identical(self) -> None: items = [_make_single(f"q{i}", task_type=f"type_{i % 3}") for i in range(20)] correctness = {f"q{i}": False for i in range(20)} self._assert_identical(items, correctness, 5, 2, 42, 0.0) def test_mixed_ratio_identical(self) -> None: items = [_make_single(f"q{i}", task_type=f"type_{i % 4}") for i in range(40)] correctness = {f"q{i}": (i % 3 == 0) for i in range(40)} self._assert_identical(items, correctness, 8, 3, 7, 0.5) def test_large_round_robin_identical(self) -> None: items = [_make_single(f"q{i}", task_type="big") for i in range(30)] correctness = {f"q{i}": False for i in range(30)} self._assert_identical(items, correctness, 6, 2, 99, 0.0) def test_ar_folding_does_not_shift_nonar_draws(self) -> None: """加入 AR pair 不改变非 AR single 的抽样序列(draw 流独立)。""" singles = [_make_single(f"q{i}", task_type=f"type_{i % 3}") for i in range(24)] s_correctness = {f"q{i}": (i % 4 == 0) for i in range(24)} base, _ = build_batches(singles, s_correctness, 6, 2, 5, 0.5) base_ids = {q.question_id for b in base for q in b} # 追加若干 AR pair(不同 task_type),非 AR single 的入选集合应不变 items = list(singles) correctness = dict(s_correctness) for k in range(3): po, pm = _make_pair(f"p{k}", task_type="AR") items.extend([po, pm]) correctness[po.question_id] = False correctness[pm.question_id] = False with_ar, _ = build_batches(items, correctness, 6, 2, 5, 0.5) with_ar_single_ids = {q.question_id for b in with_ar for q in b if q.pair_id is None} assert with_ar_single_ids == base_ids