Files
giant/scripts/check_migration_v02_v03.py
T
lars 9ce55e5013
CI / Format (ruff format) (push) Failing after 25s
CI / Lint (ruff check) (push) Successful in 27s
CI / Sync project version with tag (push) Has been skipped
CI / Type check (ty) (push) Failing after 24s
CI / Tests (push) Has been skipped
v0.3.0 step 2: network.py refactor to composable stage models
Decomposes the ten permutation classes in giant/model/network.py into
the reusable parts from docs/v0.3.0-design.md §5: ConditionEncoder (now
independently configurable per particle/material axis), ContextAdapter,
Trunk/MonolithicTrunk/RoutedTrunk/ExpertTrunk, and the stage classes
Stage1Model/Stage2OneShot/CriticModel (Stage2Autoregressive stubbed,
raises NotImplementedError until step 4/5). build_models/build_critics
now return a dict keyed by stage and accept the new nested config shape,
with routed WGAN reachable for the first time (the old --mode wgan
--router rejection is gone) and stage2_model.router.tie_to_stage1
sharing a literal Router instance.

A v0.2 checkpoint's flat model_config auto-migrates via
_migrate_legacy_model_config + migrate_legacy_state_dict, preserving the
n_sec_head's attachment to Stage1Model (legacy_owner="stage1", design
doc §4.1). tests/test_migration_v02_v03.py proves this bit-identical
against a frozen v0.2 snapshot (tests/legacy/network_v02_snapshot.py)
for both flow and wgan, both conditioning modes.
scripts/check_migration_v02_v03.py is the real-checkpoint counterpart
for a portal machine with /ceph access.

giant/model/schedule.py's flow-matching/DDPM loss helpers are updated
to the new model-call convention (t as a keyword). giant/sample.py,
giant/rollout.py, and giant/validate.py are not yet updated (deferred
to design doc step 6) — their exercising tests are marked xfail with
that reasoning rather than silently broken.

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
2026-08-06 10:55:29 +02:00

175 lines
6.9 KiB
Python

