diff --git a/giant/analysis/router_gating.py b/giant/analysis/router_gating.py index a13c287..2a6a70f 100644 --- a/giant/analysis/router_gating.py +++ b/giant/analysis/router_gating.py @@ -66,27 +66,11 @@ class _RouterHandle: router_type: str -def _conditioning_axes(model_cfg: dict, default: str = "embedding") -> tuple[str, str]: - """(particle_conditioning, material_conditioning) for - `giant.data.transforms.build_cond_features` — from either a v0.2 - checkpoint's flat `model_config["conditioning"]` (one shared string, same - for both axes) or a new-format one (independent - `model_config["conditioning"]["particle"/"material"]["type"]` — the two - axes may differ). Mirrors - `giant.cli._conditioning_axes`.""" - raw = model_cfg.get("conditioning", default) - if isinstance(raw, dict): - return ( - raw.get("particle", {}).get("type", default), - raw.get("material", {}).get("type", default), - ) - return raw, raw - - def load_router(checkpoint: str | Path) -> _RouterHandle | None: """Load a checkpoint's Stage-1 router, or None if it isn't a MoE checkpoint.""" import torch + from giant.checkpoint_io import conditioning_axes from giant.data.transforms import Normalizer from giant.model.network import build_models @@ -110,7 +94,7 @@ def load_router(checkpoint: str | Path) -> _RouterHandle | None: if router is None: return None - particle_conditioning, material_conditioning = _conditioning_axes(model_cfg) + particle_conditioning, material_conditioning = conditioning_axes(model_cfg) return _RouterHandle( router=router, pdg_map={int(k): v for k, v in ckpt["pdg_map"].items()}, diff --git a/giant/checkpoint_io.py b/giant/checkpoint_io.py new file mode 100644 index 0000000..4ee92cc --- /dev/null +++ b/giant/checkpoint_io.py @@ -0,0 +1,208 @@ +"""Load a trained checkpoint into ready-to-run models (giant.cli's `predict`/`rollout`). + +Both commands need the same ~15 steps to go from a checkpoint path to two +`eval()`-mode models plus their normalizers/vocab maps: load the pickle, +validate it carries what current code expects, resolve which conditioning +mode each axis was trained with, restore the top-N vocab maps (if the +checkpoint used one-hot conditioning), rebuild the normalizers, construct the +model from `model_config`, and load the requested (raw or EMA) weights. This +used to be duplicated near-verbatim in both commands (issues.md Issue 5) — +`load_for_inference` is the single implementation. + +This module intentionally has no Typer dependency, so it can be unit-tested +directly and imported from non-CLI code (`giant.analysis.router_gating`, +lazily — see that module's docstring for why). Failures raise +`CheckpointCompatibilityError` with the same wording the CLI has always +shown; the CLI layer catches it and does the `typer.echo`/`Exit(1)`. +""" + +from __future__ import annotations + +from dataclasses import dataclass +from pathlib import Path + +import torch +from torch import nn + +from giant import config as gconfig +from giant.constants import K_MAX +from giant.data.loader import TopNMap +from giant.data.setup_cache import topnmap_from_json +from giant.data.transforms import Normalizer +from giant.model.network import build_models + + +class CheckpointCompatibilityError(Exception): + """Checkpoint is missing something `load_for_inference` needs.""" + + +def conditioning_axes(model_cfg: dict, default: str = "embedding") -> tuple[str, str]: + """(particle_conditioning, material_conditioning) for + `giant.data.transforms.build_cond_features`/`build_features` — from + either a v0.2 checkpoint's flat `model_config["conditioning"]` (one + shared string, same for both axes) or a new-format one (independent + `model_config["conditioning"]["particle"/"material"]["type"]` — the two + axes are configured independently and may differ).""" + raw = model_cfg.get("conditioning", default) + if isinstance(raw, dict): + return ( + raw.get("particle", {}).get("type", default), + raw.get("material", {}).get("type", default), + ) + return raw, raw + + +def stage_cfg(model_cfg: dict, stage: str) -> dict: + """`model_cfg[f"{stage}_model"]` for a new-format model_config, `{}` for + a v0.2 flat one (whose ddpm schedule always used `CosineSchedule`'s own + default `T=1000` — never a config key — and which never had + `particle_type` at all, so `{}` is the correct fallback for both + `ddpm_steps`/`particle_type_other_policy` below).""" + val = model_cfg.get(f"{stage}_model") + return val if isinstance(val, dict) else {} + + +def ddpm_steps(model_cfg: dict, stage: str) -> int: + return stage_cfg(model_cfg, stage).get("ddpm", {}).get("n_steps", 1000) + + +def particle_type_other_policy(model_cfg: dict) -> str: + return stage_cfg(model_cfg, "stage2").get("particle_type", {}).get("other_policy", "sample") + + +def load_pdg_topn_map(ckpt: dict) -> TopNMap | None: + """`ckpt["pdg_topn_map"]` as a `giant.data.loader.TopNMap`, or `None` if + this checkpoint's conditioning/particle_type never needed one (see + `giant.pipeline.run_setup_stage`, which only populates it when + `conditioning.particle.type` or `stage2_model.particle_type.target` is + `"onehot"`).""" + raw = ckpt.get("pdg_topn_map") + return topnmap_from_json(raw, axis="pdg") if raw is not None else None + + +def load_mat_topn_map(ckpt: dict) -> TopNMap | None: + """`ckpt["mat_topn_map"]` as a `giant.data.loader.TopNMap`, or `None` if + this checkpoint's `conditioning.material.type` was never `"onehot"` (see + `giant.pipeline.run_setup_stage`).""" + raw = ckpt.get("mat_topn_map") + return topnmap_from_json(raw, axis="material") if raw is not None else None + + +@dataclass(frozen=True) +class InferenceContext: + """Everything needed to run a trained checkpoint forward, resolved once.""" + + stage1: nn.Module | None + stage2: nn.Module | None + cond_norm: Normalizer + tgt_norm: Normalizer + sec_phys_norm: Normalizer + pdg_map: dict[int, int] + mat_map: dict[str, int] + pdg_topn_map: TopNMap | None + mat_topn_map: TopNMap | None + particle_conditioning: str + material_conditioning: str + k_max: int + stage1_ddpm_steps: int + stage2_ddpm_steps: int + other_policy: str + model_config: dict + epoch: int | None + best_val_loss: float | None + + +def load_for_inference( + checkpoint: Path, + device: torch.device, + command_name: str, + weights: str = "raw", + require_stage2: bool = True, +) -> InferenceContext: + """Load *checkpoint* and reconstruct everything `predict`/`rollout` need + to run it forward, on *device*, in `eval()` mode. + + *command_name* (e.g. `"predict"`/`"rollout"`) only feeds the "needs both" + error message below. *weights* is `"raw"` (the live training weights) or + `"ema"` (the EMA shadow copy, see `--ema-decay`). *require_stage2* + controls whether a checkpoint with an inactive stage 2 + (`stage2_model.active = false`) is an error (both current callers need + both stages) or an acceptable `stage2 = None` result — kept as a real + parameter since `stage{1,2}_model.active` is a real, if currently + stage1+stage2-only-in-practice, config option. + """ + ckpt = torch.load(checkpoint, map_location="cpu", weights_only=False) + for key in ("model_config", "sec_decoder"): + if key not in ckpt: + raise CheckpointCompatibilityError(f"checkpoint has no {key} — retrain with the current code") + + if "sec_phys" not in ckpt.get("normalizer", {}): + raise CheckpointCompatibilityError("checkpoint has no normalizer.sec_phys — retrain with the current code") + + gconfig.warn_if_checkpoint_config_mismatch(checkpoint) + + model_cfg = ckpt["model_config"] + particle_conditioning, material_conditioning = conditioning_axes(model_cfg) + pdg_topn_map = load_pdg_topn_map(ckpt) + mat_topn_map = load_mat_topn_map(ckpt) + if particle_conditioning == "onehot" and pdg_topn_map is None: + raise CheckpointCompatibilityError( + "checkpoint's conditioning.particle.type='onehot' but has no pdg_topn_map — retrain with the current code" + ) + if material_conditioning == "onehot" and mat_topn_map is None: + raise CheckpointCompatibilityError( + "checkpoint's conditioning.material.type='onehot' but has no mat_topn_map — retrain with the current code" + ) + other_policy = particle_type_other_policy(model_cfg) + stage1_ddpm_steps = ddpm_steps(model_cfg, "stage1") + stage2_ddpm_steps = ddpm_steps(model_cfg, "stage2") + k_max = stage_cfg(model_cfg, "stage2").get("k_max", K_MAX) + pdg_map = {int(k): v for k, v in ckpt["pdg_map"].items()} + mat_map = {str(k): v for k, v in ckpt["mat_map"].items()} + cond_norm = Normalizer.from_dict(ckpt["normalizer"]["cond"]) + tgt_norm = Normalizer.from_dict(ckpt["normalizer"]["target"]) + sec_phys_norm = Normalizer.from_dict(ckpt["normalizer"]["sec_phys"]) + + built = build_models(model_cfg) + stage1, stage2 = built["stage1"], built["stage2"] + if require_stage2 and (stage1 is None or stage2 is None): + raise CheckpointCompatibilityError( + f"checkpoint has an inactive stage1 or stage2 — {command_name} needs both (see stage{{1,2}}_model.active)" + ) + + if weights == "raw": + model_key, sec_key = "model", "sec_decoder" + else: + model_key, sec_key = "model_ema", "sec_decoder_ema" + if model_key not in ckpt or sec_key not in ckpt: + raise CheckpointCompatibilityError( + f"{checkpoint} has no EMA weights (trained before --ema-decay, " + "or with --ema-decay 0) — use --weights raw" + ) + if stage1 is not None: + stage1.load_state_dict(ckpt[model_key]) + stage1.to(device).eval() + if stage2 is not None: + stage2.load_state_dict(ckpt[sec_key]) + stage2.to(device).eval() + + return InferenceContext( + stage1=stage1, + stage2=stage2, + cond_norm=cond_norm, + tgt_norm=tgt_norm, + sec_phys_norm=sec_phys_norm, + pdg_map=pdg_map, + mat_map=mat_map, + pdg_topn_map=pdg_topn_map, + mat_topn_map=mat_topn_map, + particle_conditioning=particle_conditioning, + material_conditioning=material_conditioning, + k_max=k_max, + stage1_ddpm_steps=stage1_ddpm_steps, + stage2_ddpm_steps=stage2_ddpm_steps, + other_policy=other_policy, + model_config=model_cfg, + epoch=ckpt.get("epoch"), + best_val_loss=ckpt.get("best_val_loss"), + ) diff --git a/giant/cli.py b/giant/cli.py index 533034c..ea254bb 100644 --- a/giant/cli.py +++ b/giant/cli.py @@ -19,7 +19,6 @@ from tqdm import tqdm from giant import config as gconfig from giant.constants import ( - K_MAX, LOCAL_TARGET_NAMES, PREDICT_COORD_METADATA_KEY, PREDICT_SCHEMA_VERSION, @@ -39,11 +38,9 @@ from giant.data.transforms import ( inv_local_frame_rotation, inv_log_transform, reconstruct_post_pos, - Normalizer, ) -from giant.data.setup_cache import topnmap_from_json +from giant.checkpoint_io import CheckpointCompatibilityError, load_for_inference from giant.geometry import GeometryOracle -from giant.model.network import build_models from giant.pipeline import run_train_job from giant.rollout import ( L1DistCollector, @@ -71,58 +68,6 @@ def _router_total_experts(router_cfg: dict) -> int: return int(router_cfg.get("n_experts", 1)) -def _conditioning_axes(model_cfg: dict, default: str = "embedding") -> tuple[str, str]: - """(particle_conditioning, material_conditioning) for - `giant.data.transforms.build_cond_features`/`build_features` — from - either a v0.2 checkpoint's flat `model_config["conditioning"]` (one - shared string, same for both axes) or a new-format one (independent - `model_config["conditioning"]["particle"/"material"]["type"]` — the two - axes are configured independently and may differ).""" - raw = model_cfg.get("conditioning", default) - if isinstance(raw, dict): - return ( - raw.get("particle", {}).get("type", default), - raw.get("material", {}).get("type", default), - ) - return raw, raw - - -def _stage_cfg(model_cfg: dict, stage: str) -> dict: - """`model_cfg[f"{stage}_model"]` for a new-format model_config, `{}` for - a v0.2 flat one (whose ddpm schedule always used `CosineSchedule`'s own - default `T=1000` — never a config key — and which never had - `particle_type` at all, so `{}` is the correct fallback for both - `_ddpm_steps`/`_particle_type_other_policy` below).""" - val = model_cfg.get(f"{stage}_model") - return val if isinstance(val, dict) else {} - - -def _ddpm_steps(model_cfg: dict, stage: str) -> int: - return _stage_cfg(model_cfg, stage).get("ddpm", {}).get("n_steps", 1000) - - -def _particle_type_other_policy(model_cfg: dict) -> str: - return _stage_cfg(model_cfg, "stage2").get("particle_type", {}).get("other_policy", "sample") - - -def _load_pdg_topn_map(ckpt: dict): - """`ckpt["pdg_topn_map"]` as a `giant.data.loader.TopNMap`, or `None` if - this checkpoint's conditioning/particle_type never needed one (see - `giant.pipeline.run_setup_stage`, which only populates it when - `conditioning.particle.type` or `stage2_model.particle_type.target` is - `"onehot"`).""" - raw = ckpt.get("pdg_topn_map") - return topnmap_from_json(raw, axis="pdg") if raw is not None else None - - -def _load_mat_topn_map(ckpt: dict): - """`ckpt["mat_topn_map"]` as a `giant.data.loader.TopNMap`, or `None` if - this checkpoint's `conditioning.material.type` was never `"onehot"` (see - `giant.pipeline.run_setup_stage`).""" - raw = ckpt.get("mat_topn_map") - return topnmap_from_json(raw, axis="material") if raw is not None else None - - def _batch_size_estimate_dims(model_cfg: dict, training: bool, stage: str = "stage1") -> tuple[int, int]: """Pick the (hidden_dim, n_blocks) that dominate per-call activation memory. @@ -298,34 +243,6 @@ class Weights(str, Enum): ema = "ema" -def _load_model_weights( - model: torch.nn.Module, - sec_decoder: torch.nn.Module, - ckpt: dict, - weights: "Weights", - checkpoint_path: Path, -) -> None: - """Load either the raw or EMA state dicts from a training checkpoint. - - EMA weights (giant.training's shadow copy, see --ema-decay) only exist in - checkpoints written after that feature landed, so `ema` fails loudly - rather than silently falling back to raw weights a caller didn't ask for. - """ - if weights == Weights.raw: - model_key, sec_key = "model", "sec_decoder" - else: - model_key, sec_key = "model_ema", "sec_decoder_ema" - if model_key not in ckpt or sec_key not in ckpt: - typer.echo( - f"error: {checkpoint_path} has no EMA weights (trained before " - "--ema-decay, or with --ema-decay 0) — use --weights raw", - err=True, - ) - raise typer.Exit(1) - model.load_state_dict(ckpt[model_key]) - sec_decoder.load_state_dict(ckpt[sec_key]) - - @app.command() def train( data: Annotated[Path, typer.Argument(help="Parquet file or directory of parquet files")], @@ -992,37 +909,25 @@ def predict( typer.echo(f"device: {_device}") # --- Load checkpoint --- - ckpt = torch.load(checkpoint, map_location="cpu", weights_only=False) - if "model_config" not in ckpt: - typer.echo( - "error: checkpoint has no model_config — retrain with the current code", - err=True, - ) + try: + ctx = load_for_inference(checkpoint, _device, "predict", weights=weights.value) + except CheckpointCompatibilityError as exc: + typer.echo(f"error: {exc}", err=True) raise typer.Exit(1) + typer.echo(f"loaded checkpoint: {checkpoint} (weights: {weights.value})") - if "sec_decoder" not in ckpt: - typer.echo( - "error: checkpoint has no sec_decoder — retrain with the current code", - err=True, - ) - raise typer.Exit(1) - - if "sec_phys" not in ckpt.get("normalizer", {}): - typer.echo( - "error: checkpoint has no normalizer.sec_phys — retrain with the current code", - err=True, - ) - raise typer.Exit(1) - - model_cfg = ckpt["model_config"] - pdg_topn_map = _load_pdg_topn_map(ckpt) - mat_topn_map = _load_mat_topn_map(ckpt) - other_policy = _particle_type_other_policy(model_cfg) - stage1_ddpm_steps = _ddpm_steps(model_cfg, "stage1") - stage2_k_max = _stage_cfg(model_cfg, "stage2").get("k_max", K_MAX) + assert ctx.stage1 is not None and ctx.stage2 is not None # require_stage2=True (default) guarantees this + model, sec_decoder = ctx.stage1, ctx.stage2 + cond_norm, tgt_norm, sec_phys_norm = ctx.cond_norm, ctx.tgt_norm, ctx.sec_phys_norm + pdg_map, mat_map = ctx.pdg_map, ctx.mat_map + pdg_topn_map, mat_topn_map = ctx.pdg_topn_map, ctx.mat_topn_map + particle_conditioning, material_conditioning = ctx.particle_conditioning, ctx.material_conditioning + other_policy = ctx.other_policy + stage1_ddpm_steps = ctx.stage1_ddpm_steps + stage2_k_max = ctx.k_max if batch_size_auto: - est_hidden_dim, est_n_blocks = _batch_size_estimate_dims(model_cfg, training=False) + est_hidden_dim, est_n_blocks = _batch_size_estimate_dims(ctx.model_config, training=False) try: batch_size_value = gconfig.estimate_batch_size( est_hidden_dim, @@ -1037,42 +942,6 @@ def predict( assert batch_size_value is not None bs = batch_size_value - particle_conditioning, material_conditioning = _conditioning_axes(model_cfg) - if particle_conditioning == "onehot" and pdg_topn_map is None: - typer.echo( - "error: checkpoint's conditioning.particle.type='onehot' but has " - "no pdg_topn_map — retrain with the current code", - err=True, - ) - raise typer.Exit(1) - if material_conditioning == "onehot" and mat_topn_map is None: - typer.echo( - "error: checkpoint's conditioning.material.type='onehot' but has " - "no mat_topn_map — retrain with the current code", - err=True, - ) - raise typer.Exit(1) - pdg_map = {int(k): v for k, v in ckpt["pdg_map"].items()} - mat_map = {str(k): v for k, v in ckpt["mat_map"].items()} - cond_norm = Normalizer.from_dict(ckpt["normalizer"]["cond"]) - tgt_norm = Normalizer.from_dict(ckpt["normalizer"]["target"]) - sec_phys_norm = Normalizer.from_dict(ckpt["normalizer"]["sec_phys"]) - - built = build_models(model_cfg) - model, sec_decoder = built["stage1"], built["stage2"] - if model is None or sec_decoder is None: - typer.echo( - "error: checkpoint has an inactive stage1 or stage2 — giant " - "predict needs both (see stage{1,2}_model.active)", - err=True, - ) - raise typer.Exit(1) - _load_model_weights(model, sec_decoder, ckpt, weights, checkpoint) - model.to(_device).eval() - sec_decoder.to(_device).eval() - - typer.echo(f"loaded checkpoint: {checkpoint} (weights: {weights.value})") - gconfig.warn_if_checkpoint_config_mismatch(checkpoint) # --- Output path --- out, dataset_path, pred_uuid = _resolve_prediction_output(data, out) @@ -1094,8 +963,12 @@ def predict( return iter_file_chunks(path, offset=offset, k_max=stage2_k_max) return iter_cond_chunks(path, offset=offset) - cond_pdg_topn = pdg_topn_map.class_map if particle_conditioning == "onehot" else None - cond_mat_topn = mat_topn_map.class_map if material_conditioning == "onehot" else None + # load_for_inference already guarantees pdg_topn_map/mat_topn_map are not + # None whenever the matching conditioning axis is "onehot" — the extra + # `is not None` conjuncts below are redundant at runtime, just narrowing + # for the type checker. + cond_pdg_topn = pdg_topn_map.class_map if pdg_topn_map is not None and particle_conditioning == "onehot" else None + cond_mat_topn = mat_topn_map.class_map if mat_topn_map is not None and material_conditioning == "onehot" else None def _concat(a: dict[str, np.ndarray], b: dict[str, np.ndarray]) -> dict[str, np.ndarray]: return {k: np.concatenate([a[k], b[k]], axis=0) for k in a} @@ -1405,65 +1278,24 @@ def rollout( _device = torch.device(device) if device else gconfig.auto_device() typer.echo(f"device: {_device}") - ckpt = torch.load(checkpoint, map_location="cpu", weights_only=False) - for key in ("model_config", "sec_decoder"): - if key not in ckpt: - typer.echo( - f"error: checkpoint has no {key} — retrain with the current code", - err=True, - ) - raise typer.Exit(1) - - if "sec_phys" not in ckpt.get("normalizer", {}): - typer.echo( - "error: checkpoint has no normalizer.sec_phys — retrain with the current code", - err=True, - ) + try: + ctx = load_for_inference(checkpoint, _device, "rollout", weights=weights.value) + except CheckpointCompatibilityError as exc: + typer.echo(f"error: {exc}", err=True) raise typer.Exit(1) + typer.echo(f"loaded checkpoint: {checkpoint} (weights: {weights.value})") - gconfig.warn_if_checkpoint_config_mismatch(checkpoint) training_cfg = gconfig.load_checkpoint_config(checkpoint) - model_cfg = ckpt["model_config"] - particle_conditioning, material_conditioning = _conditioning_axes(model_cfg) - pdg_topn_map = _load_pdg_topn_map(ckpt) - mat_topn_map = _load_mat_topn_map(ckpt) - if particle_conditioning == "onehot" and pdg_topn_map is None: - typer.echo( - "error: checkpoint's conditioning.particle.type='onehot' but has " - "no pdg_topn_map — retrain with the current code", - err=True, - ) - raise typer.Exit(1) - if material_conditioning == "onehot" and mat_topn_map is None: - typer.echo( - "error: checkpoint's conditioning.material.type='onehot' but has " - "no mat_topn_map — retrain with the current code", - err=True, - ) - raise typer.Exit(1) - other_policy = _particle_type_other_policy(model_cfg) - stage1_ddpm_steps = _ddpm_steps(model_cfg, "stage1") - stage2_ddpm_steps = _ddpm_steps(model_cfg, "stage2") - pdg_map = {int(k): v for k, v in ckpt["pdg_map"].items()} - mat_map = {str(k): v for k, v in ckpt["mat_map"].items()} - cond_norm = Normalizer.from_dict(ckpt["normalizer"]["cond"]) - tgt_norm = Normalizer.from_dict(ckpt["normalizer"]["target"]) - sec_phys_norm = Normalizer.from_dict(ckpt["normalizer"]["sec_phys"]) - - built = build_models(model_cfg) - model, sec_decoder = built["stage1"], built["stage2"] - if model is None or sec_decoder is None: - typer.echo( - "error: checkpoint has an inactive stage1 or stage2 — giant " - "rollout needs both (see stage{1,2}_model.active)", - err=True, - ) - raise typer.Exit(1) - _load_model_weights(model, sec_decoder, ckpt, weights, checkpoint) - model.to(_device).eval() - sec_decoder.to(_device).eval() - typer.echo(f"loaded checkpoint: {checkpoint} (weights: {weights.value})") + assert ctx.stage1 is not None and ctx.stage2 is not None # require_stage2=True (default) guarantees this + model, sec_decoder = ctx.stage1, ctx.stage2 + cond_norm, tgt_norm, sec_phys_norm = ctx.cond_norm, ctx.tgt_norm, ctx.sec_phys_norm + pdg_map, mat_map = ctx.pdg_map, ctx.mat_map + pdg_topn_map, mat_topn_map = ctx.pdg_topn_map, ctx.mat_topn_map + particle_conditioning, material_conditioning = ctx.particle_conditioning, ctx.material_conditioning + other_policy = ctx.other_policy + stage1_ddpm_steps, stage2_ddpm_steps = ctx.stage1_ddpm_steps, ctx.stage2_ddpm_steps + model_cfg = ctx.model_config oracle = GeometryOracle.load(geometry) typer.echo(f"loaded geometry oracle: {geometry} (escape_threshold={oracle.escape_threshold:.3f})") @@ -1560,8 +1392,8 @@ def rollout( # model knob (router type/n_experts, noise_dim, vocab sizes, ...) # is available downstream without touching this command again. "model_config": dict(model_cfg), - "training_epoch": ckpt.get("epoch"), - "best_val_loss": ckpt.get("best_val_loss"), + "training_epoch": ctx.epoch, + "best_val_loss": ctx.best_val_loss, # [train]/[meta] from the sibling config.toml (giant.config.save_config) # — empty dicts if the checkpoint has no config.toml next to it. "training_config": dict(training_cfg.get("train", {})), diff --git a/tests/test_checkpoint_io.py b/tests/test_checkpoint_io.py new file mode 100644 index 0000000..8107945 --- /dev/null +++ b/tests/test_checkpoint_io.py @@ -0,0 +1,236 @@ +"""Tests for giant.checkpoint_io.load_for_inference (issues.md Issue 5) — +the shared bootstrap `giant predict`/`giant rollout` use to go from a +checkpoint path to ready-to-run models.""" + +from __future__ import annotations + +import copy + +import numpy as np +import pytest +import torch + +from giant import config as gconfig +from giant.checkpoint_io import ( + CheckpointCompatibilityError, + InferenceContext, + conditioning_axes, + load_for_inference, + stage_cfg, +) +from giant.data.loader import TopNMap +from giant.data.setup_cache import topnmap_to_json +from giant.data.transforms import Normalizer +from giant.model.network import build_models + +PDG_MAP = {11: 0, 22: 1, -11: 2} +MAT_MAP = {"G4_PbWO4": 0, "G4_AIR": 1} + + +def _model_cfg(stage2_active: bool = True) -> dict: + """DEFAULT_CONFIG-derived, shrunk for speed — same pattern as + tests/test_network.py::_minimal_model_config. Default `conditioning` + (both axes "physical") needs no top-N vocab map, so this is a cheap, + fully self-contained happy-path config.""" + cfg = copy.deepcopy(gconfig.DEFAULT_CONFIG) + cfg["conditioning"]["particle"]["emb_dim"] = 4 + cfg["conditioning"]["material"]["emb_dim"] = 4 + cfg["stage1_model"].update({"hidden_dim": 8, "n_res_blocks": 1}) + cfg["stage2_model"].update({"hidden_dim": 8, "n_res_blocks": 1, "k_max": 3}) + cfg["stage2_model"]["active"] = stage2_active + return { + "pdg_vocab": len(PDG_MAP), + "mat_vocab": len(MAT_MAP), + "conditioning": cfg["conditioning"], + "stage1_model": cfg["stage1_model"], + "stage2_model": cfg["stage2_model"], + } + + +def _norms() -> tuple[Normalizer, Normalizer, Normalizer]: + rng = np.random.default_rng(0) + cond = Normalizer().fit(rng.standard_normal((100, 15)).astype(np.float32)) + tgt = Normalizer().fit(rng.standard_normal((100, 9)).astype(np.float32)) + sec_phys = Normalizer().fit(rng.standard_normal((100, 2)).astype(np.float32)) + return cond, tgt, sec_phys + + +def _write_checkpoint(tmp_path, model_cfg=None, ema: bool = False, **ckpt_overrides): + cfg = model_cfg if model_cfg is not None else _model_cfg() + built = build_models(cfg) + stage1, stage2 = built["stage1"], built["stage2"] + cond, tgt, sec_phys = _norms() + ckpt: dict = { + "model_config": cfg, + "model": stage1.state_dict() if stage1 is not None else {}, + "sec_decoder": stage2.state_dict() if stage2 is not None else {}, + "pdg_map": PDG_MAP, + "mat_map": MAT_MAP, + "normalizer": {"cond": cond.to_dict(), "target": tgt.to_dict(), "sec_phys": sec_phys.to_dict()}, + "epoch": 3, + "best_val_loss": 0.5, + } + if ema: + ckpt["model_ema"] = stage1.state_dict() if stage1 is not None else {} + ckpt["sec_decoder_ema"] = stage2.state_dict() if stage2 is not None else {} + ckpt.update(ckpt_overrides) + path = tmp_path / "ckpt.pt" + torch.save(ckpt, path) + return path + + +def _onehot_model_cfg() -> dict: + cfg = _model_cfg() + cfg["conditioning"]["particle"]["type"] = "onehot" + return cfg + + +# --------------------------------------------------------------------------- +# Happy path +# --------------------------------------------------------------------------- + + +def test_happy_path_returns_populated_context(tmp_path): + checkpoint = _write_checkpoint(tmp_path) + ctx = load_for_inference(checkpoint, torch.device("cpu"), "predict") + + assert isinstance(ctx, InferenceContext) + assert ctx.stage1 is not None and ctx.stage2 is not None + assert not ctx.stage1.training + assert not ctx.stage2.training + assert next(ctx.stage1.parameters()).device == torch.device("cpu") + assert ctx.pdg_map == PDG_MAP + assert ctx.mat_map == MAT_MAP + assert all(isinstance(k, int) for k in ctx.pdg_map) + assert all(isinstance(k, str) for k in ctx.mat_map) + assert ctx.particle_conditioning == "physical" + assert ctx.material_conditioning == "physical" + assert ctx.k_max == 3 + assert ctx.epoch == 3 + assert ctx.best_val_loss == 0.5 + assert ctx.model_config["stage1_model"]["hidden_dim"] == 8 + + +def test_happy_path_normalizer_values_round_trip(tmp_path): + cond, tgt, sec_phys = _norms() + checkpoint = _write_checkpoint(tmp_path) + ctx = load_for_inference(checkpoint, torch.device("cpu"), "predict") + + assert ctx.cond_norm.mean is not None and cond.mean is not None + assert ctx.tgt_norm.mean is not None and tgt.mean is not None + assert ctx.sec_phys_norm.mean is not None and sec_phys.mean is not None + np.testing.assert_allclose(ctx.cond_norm.mean, cond.mean) + np.testing.assert_allclose(ctx.tgt_norm.mean, tgt.mean) + np.testing.assert_allclose(ctx.sec_phys_norm.mean, sec_phys.mean) + + +# --------------------------------------------------------------------------- +# Guards +# --------------------------------------------------------------------------- + + +def test_missing_model_config_raises(tmp_path): + checkpoint = _write_checkpoint(tmp_path) + ckpt = torch.load(checkpoint, weights_only=False) + del ckpt["model_config"] + torch.save(ckpt, checkpoint) + + with pytest.raises(CheckpointCompatibilityError, match="no model_config"): + load_for_inference(checkpoint, torch.device("cpu"), "predict") + + +def test_missing_sec_decoder_raises(tmp_path): + checkpoint = _write_checkpoint(tmp_path) + ckpt = torch.load(checkpoint, weights_only=False) + del ckpt["sec_decoder"] + torch.save(ckpt, checkpoint) + + with pytest.raises(CheckpointCompatibilityError, match="no sec_decoder"): + load_for_inference(checkpoint, torch.device("cpu"), "predict") + + +def test_missing_sec_phys_normalizer_raises(tmp_path): + checkpoint = _write_checkpoint(tmp_path) + ckpt = torch.load(checkpoint, weights_only=False) + del ckpt["normalizer"]["sec_phys"] + torch.save(ckpt, checkpoint) + + with pytest.raises(CheckpointCompatibilityError, match="no normalizer.sec_phys"): + load_for_inference(checkpoint, torch.device("cpu"), "predict") + + +def test_onehot_particle_conditioning_without_topn_map_raises(tmp_path): + checkpoint = _write_checkpoint(tmp_path, model_cfg=_onehot_model_cfg()) + + with pytest.raises(CheckpointCompatibilityError, match="pdg_topn_map"): + load_for_inference(checkpoint, torch.device("cpu"), "predict") + + +def test_onehot_particle_conditioning_with_topn_map_succeeds(tmp_path): + topn = TopNMap(class_map={11: 0, 22: 1}, other_members={}) + checkpoint = _write_checkpoint( + tmp_path, + model_cfg=_onehot_model_cfg(), + pdg_topn_map=topnmap_to_json(topn), + ) + ctx = load_for_inference(checkpoint, torch.device("cpu"), "predict") + assert ctx.particle_conditioning == "onehot" + assert ctx.pdg_topn_map is not None + assert ctx.pdg_topn_map.class_map == {11: 0, 22: 1} + + +def test_ema_weights_requested_but_missing_raises(tmp_path): + checkpoint = _write_checkpoint(tmp_path, ema=False) + + with pytest.raises(CheckpointCompatibilityError, match="no EMA weights"): + load_for_inference(checkpoint, torch.device("cpu"), "predict", weights="ema") + + +def test_ema_weights_requested_and_present_succeeds(tmp_path): + checkpoint = _write_checkpoint(tmp_path, ema=True) + ctx = load_for_inference(checkpoint, torch.device("cpu"), "predict", weights="ema") + assert ctx.stage1 is not None and ctx.stage2 is not None + + +@pytest.mark.parametrize("command_name", ["predict", "rollout"]) +def test_inactive_stage_with_require_stage2_raises_with_command_name(tmp_path, command_name): + checkpoint = _write_checkpoint(tmp_path, model_cfg=_model_cfg(stage2_active=False)) + + with pytest.raises(CheckpointCompatibilityError, match=f"{command_name} needs both"): + load_for_inference(checkpoint, torch.device("cpu"), command_name) + + +def test_inactive_stage_with_require_stage2_false_succeeds_with_stage2_none(tmp_path): + checkpoint = _write_checkpoint(tmp_path, model_cfg=_model_cfg(stage2_active=False)) + + ctx = load_for_inference(checkpoint, torch.device("cpu"), "predict", require_stage2=False) + assert ctx.stage1 is not None + assert ctx.stage2 is None + + +# --------------------------------------------------------------------------- +# conditioning_axes / stage_cfg +# --------------------------------------------------------------------------- + + +def test_conditioning_axes_v02_flat_string_applies_to_both_axes(): + assert conditioning_axes({"conditioning": "embedding"}) == ("embedding", "embedding") + + +def test_conditioning_axes_v03_nested_dict_independent_per_axis(): + model_cfg = {"conditioning": {"particle": {"type": "onehot"}, "material": {"type": "physical"}}} + assert conditioning_axes(model_cfg) == ("onehot", "physical") + + +def test_conditioning_axes_missing_key_uses_default(): + assert conditioning_axes({}, default="embedding") == ("embedding", "embedding") + + +def test_stage_cfg_new_shape_returns_subdict(): + model_cfg = {"stage2_model": {"k_max": 7}} + assert stage_cfg(model_cfg, "stage2") == {"k_max": 7} + + +def test_stage_cfg_v02_flat_shape_returns_empty_dict(): + model_cfg = {"hidden_dim": 32, "n_blocks": 4} + assert stage_cfg(model_cfg, "stage2") == {} diff --git a/tests/test_cli_predict.py b/tests/test_cli_predict.py index 1c18ae8..7c8e237 100644 --- a/tests/test_cli_predict.py +++ b/tests/test_cli_predict.py @@ -1,13 +1,18 @@ import uuid +import torch import yaml +from typer.testing import CliRunner from giant.cli import ( _CEPH_PREDICTIONS, _resolve_prediction_output, _write_prediction_ref, + app, ) +runner = CliRunner() + # --------------------------------------------------------------------------- # _resolve_prediction_output @@ -150,3 +155,20 @@ def test_ref_checkpoint_path_is_absolute(tmp_path): data = yaml.safe_load(ref_path.read_text()) assert data["checkpoint"].startswith("/") + + +# --------------------------------------------------------------------------- +# Bootstrap failure surfaces via the CLI (issues.md Issue 5 — confirms +# CheckpointCompatibilityError -> typer.Exit(1) actually wires up end-to-end, +# not just at the giant.checkpoint_io unit level). +# --------------------------------------------------------------------------- + + +def test_predict_exits_1_on_checkpoint_missing_model_config(tmp_path): + checkpoint = tmp_path / "bad.pt" + torch.save({"sec_decoder": {}, "normalizer": {"sec_phys": {}}}, checkpoint) + + result = runner.invoke(app, ["predict", "dummy.parquet", "--checkpoint", str(checkpoint)]) + + assert result.exit_code == 1 + assert "checkpoint has no model_config" in result.output diff --git a/tests/test_cli_rollout.py b/tests/test_cli_rollout.py new file mode 100644 index 0000000..b5bd2a6 --- /dev/null +++ b/tests/test_cli_rollout.py @@ -0,0 +1,33 @@ +"""Thin CLI smoke coverage for `giant rollout` (issues.md Issue 5) — confirms +the CheckpointCompatibilityError raised by giant.checkpoint_io.load_for_inference +surfaces as a clean typer.Exit(1) with the expected message, end-to-end +through the CLI, not just at the giant.checkpoint_io unit level.""" + +from __future__ import annotations + +import torch +from typer.testing import CliRunner + +from giant.cli import app + +runner = CliRunner() + + +def test_rollout_exits_1_on_checkpoint_missing_model_config(tmp_path): + checkpoint = tmp_path / "bad.pt" + torch.save({"sec_decoder": {}, "normalizer": {"sec_phys": {}}}, checkpoint) + + result = runner.invoke( + app, + [ + "rollout", + "dummy.parquet", + "--checkpoint", + str(checkpoint), + "--geometry", + "dummy_geometry.pkl", + ], + ) + + assert result.exit_code == 1 + assert "checkpoint has no model_config" in result.output