Files
iomgaa 8958eee11b refactor: remove block-sequential gate path and gate_block knob (algo #6)
config/train_videomme.yaml 同时收录待入库的实验配置变更(run_id v2 /
concurrency 32 / batch_size 40)。tests/integration/test_v3_contract_e2e.py
的 run_id 断言按 Task 5 显式契约同步修正(原断言依赖旧隐式实例注入)。
2026-07-17 04:40:30 -04:00

255 lines
9.7 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""断点续跑 checkpoint 的 unit_id 粒度契约集成测试(Task 10)。
覆盖三条铁律:
- checkpoint 存 **unit_id 序列**而非逐题 question_id(孪生对折叠为单个 unit_id);
- 续跑恢复以 **完整 unit** 展开(pair 两成员同进同出,绝不被劈开);
- 真实经 write_checkpoint → load_checkpoint 落盘往返后,unit_id 序列无损存活、
重建的 batch 与原 batch 逐字节一致(core 算法保真 #3 断点续跑)。
"""
from __future__ import annotations
from dataclasses import dataclass, field
from typing import TYPE_CHECKING
from app.harness.batching import build_batches
from app.harness.checkpoint import load_checkpoint, write_checkpoint
from app.harness.pools import Pools
from app.harness.question_units import build_units, validate_units
from app.harness.runner import _batch_from_ids, _batch_unit_ids
from core.types import GeneratedQuestion
if TYPE_CHECKING:
from pathlib import Path
# ---------------------------------------------------------------------------
# 辅助构造
# ---------------------------------------------------------------------------
def _make_single(qid: str, task_type: str = "RETRIEVAL", video_id: str = "v1") -> GeneratedQuestion:
"""构造最小 single 题目。"""
return GeneratedQuestion(
question_id=qid,
video_id=video_id,
task_type=task_type,
question=f"q_{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 _pools_with(diagnosis: list[GeneratedQuestion]) -> Pools:
"""用给定诊断池构造最小 Pools(其余池置空,_batch_from_ids 只读 diagnosis)。"""
return Pools(
diagnosis=diagnosis,
validation=[],
test=[],
baseline_run_id="baseline",
baseline_val_accuracy=0.0,
correctness={q.question_id: False for q in diagnosis},
)
# ---------------------------------------------------------------------------
# checkpoint 落盘所需的最小 _TrainState / RunConfig 替身
# ---------------------------------------------------------------------------
@dataclass
class _FakeState:
"""serialize_state 只读的可持久化字段(本测试用空累加包即可)。"""
correctness: dict[str, bool] = field(default_factory=dict)
eval_prev_acc: float = 0.0
eval_prev_run_id: str = "run-0"
baseline_skills_version: str = "v1"
baseline_prompts_version: str = "v1"
epochs_since_best_improved: int = 0
epoch_start_skills: str = "v1"
changed_task_types_this_epoch: set[str] = field(default_factory=set)
rejected_buffer: dict = field(default_factory=dict)
system_packs: list = field(default_factory=list)
tool_packs: list = field(default_factory=list)
probations: dict = field(default_factory=dict)
gate_cooldown: dict = field(default_factory=dict)
gate_epoch_observed: dict = field(default_factory=dict)
@dataclass(frozen=True)
class _FakeConfig:
"""compute_fingerprint 读取的结构性 + 决策性字段。"""
batch_size: int = 4
min_class_per_batch: int = 2
epochs: int = 3
diag_size: int = 30
val_size: int = 50
batch_correct_ratio: float = 0.0
trainable_min_units: int = 8
edit_budget_start: int = 6
edit_budget_end: int = 3
early_stop_patience: int = 3
use_slow_momentum: bool = True
skill_update_mode: str = "patch"
appendix_consolidate_threshold: int = 10
momentum_samples: int = 20
gate_e_confirm: float = 20.0
gate_e_provisional: float = 6.0
gate_w_net_min: int = 2
gate_delta_min: float = 0.02
gate_lambda_dir: float = -3.0
gate_e_rollback: float = 10.0
gate_n_max: int = 40
gate_p_low: float = 0.1
gate_p_high: float = 0.9
gate_probe_quota: float = 0.2
gate_gamma_decay: float = 0.9
gate_cooldown_steps: int = 2
gate_guard_err: float = 0.3
# ---------------------------------------------------------------------------
# checkpoint 存 unit_id 序列(孪生对折叠为单个 unit_id)
# ---------------------------------------------------------------------------
class TestCheckpointStoresUnitIds:
"""一个含孪生对的 batchcheckpoint 序列应折叠为 unit_idpair 只记一次)。"""
def test_pair_folds_to_single_unit_id(self) -> None:
po, pm = _make_pair("pA")
s1, s2 = _make_single("s1"), _make_single("s2")
batch = [s1, po, pm, s2] # pair 两成员相邻
unit_ids = _batch_unit_ids(batch)
# 4 题折叠为 3 个 units1 / pA(pair) / s2
assert unit_ids == ["s1", "pA", "s2"]
# 逐题 question_id 序列(旧行为)会含 4 项且把 pair 拆成两条
assert [q.question_id for q in batch] == ["s1", "pA_o", "pA_m", "s2"]
def test_pure_single_unit_ids_equal_question_ids(self) -> None:
"""纯 single batch 的 unit_id 序列与逐题 question_id 序列逐字节一致。"""
batch = [_make_single(f"s{i}") for i in range(5)]
assert _batch_unit_ids(batch) == [q.question_id for q in batch]
# ---------------------------------------------------------------------------
# 恢复以完整 unit 展开:pair 不拆
# ---------------------------------------------------------------------------
class TestResumeExpandsFullUnit:
"""按 unit_id 序列恢复时,孪生对整体展开、绝不被劈开。"""
def test_pair_not_split_on_resume(self) -> None:
po, pm = _make_pair("pA")
s1, s2 = _make_single("s1"), _make_single("s2")
batch = [s1, po, pm, s2]
# 诊断池顺序故意与 batch 不同,验证恢复以池折叠的 unit 为准
pools = _pools_with([s2, pm, s1, po])
unit_ids = _batch_unit_ids(batch)
restored = _batch_from_ids(pools, unit_ids)
# 逐字节一致(顺序保留)
assert [q.question_id for q in restored] == ["s1", "pA_o", "pA_m", "s2"]
# 恢复后可无损重建为完整孪生对单元
units = validate_units(build_units(restored))
pair_units = [u for u in units if u.kind == "pair"]
assert len(pair_units) == 1
assert pair_units[0].unit_id == "pA"
assert pair_units[0].size == 2
def test_roundtrip_preserves_order_mixed(self) -> None:
"""混格 batch 经 unit_id 折叠 → 展开 round-trip 逐字节还原原顺序。"""
po, pm = _make_pair("pA")
qo, qm = _make_pair("pB")
singles = [_make_single(f"s{i}") for i in range(3)]
batch = [singles[0], po, pm, singles[1], qo, qm, singles[2]]
pools = _pools_with(list(batch))
restored = _batch_from_ids(pools, _batch_unit_ids(batch))
assert [q.question_id for q in restored] == [q.question_id for q in batch]
def test_missing_unit_id_raises(self) -> None:
"""checkpoint 引用了诊断池不存在的 unit_id 时硬失败(防静默兜底)。"""
pools = _pools_with([_make_single("s1")])
try:
_batch_from_ids(pools, ["s1", "ghost"])
except KeyError:
return
raise AssertionError("引用缺失 unit_id 应触发 KeyError")
# ---------------------------------------------------------------------------
# 真实落盘往返:write_checkpoint → load_checkpoint → 重建 batch
# ---------------------------------------------------------------------------
class TestCheckpointPersistenceRoundtrip:
"""unit_id 序列经真实 checkpoint.json 原子写/读回后无损、重建逐字节一致。"""
def test_persisted_unit_ids_rebuild_batches(self, tmp_path: Path) -> None:
po, pm = _make_pair("pA")
singles = [_make_single(f"s{i}", task_type="RETRIEVAL") for i in range(6)]
items = [*singles, po, pm]
correctness = {q.question_id: False for q in items}
batches, _ = build_batches(items, correctness, batch_size=4, min_class_per_batch=2, seed=7)
pools = _pools_with(list(items))
# 存 unit_id 序列(train() 落盘路径)
epoch_batches = [_batch_unit_ids(b) for b in batches]
write_checkpoint(
tmp_path,
state=_FakeState(correctness=correctness),
epoch=1,
step_completed=0,
phase="in_epoch",
global_step=1,
total_steps=len(batches),
version_snapshot={"skills": "skills/v1", "prompts": "prompts/v1"},
epoch_batches=epoch_batches,
config=_FakeConfig(),
)
ckpt = load_checkpoint(tmp_path)
assert ckpt is not None
rebuilt = [_batch_from_ids(pools, ids) for ids in ckpt["epoch_batches"]]
# 落盘往返后逐批逐字节还原
assert [[q.question_id for q in b] for b in rebuilt] == [
[q.question_id for q in b] for b in batches
]
# 含孪生对的批次恢复后 pair 仍成对
for b in rebuilt:
for u in validate_units(build_units(b)):
if u.kind == "pair":
assert u.size == 2