"""Portal-machine follow-up for v0.3.0 step 2 (docs/v0.3.0-design.md §4.3):
diff a real v0.2 checkpoint's outputs against the new `build_models` on the
same input batch.
`tests/test_migration_v02_v03.py` already proves this bit-identical with
synthetic random weights, but that test can't run where it matters (no
`/ceph` on local dev machines — see CLAUDE.md's Compute environment
section). This script is the real-checkpoint counterpart: run it on a portal
machine against an actual trained checkpoint before merging
`v0.3.0-stage2-autoregressive` to `master`.
Usage (from the repo root, on a portal machine):
uv run python scripts/check_migration_v02_v03.py /ceph/lbogner/.../best.pt
uv run python scripts/check_migration_v02_v03.py /ceph/lbogner/.../best.pt --ema
uv run python scripts/check_migration_v02_v03.py /ceph/lbogner/.../best.pt --batch 32 --seed 1
Run it once against a flow (or ddpm) checkpoint and once against a wgan
checkpoint (design doc §4.3's "one flow checkpoint and one WGAN checkpoint").
A routed checkpoint (`model_config["router"]["enabled"]`) is only checked for
successful construction — `giant.model.network.migrate_legacy_state_dict`
doesn't yet remap routed (Expert-per-router) state dicts, so the
bit-identical assertion is skipped with a clear warning in that case (see the
function's own docstring for why).
"""
import argparse
import sys
from pathlib import Path
import torch
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
from giant.constants import COND_DIM, SEC_SLOT_DIM, X_DIM # noqa: E402
from giant.model import network as net # noqa: E402
from tests.legacy import network_v02_snapshot as legacy # noqa: E402
def _random_batch(model_config: dict, batch: int, seed: int):
g = torch.Generator().manual_seed(seed)
pdg_vocab = model_config["pdg_vocab"]
mat_vocab = model_config["mat_vocab"]
k_max = model_config.get("k_max", 15)
noise_dim = model_config.get("noise_dim", 64)
cond_cont = torch.randn(batch, COND_DIM, generator=g)
cond_cat = torch.stack(
[
torch.randint(0, pdg_vocab, (batch,), generator=g),
torch.randint(0, mat_vocab, (batch,), generator=g),
],
dim=1,
)
x1 = torch.randn(batch, X_DIM, generator=g)
x2 = torch.randn(batch, k_max * SEC_SLOT_DIM, generator=g)
t = torch.rand(batch, generator=g)
z1 = torch.randn(batch, noise_dim, generator=g)
z2 = torch.randn(batch, noise_dim, generator=g)
return cond_cont, cond_cat, x1, x2, t, z1, z2
def _max_abs_diff(a: torch.Tensor, b: torch.Tensor) -> float:
return (a - b).abs().max().item()
def main() -> int:
p = argparse.ArgumentParser(description=__doc__)
p.add_argument("checkpoint", type=Path, help="Path to a v0.2 best.pt/last.pt")
p.add_argument(
"--ema",
action="store_true",
help="Use the checkpoint's EMA weights (model_ema/sec_decoder_ema) — "
"what predict/rollout actually sample from — instead of raw weights.",
)
p.add_argument("--batch", type=int, default=16)
p.add_argument("--seed", type=int, default=0)
args = p.parse_args()
ckpt = torch.load(args.checkpoint, map_location="cpu", weights_only=False)
if "model_config" not in ckpt:
print(f"FAIL: {args.checkpoint} has no 'model_config' key — can't migrate it")
return 1
model_config = ckpt["model_config"]
mode = model_config.get("mode", "flow")
routed = bool((model_config.get("router") or {}).get("enabled"))
print(f"checkpoint: {args.checkpoint}")
print(
f" mode={mode!r} conditioning={model_config.get('conditioning')!r} "
f"routed={routed} ema={args.ema}"
)
stage1_key = "model_ema" if args.ema and "model_ema" in ckpt else "model"
stage2_key = (
"sec_decoder_ema" if args.ema and "sec_decoder_ema" in ckpt else "sec_decoder"
)
if args.ema and stage1_key == "model":
print(
" warning: --ema requested but no model_ema in checkpoint, using raw weights"
)
# --- old side: the frozen v0.2 snapshot, loaded with the checkpoint's own weights ---
old_stage1, old_stage2 = legacy.build_models(model_config)
old_stage1.load_state_dict(ckpt[stage1_key])
old_stage2.load_state_dict(ckpt[stage2_key])
old_stage1.eval()
old_stage2.eval()
# --- new side: migrated config + remapped state dict, through the new build_models ---
new_models = net.build_models(model_config)
new_stage1, new_stage2 = new_models["stage1"], new_models["stage2"]
assert new_stage1 is not None and new_stage2 is not None
if routed:
print(
" routed checkpoint: migrate_legacy_state_dict only handles the "
"monolithic trunk shape — verifying construction only, skipping "
"the bit-identical weight/output comparison. See "
"docs/v0.3.0-design.md §2.4's scope note."
)
print("PASS (construction only, routed checkpoint)")
return 0
remapped1, remapped2 = net.migrate_legacy_state_dict(
ckpt[stage1_key], ckpt[stage2_key]
)
missing1, unexpected1 = new_stage1.load_state_dict(remapped1, strict=True)
missing2, unexpected2 = new_stage2.load_state_dict(remapped2, strict=True)
if missing1 or unexpected1 or missing2 or unexpected2:
print("FAIL: state dict mismatch after remap")
print(f" stage1 missing={missing1} unexpected={unexpected1}")
print(f" stage2 missing={missing2} unexpected={unexpected2}")
return 1
new_stage1.eval()
new_stage2.eval()
cond_cont, cond_cat, x1, x2, t, z1, z2 = _random_batch(
model_config, args.batch, args.seed
)
ok = True
with torch.no_grad():
if mode == "wgan":
old_out1 = old_stage1(z1, cond_cont, cond_cat)
new_out1 = new_stage1(z1, cond_cont, cond_cat)
else:
old_out1 = old_stage1(x1, t, cond_cont, cond_cat)
new_out1 = new_stage1(x1, cond_cont, cond_cat, t=t)
old_n_sec = old_stage1.predict_n_sec(cond_cont, cond_cat)
new_n_sec = new_stage1.predict_n_sec(cond_cont, cond_cat)
if mode == "wgan":
old_out2 = old_stage2(z2, cond_cont, cond_cat, old_out1)
new_out2 = new_stage2(z2, cond_cont, cond_cat, new_out1)
else:
old_out2 = old_stage2(x2, t, cond_cont, cond_cat, old_out1)
new_out2 = new_stage2(x2, cond_cont, cond_cat, new_out1, t=t)
for label, old_out, new_out in [
("stage1 output", old_out1, new_out1),
("n_sec logits", old_n_sec, new_n_sec),
("stage2 output", old_out2, new_out2),
]:
identical = torch.equal(old_out, new_out)
diff = _max_abs_diff(old_out, new_out)
status = "OK" if identical else "MISMATCH"
print(f" {label}: {status} (max abs diff = {diff:.3e})")
ok = ok and identical
print("PASS" if ok else "FAIL")
return 0 if ok else 1
if __name__ == "__main__":
raise SystemExit(main())