feat(harness): pools.json 序列化 pair 四字段防孤儿 single

_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,序列化不破坏该原子性。
This commit is contained in:
2026-07-15 08:05:15 -04:00
parent 273984674b
commit bd1f7a22a2
2 changed files with 189 additions and 0 deletions
+17
View File
@@ -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", ""),
)
+172
View File
@@ -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"