From bd1f7a22a2db90353d29faa836039ca6b4d9632b Mon Sep 17 00:00:00 2001 From: iomgaa Date: Wed, 15 Jul 2026 08:05:15 -0400 Subject: [PATCH] =?UTF-8?q?feat(harness):=20pools.json=20=E5=BA=8F?= =?UTF-8?q?=E5=88=97=E5=8C=96=20pair=20=E5=9B=9B=E5=AD=97=E6=AE=B5?= =?UTF-8?q?=E9=98=B2=E5=AD=A4=E5=84=BF=20single?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit _q_to_dict 写出 pair_id/question_role/flip_axis/unit_id,_dict_to_q 用 .get 兼容旧 workspace 的 pools.json 读回并回填(unit_id 缺省交 __post_init__)。 pools.json 是训练主回路读回题目处,此前漏写会让孪生对解冻后退化成孤儿 single、配对指标失真。categories 块沿用 per-qid 记录,Task 3 的 unit 原子 切分已保证两 pair 成员同池同 key,序列化不破坏该原子性。 --- app/harness/pools.py | 17 +++ tests/unit/test_pools_serialization.py | 172 +++++++++++++++++++++++++ 2 files changed, 189 insertions(+) create mode 100644 tests/unit/test_pools_serialization.py diff --git a/app/harness/pools.py b/app/harness/pools.py index 58859a1..32e3fb6 100644 --- a/app/harness/pools.py +++ b/app/harness/pools.py @@ -242,6 +242,10 @@ def _q_to_dict(q: GeneratedQuestion) -> dict: 返回: 包含全部字段的字典(options/source_nodes 从 tuple 转为 list)。 + + 关键实现细节: + pair 四字段(pair_id/question_role/flip_axis/unit_id)必须写出——pools.json + 是训练主回路读回题目的地方,漏写会让孪生对解冻后退化成孤儿 single。 """ return { "question_id": q.question_id, @@ -255,6 +259,10 @@ def _q_to_dict(q: GeneratedQuestion) -> dict: "family": q.family, "skill_target": q.skill_target, "difficulty_steps": q.difficulty_steps, + "pair_id": q.pair_id, + "question_role": q.question_role, + "flip_axis": q.flip_axis, + "unit_id": q.unit_id, } @@ -266,6 +274,11 @@ def _dict_to_q(d: dict) -> GeneratedQuestion: 返回: 恢复的 GeneratedQuestion 实例(options/source_nodes 恢复为 tuple)。 + + 关键实现细节: + pair 四字段用 .get 兼容旧 workspace 的 pools.json(无这些字段不崩,默认退化 + 为 single)——断点续跑铁律。unit_id 缺省时传 "",交给 GeneratedQuestion + 的 __post_init__ 回填为 pair_id 或 question_id,避免孤儿 single。 """ return GeneratedQuestion( question_id=d["question_id"], @@ -279,6 +292,10 @@ def _dict_to_q(d: dict) -> GeneratedQuestion: family=d.get("family"), skill_target=d.get("skill_target"), difficulty_steps=d.get("difficulty_steps"), + pair_id=d.get("pair_id"), + question_role=d.get("question_role", "single"), + flip_axis=d.get("flip_axis"), + unit_id=d.get("unit_id", ""), ) diff --git a/tests/unit/test_pools_serialization.py b/tests/unit/test_pools_serialization.py new file mode 100644 index 0000000..1cd53ef --- /dev/null +++ b/tests/unit/test_pools_serialization.py @@ -0,0 +1,172 @@ +"""pools.json 冻结/解冻序列化 pair 字段回归测试。 + +pools.json 是训练主回路真正读回题目的地方——若 _q_to_dict/_dict_to_q 不序列化 +pair 四字段(pair_id/question_role/flip_axis/unit_id),孪生对一旦冻结再 load_pools +就退化成孤儿 single,配对指标(collapse 等)全部失真。本测试锁死: + +- pair 冻结进 pools.json 再 load_pools,四字段不丢; +- per_category 的 categories 块以 unit 原子记 train/val(两 pair 成员同 key); +- 旧 JSON(无这些字段)用 .get 兼容不崩、默认退化为 single。 +""" + +from __future__ import annotations + +import json +from typing import TYPE_CHECKING + +from app.harness.pools import ( + PerCategoryPoolStrategy, + Pools, + load_pools, + save_pools, +) +from core.types import GeneratedQuestion, PoolConfig + +if TYPE_CHECKING: + from pathlib import Path + + +def _pair(pid: str) -> tuple[GeneratedQuestion, GeneratedQuestion]: + """构造一个合法孪生对(original + mirror),共享 pair_id / flip_axis。 + + 参数: + pid: 该孪生对的共享标识。 + + 返回: + (original, mirror) 两条题目,unit_id 由 __post_init__ 回填为 pid。 + """ + 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", + } + original = GeneratedQuestion(question_id=f"{pid}_o", question_role="pair_original", **base) + mirror = GeneratedQuestion(question_id=f"{pid}_m", question_role="pair_mirror", **base) + return original, mirror + + +def test_pair_fields_survive_freeze_thaw(tmp_path: Path) -> None: + """pair 冻结进 pools.json 再 load_pools,四字段逐一不丢。""" + o, m = _pair("p00") + single = GeneratedQuestion( + question_id="s0", + video_id="v", + task_type="Action Reasoning", + question="?", + options=("A. a", "B. b", "C. c", "D. d"), + answer="B", + source_nodes=("n",), + difficulty="medium", + ) + pools = Pools( + diagnosis=[o, m], + validation=[single], + test=[], + baseline_run_id="run_001", + baseline_val_accuracy=0.0, + correctness={}, + ) + pools_path = tmp_path / "pools.json" + save_pools(pools, pools_path) + restored = load_pools(pools_path) + + ro, rm = restored.diagnosis + assert ro.pair_id == "p00" + assert rm.pair_id == "p00" + assert ro.question_role == "pair_original" + assert rm.question_role == "pair_mirror" + assert ro.flip_axis == "before_after" + assert rm.flip_axis == "before_after" + assert ro.unit_id == "p00" + assert rm.unit_id == "p00" + + # single 侧默认值也需正确往返 + rs = restored.validation[0] + assert rs.pair_id is None + assert rs.question_role == "single" + assert rs.flip_axis is None + assert rs.unit_id == "s0" + + +def test_per_category_categories_record_pair_atomically(tmp_path: Path) -> None: + """per_category 的 categories 块以 unit 原子记:两 pair 成员必落同一 train/val key。""" + questions: list[GeneratedQuestion] = [] + correctness: dict[str, bool] = {} + for i in range(15): + o, m = _pair(f"p{i:02d}") + questions.extend((o, m)) + # 部分对错混合以触发分层划分而非退化随机 + val = i < 9 + correctness[o.question_id] = val + correctness[m.question_id] = val + + config = PoolConfig( + task_types=("Action Reasoning",), + seed=42, + baseline_run_id="baseline_v2", + 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(questions, correctness, config) + pools_path = tmp_path / "pools.json" + save_pools(pools, pools_path, split_mode="per_category", config=config) + + data = json.loads(pools_path.read_text(encoding="utf-8")) + cat = data["categories"]["Action Reasoning"] + train_ids = set(cat["train"]) + val_ids = set(cat["val"]) + + # 每个 pair 的两个成员必须同在 train 或同在 val,绝不被劈开 + for i in range(15): + o_id, m_id = f"p{i:02d}_o", f"p{i:02d}_m" + in_train = {o_id, m_id} <= train_ids + in_val = {o_id, m_id} <= val_ids + assert in_train ^ in_val, f"pair p{i:02d} 被 categories 劈开: train={o_id in train_ids},{m_id in train_ids} val={o_id in val_ids},{m_id in val_ids}" + + +def test_old_json_without_pair_fields_defaults_single(tmp_path: Path) -> None: + """旧 JSON(无 pair 四字段)用 .get 兼容加载不崩,退化为 single。""" + old_format = { + "baseline_run_id": "run_old", + "baseline_val_accuracy": 0.5, + "correctness": {}, + "diagnosis": [ + { + "question_id": "q_old", + "video_id": "v", + "task_type": "Action Reasoning", + "question": "?", + "options": ["A. a", "B. b", "C. c", "D. d"], + "answer": "A", + "source_nodes": ["n"], + "difficulty": "medium", + "family": None, + "skill_target": None, + "difficulty_steps": None, + } + ], + "validation": [], + "test": [], + } + pools_path = tmp_path / "pools.json" + pools_path.write_text(json.dumps(old_format), encoding="utf-8") + + restored = load_pools(pools_path) + q = restored.diagnosis[0] + assert q.pair_id is None + assert q.question_role == "single" + assert q.flip_axis is None + # unit_id 缺省 → __post_init__ 回填为 question_id + assert q.unit_id == "q_old"