From bd255419e1aab6e4930b81ba82b27d79245df789 Mon Sep 17 00:00:00 2001 From: Lars Bogner Date: Fri, 28 Aug 2026 11:21:36 +0200 Subject: [PATCH] Add inference-time model_config overrides with a sampling-key allowlist (gitea #87) giant predict/rollout rebuilt models straight from ckpt["model_config"] with no way to change sampling-only keys (e.g. stage2_model.n_sec.stop_sampling) without retraining. Adds config_overrides to load_for_inference, validated against giant.config.INFERENCE_OVERRIDES so a typo or shape-bearing key raises CheckpointCompatibilityError up front instead of an opaque load_state_dict mismatch. Wired as a repeatable --set dotted.path=value on both CLI commands, recorded in the rollout YAML sidecar, and surfaced in `giant model summary`'s output. Co-Authored-By: Claude Sonnet 5 --- giant/checkpoint_io.py | 55 ++++++++++++++++- giant/cli.py | 46 +++++++++++++- giant/config.py | 78 ++++++++++++++++++++++- giant/model/summary.py | 12 +++- tests/test_checkpoint_io.py | 120 ++++++++++++++++++++++++++++++++++++ tests/test_cli_predict.py | 40 ++++++++++++ tests/test_cli_rollout.py | 25 ++++++++ tests/test_model_summary.py | 14 +++++ 8 files changed, 383 insertions(+), 7 deletions(-) diff --git a/giant/checkpoint_io.py b/giant/checkpoint_io.py index fe6ce63..9bf0e9a 100644 --- a/giant/checkpoint_io.py +++ b/giant/checkpoint_io.py @@ -14,11 +14,18 @@ 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)`. + +`load_for_inference`'s `config_overrides` (gitea #87) lets a caller change a +checkpoint's `model_config` at load time, restricted to +`giant.config.INFERENCE_OVERRIDES` — the allowlist of keys that only affect +sampling, never module construction/shapes or the preprocessing normalizers/ +vocab maps were fit under. """ from __future__ import annotations -from dataclasses import dataclass +import copy +from dataclasses import dataclass, field from pathlib import Path import torch @@ -29,13 +36,45 @@ 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 +from giant.model.network import _migrate_legacy_model_config, build_models class CheckpointCompatibilityError(Exception): """Checkpoint is missing something `load_for_inference` needs.""" +def apply_config_overrides(model_cfg: dict, overrides: dict[str, object] | None) -> dict: + """Deep-merge dotted-path *overrides* into a checkpoint's `model_config`, + validated against `giant.config.INFERENCE_OVERRIDES` — the allowlist of + keys that only affect sampling, not module construction/shapes or the + preprocessing normalizers/vocab maps were fit under (gitea #87). + + Migrates a v0.2 flat `model_config` to the nested v0.3 shape first: a + dotted path like "stage1_model.ddpm.n_steps" would otherwise silently + write into a dict that `build_models` still reads as flat (it decides + v0.2-vs-v0.3 by `"stage1_model" in model_config`), suppressing migration. + + Raises `CheckpointCompatibilityError` — never a bare `ValueError` or a + downstream `load_state_dict` size mismatch — for an unknown/disallowed + path or a value that fails its allowlisted check. + """ + if not overrides: + return model_cfg + cfg = model_cfg if "stage1_model" in model_cfg else _migrate_legacy_model_config(model_cfg) + cfg = copy.deepcopy(cfg) + for path, value in overrides.items(): + spec = gconfig.INFERENCE_OVERRIDES.get(path) + if spec is None: + allowed = ", ".join(sorted(gconfig.INFERENCE_OVERRIDES)) + raise CheckpointCompatibilityError(f"{path!r} is not an inference-safe override — allowed paths: {allowed}") + try: + spec.check(path, value) + except ValueError as exc: + raise CheckpointCompatibilityError(str(exc)) from exc + gconfig._set_path(cfg, path, value) + return cfg + + 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 @@ -128,6 +167,7 @@ class InferenceContext: model_config: dict epoch: int | None best_val_loss: float | None + config_overrides: dict[str, object] = field(default_factory=dict) def load_for_inference( @@ -136,6 +176,7 @@ def load_for_inference( command_name: str, weights: str = "raw", require_stage2: bool = True, + config_overrides: dict[str, object] | None = None, ) -> InferenceContext: """Load *checkpoint* and reconstruct everything `predict`/`rollout` need to run it forward, on *device*, in `eval()` mode. @@ -148,6 +189,13 @@ def load_for_inference( 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. + + *config_overrides* deep-merges dotted `model_config` paths (e.g. + `{"stage2_model.n_sec.stop_sampling": "sample"}`) before anything is + derived from `model_config` or built — see `apply_config_overrides` for + the allowlist and validation. Every derived `InferenceContext` field + (`other_policy`, `stage{1,2}_ddpm_steps`, the built modules, ...) + reflects the overridden config. """ ckpt = torch.load(checkpoint, map_location="cpu", weights_only=False) for key in ("model_config", "sec_decoder"): @@ -159,7 +207,7 @@ def load_for_inference( gconfig.warn_if_checkpoint_config_mismatch(checkpoint) - model_cfg = ckpt["model_config"] + model_cfg = apply_config_overrides(ckpt["model_config"], config_overrides) particle_conditioning, material_conditioning = conditioning_axes(model_cfg) pdg_topn_map = load_pdg_topn_map(ckpt) mat_topn_map = load_mat_topn_map(ckpt) @@ -231,4 +279,5 @@ def load_for_inference( model_config=model_cfg, epoch=ckpt.get("epoch"), best_val_loss=ckpt.get("best_val_loss"), + config_overrides=dict(config_overrides) if config_overrides else {}, ) diff --git a/giant/cli.py b/giant/cli.py index 9f7433c..5d2f2ac 100644 --- a/giant/cli.py +++ b/giant/cli.py @@ -141,6 +141,23 @@ def _parse_router_axis_flags(specs: list[str]) -> dict[str, object]: return out +def _parse_set_flags(specs: Optional[list[str]]) -> dict[str, object]: + """Parse repeated `--set dotted.path=value` flags into a dict, typing + each value with `_coerce_scalar` the same way a TOML file's native types + would arrive. Validation against the inference-safe allowlist happens + downstream in `giant.checkpoint_io.apply_config_overrides` — this only + parses syntax. + """ + out: dict[str, object] = {} + for spec in specs or []: + path, sep, val = spec.partition("=") + if not sep: + typer.echo(f"error: --set {spec!r} must be 'dotted.path=value'", err=True) + raise typer.Exit(1) + out[path] = _coerce_scalar(val) + return out + + def _router_cli_overrides( router: bool | None, router_type: str | None, @@ -1020,6 +1037,15 @@ def predict( help="Free-text note recorded in the prediction's YAML sidecar", ), ] = None, + set_: Annotated[ + Optional[list[str]], + typer.Option( + "--set", + help="Override a sampling-only model_config key on this checkpoint, " + "'dotted.path=value' (repeatable) — see giant.config.INFERENCE_OVERRIDES " + "for the allowlist, e.g. --set stage2_model.n_sec.stop_sampling=sample", + ), + ] = None, ) -> None: """Run trained model on a parquet file and save predictions.""" batch_size_auto = False @@ -1040,8 +1066,11 @@ def predict( typer.echo(f"device: {_device}") # --- Load checkpoint --- + config_overrides = _parse_set_flags(set_) try: - ctx = load_for_inference(checkpoint, _device, "predict", weights=weights.value) + ctx = load_for_inference( + checkpoint, _device, "predict", weights=weights.value, config_overrides=config_overrides + ) except CheckpointCompatibilityError as exc: typer.echo(f"error: {exc}", err=True) raise typer.Exit(1) @@ -1407,6 +1436,15 @@ def rollout( Optional[int], typer.Option("--seed", help="Torch/numpy seed for reproducibility"), ] = None, + set_: Annotated[ + Optional[list[str]], + typer.Option( + "--set", + help="Override a sampling-only model_config key on this checkpoint, " + "'dotted.path=value' (repeatable) — see giant.config.INFERENCE_OVERRIDES " + "for the allowlist, e.g. --set stage2_model.n_sec.stop_sampling=sample", + ), + ] = None, ) -> None: """Roll the surrogate forward into full showers (autoregressive).""" if seed is not None: @@ -1416,8 +1454,11 @@ def rollout( _device = torch.device(device) if device else gconfig.auto_device() typer.echo(f"device: {_device}") + config_overrides = _parse_set_flags(set_) try: - ctx = load_for_inference(checkpoint, _device, "rollout", weights=weights.value) + ctx = load_for_inference( + checkpoint, _device, "rollout", weights=weights.value, config_overrides=config_overrides + ) except CheckpointCompatibilityError as exc: typer.echo(f"error: {exc}", err=True) raise typer.Exit(1) @@ -1532,6 +1573,7 @@ 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), + "config_overrides": dict(ctx.config_overrides), "training_epoch": ctx.epoch, "best_val_loss": ctx.best_val_loss, # [train]/[meta] from the sibling config.toml (giant.config.save_config) diff --git a/giant/config.py b/giant/config.py index 5ee2eaa..6f32931 100644 --- a/giant/config.py +++ b/giant/config.py @@ -404,6 +404,16 @@ class Stage2RouterConfig(RouterConfig): return {"tie_to_stage1": self.tie_to_stage1, **super().to_dict()} +# stage2_model.n_sec.stop_sampling choices — single source of truth for both +# validate_config's train-time check and INFERENCE_OVERRIDES below. +STOP_SAMPLING_CHOICES = ("greedy", "sample") + +# stage2_model.particle_type.other_policy choices — see ParticleTypeConfig's +# docstring for what each means; only documented there until now, since +# nothing validated it at train time. +OTHER_POLICY_CHOICES = ("sample", "modal", "drop") + + @dataclass(frozen=True) class NSecConfig: # "head": a classifier over {0..k_max} on the condition encoding alone @@ -1093,6 +1103,72 @@ def _deep_merge(base: dict, override: dict) -> dict: return result +@dataclass(frozen=True) +class InferenceOverride: + """One dotted `model_config` path that `giant.checkpoint_io.load_for_inference` + is allowed to change on an already-trained checkpoint, without retraining. + + A path only belongs here if it affects neither module construction/tensor + shapes nor the data preprocessing the normalizers/vocab maps were fit + under — see the module docstring on `giant.model.summary` for the class + of key this targets (`_fingerprint`'s "plain scalar attribute" leaves), + and `giant.checkpoint_io.apply_config_overrides` for where this is used. + """ + + why: str + choices: tuple[str, ...] | None = None + minimum: float | None = None + numeric: bool = False # int/float leaf (vs. str, the default) + + def check(self, path: str, value: object) -> None: + if self.choices is not None: + if value not in self.choices: + raise ValueError(f"{path} = {value!r} — must be one of {self.choices}") + return + if self.numeric: + if isinstance(value, bool) or not isinstance(value, (int, float)): + raise ValueError(f"{path} = {value!r} — must be a number") + if self.minimum is not None and value < self.minimum: + raise ValueError(f"{path} = {value!r} — must be >= {self.minimum}") + + +# Inference-safe dotted `model_config` paths — the allowlist gitea #87 asked +# for, so a typo or a shape-bearing key (e.g. "stage1_model.hidden_dim") +# raises a clear CheckpointCompatibilityError instead of surfacing as an +# opaque load_state_dict size mismatch later. Extend this table, not a +# per-call bypass, when a new inference-only key needs the capability. +INFERENCE_OVERRIDES: dict[str, InferenceOverride] = { + "stage2_model.n_sec.stop_sampling": InferenceOverride( + why="giant.sample's stop-token loop reads this at sample time only", + choices=STOP_SAMPLING_CHOICES, + ), + "stage1_model.ddpm.n_steps": InferenceOverride( + why="giant.model.schedule.CosineSchedule's step count, resolved at sample time", + numeric=True, + minimum=1, + ), + "stage2_model.ddpm.n_steps": InferenceOverride( + why="giant.model.schedule.CosineSchedule's step count, resolved at sample time", + numeric=True, + minimum=1, + ), + "stage2_model.particle_type.other_policy": InferenceOverride( + why="giant.rollout resolves an 'other'-bucket secondary's PDG code with this at rollout time", + choices=OTHER_POLICY_CHOICES, + ), + "stage1_model.router.temperature": InferenceOverride( + why="giant.model.routers.EnergyRouter.temperature, a plain constructor attribute", + numeric=True, + minimum=1e-6, + ), + "stage2_model.router.temperature": InferenceOverride( + why="giant.model.routers.EnergyRouter.temperature, a plain constructor attribute", + numeric=True, + minimum=1e-6, + ), +} + + @dataclass(frozen=True) class FlagSpec: """One CLI flag's mapping into the config-overrides tree. @@ -1513,7 +1589,7 @@ def validate_config(cfg: dict, *, resume: bool = False) -> None: ) stop_sampling = _get_path(cfg, "stage2_model.n_sec.stop_sampling") - if stop_sampling not in ("greedy", "sample"): + if stop_sampling not in STOP_SAMPLING_CHOICES: raise ValueError(f"stage2_model.n_sec.stop_sampling = {stop_sampling!r} — must be 'greedy' or 'sample'") precision = _get_path(cfg, "train.precision") diff --git a/giant/model/summary.py b/giant/model/summary.py index 36a34d3..5e728b4 100644 --- a/giant/model/summary.py +++ b/giant/model/summary.py @@ -37,7 +37,7 @@ from dataclasses import dataclass, field import torch.nn as nn -from giant.config import _get_path, _set_path, leaf_paths +from giant.config import INFERENCE_OVERRIDES, _get_path, _set_path, leaf_paths from giant.model.builders import build_critics, build_models from giant.model.trunks import RoutedTrunk @@ -111,6 +111,7 @@ class ModelSummary: pdg_vocab: int mat_vocab: int vocab_caveats: list[str] = field(default_factory=list) + overridable: list[str] = field(default_factory=list) def _build_model_config(cfg: dict, pdg_vocab: int, mat_vocab: int) -> dict: @@ -234,6 +235,8 @@ def summarize_model(cfg: dict, pdg_vocab: int, mat_vocab: int) -> ModelSummary: else: inert.append(path) + overridable = sorted(p for p in in_scope if p in INFERENCE_OVERRIDES) + return ModelSummary( modules=modules, consumed=sorted(consumed), @@ -242,6 +245,7 @@ def summarize_model(cfg: dict, pdg_vocab: int, mat_vocab: int) -> ModelSummary: pdg_vocab=pdg_vocab, mat_vocab=mat_vocab, vocab_caveats=_vocab_caveats(cfg), + overridable=overridable, ) @@ -309,6 +313,12 @@ def render_summary(summary: ModelSummary) -> str: else: lines.append(" (none)") + if summary.overridable: + lines.append("") + lines.append("inference-overridable without retraining (giant predict/rollout --set):") + for path in summary.overridable: + lines.append(f" {path} ({INFERENCE_OVERRIDES[path].why})") + if summary.vocab_caveats: lines.append("") lines.append("vocab placeholder caveats:") diff --git a/tests/test_checkpoint_io.py b/tests/test_checkpoint_io.py index fcceaf8..f70c80e 100644 --- a/tests/test_checkpoint_io.py +++ b/tests/test_checkpoint_io.py @@ -14,6 +14,7 @@ from giant import config as gconfig from giant.checkpoint_io import ( CheckpointCompatibilityError, InferenceContext, + apply_config_overrides, conditioning_axes, load_for_inference, stage_cfg, @@ -278,3 +279,122 @@ def test_stage_cfg_new_shape_returns_subdict(): def test_stage_cfg_v02_flat_shape_returns_empty_dict(): model_cfg = {"hidden_dim": 32, "n_blocks": 4} assert stage_cfg(model_cfg, "stage2") == {} + + +# --------------------------------------------------------------------------- +# config_overrides (gitea #87) +# --------------------------------------------------------------------------- + + +def _router_model_cfg() -> dict: + cfg = _model_cfg() + cfg["stage1_model"]["router"] = {"enabled": True, "type": "energy", "n_experts": 2} + return cfg + + +def test_config_override_stop_sampling_changes_stage2_attribute(tmp_path): + checkpoint = _write_checkpoint(tmp_path) + ctx = load_for_inference( + checkpoint, + torch.device("cpu"), + "predict", + config_overrides={"stage2_model.n_sec.stop_sampling": "sample"}, + ) + assert ctx.stage2 is not None + assert ctx.stage2.stop_sampling == "sample" + assert ctx.config_overrides == {"stage2_model.n_sec.stop_sampling": "sample"} + + +def test_config_override_ddpm_n_steps_changes_context_fields(tmp_path): + checkpoint = _write_checkpoint(tmp_path) + ctx = load_for_inference( + checkpoint, + torch.device("cpu"), + "predict", + config_overrides={"stage1_model.ddpm.n_steps": 42, "stage2_model.ddpm.n_steps": 7}, + ) + assert ctx.stage1_ddpm_steps == 42 + assert ctx.stage2_ddpm_steps == 7 + + +def test_config_override_other_policy_changes_context_field(tmp_path): + checkpoint = _write_checkpoint(tmp_path) + ctx = load_for_inference( + checkpoint, + torch.device("cpu"), + "predict", + config_overrides={"stage2_model.particle_type.other_policy": "modal"}, + ) + assert ctx.other_policy == "modal" + + +def test_config_override_router_temperature_changes_router_attribute(tmp_path): + checkpoint = _write_checkpoint(tmp_path, model_cfg=_router_model_cfg()) + ctx = load_for_inference( + checkpoint, + torch.device("cpu"), + "predict", + config_overrides={"stage1_model.router.temperature": 1.5}, + ) + assert ctx.stage1 is not None + assert ctx.stage1.trunk.router.temperature == pytest.approx(1.5) + + +def test_config_override_no_overrides_defaults_to_empty_dict(tmp_path): + checkpoint = _write_checkpoint(tmp_path) + ctx = load_for_inference(checkpoint, torch.device("cpu"), "predict") + assert ctx.config_overrides == {} + + +def test_config_override_unknown_path_raises(tmp_path): + checkpoint = _write_checkpoint(tmp_path) + with pytest.raises(CheckpointCompatibilityError, match="not an inference-safe override"): + load_for_inference( + checkpoint, + torch.device("cpu"), + "predict", + config_overrides={"stage2_model.n_sec.typo": "sample"}, + ) + + +def test_config_override_shape_bearing_key_raises_up_front(tmp_path): + checkpoint = _write_checkpoint(tmp_path) + with pytest.raises(CheckpointCompatibilityError, match="not an inference-safe override"): + load_for_inference( + checkpoint, + torch.device("cpu"), + "predict", + config_overrides={"stage1_model.hidden_dim": 999}, + ) + + +def test_config_override_bad_value_raises(tmp_path): + checkpoint = _write_checkpoint(tmp_path) + with pytest.raises(CheckpointCompatibilityError, match="must be one of"): + load_for_inference( + checkpoint, + torch.device("cpu"), + "predict", + config_overrides={"stage2_model.n_sec.stop_sampling": "maybe"}, + ) + + +def test_apply_config_overrides_no_overrides_returns_same_object(): + cfg = _model_cfg() + assert apply_config_overrides(cfg, None) is cfg + assert apply_config_overrides(cfg, {}) is cfg + + +def test_apply_config_overrides_migrates_legacy_flat_model_config_first(): + legacy_cfg = { + "pdg_vocab": len(PDG_MAP), + "mat_vocab": len(MAT_MAP), + "hidden_dim": 32, + "n_blocks": 4, + "emb_dim": 8, + "dropout": 0.1, + "k_max": 5, + } + merged = apply_config_overrides(legacy_cfg, {"stage1_model.ddpm.n_steps": 10}) + assert merged["stage1_model"]["ddpm"]["n_steps"] == 10 + assert merged["stage1_model"]["hidden_dim"] == 32 diff --git a/tests/test_cli_predict.py b/tests/test_cli_predict.py index 7c8e237..799aa59 100644 --- a/tests/test_cli_predict.py +++ b/tests/test_cli_predict.py @@ -172,3 +172,43 @@ def test_predict_exits_1_on_checkpoint_missing_model_config(tmp_path): assert result.exit_code == 1 assert "checkpoint has no model_config" in result.output + + +# --------------------------------------------------------------------------- +# --set (gitea #87) +# --------------------------------------------------------------------------- + + +def test_predict_set_flag_without_equals_exits_1(tmp_path): + checkpoint = tmp_path / "missing.pt" + + result = runner.invoke( + app, + ["predict", "dummy.parquet", "--checkpoint", str(checkpoint), "--set", "stop_sampling"], + ) + + assert result.exit_code == 1 + assert "must be 'dotted.path=value'" in result.output + + +def test_predict_set_flag_disallowed_path_surfaces_compat_error(tmp_path): + checkpoint = tmp_path / "ckpt.pt" + torch.save( + {"model_config": {"stage1_model": {}, "stage2_model": {}}, "sec_decoder": {}, "normalizer": {"sec_phys": {}}}, + checkpoint, + ) + + result = runner.invoke( + app, + [ + "predict", + "dummy.parquet", + "--checkpoint", + str(checkpoint), + "--set", + "stage1_model.hidden_dim=999", + ], + ) + + assert result.exit_code == 1 + assert "not an inference-safe override" in result.output diff --git a/tests/test_cli_rollout.py b/tests/test_cli_rollout.py index b5bd2a6..06857fe 100644 --- a/tests/test_cli_rollout.py +++ b/tests/test_cli_rollout.py @@ -31,3 +31,28 @@ def test_rollout_exits_1_on_checkpoint_missing_model_config(tmp_path): assert result.exit_code == 1 assert "checkpoint has no model_config" in result.output + + +def test_rollout_set_flag_disallowed_path_surfaces_compat_error(tmp_path): + checkpoint = tmp_path / "ckpt.pt" + torch.save( + {"model_config": {"stage1_model": {}, "stage2_model": {}}, "sec_decoder": {}, "normalizer": {"sec_phys": {}}}, + checkpoint, + ) + + result = runner.invoke( + app, + [ + "rollout", + "dummy.parquet", + "--checkpoint", + str(checkpoint), + "--geometry", + "dummy_geometry.pkl", + "--set", + "stage2_model.n_sec.typo=sample", + ], + ) + + assert result.exit_code == 1 + assert "not an inference-safe override" in result.output diff --git a/tests/test_model_summary.py b/tests/test_model_summary.py index 5685858..b4f310d 100644 --- a/tests/test_model_summary.py +++ b/tests/test_model_summary.py @@ -12,6 +12,8 @@ from giant.cli import app from giant.materials import MATERIAL_PROPERTIES from giant.model.summary import _NOT_BUILD_TIME, _built_modules, _vocab_caveats, summarize_model +INFERENCE_OVERRIDES = gconfig.INFERENCE_OVERRIDES + runner = CliRunner() _PDG_VOCAB = 300 @@ -53,6 +55,16 @@ def test_not_build_time_allow_list_has_no_stale_entries(): assert not stale, f"_NOT_BUILD_TIME entries no longer in DEFAULT_CONFIG: {sorted(stale)}" +def test_inference_overrides_allow_list_has_no_stale_entries(): + in_scope = set(gconfig.leaf_paths(gconfig.DEFAULT_CONFIG)) + stale = set(INFERENCE_OVERRIDES) - in_scope + assert not stale, f"INFERENCE_OVERRIDES entries no longer in DEFAULT_CONFIG: {sorted(stale)}" + + +def test_default_config_overridable_lists_every_allowlisted_path(default_summary): + assert set(default_summary.overridable) == set(INFERENCE_OVERRIDES) + + def test_router_disabled_by_default_so_its_fields_are_inert(default_summary): assert "stage1_model.router.n_experts" in default_summary.inert assert "stage1_model.router.temperature" in default_summary.inert @@ -135,3 +147,5 @@ def test_cli_default_smoke(): assert "parameters" in result.output assert "trunk" in result.output assert "inert under this config" in result.output + assert "inference-overridable without retraining" in result.output + assert "stage2_model.n_sec.stop_sampling" in result.output