Merge pull request 'Add per-stage init_from/freeze (gitea #42)' (#67) from fix/issue-42 into master
CI / Format (ruff format) (push) Successful in 28s
CI / Lint (ruff check) (push) Successful in 28s
CI / Sync project version with tag (push) Has been skipped
CI / Type check (ty) (push) Successful in 30s
CI / Tests (push) Successful in 2m35s

Reviewed-on: #67
This commit was merged in pull request #67.
This commit is contained in:
2026-08-17 14:48:38 +02:00
15 changed files with 429 additions and 17 deletions
+1
View File
@@ -139,6 +139,7 @@ Useful flags on `giant train`:
- `--router` / `--router-type` / `--n-experts` / `--router-axis` — MoE routing
- `--wandb` — log per-epoch metrics to Weights & Biases (needs `uv sync --extra wandb`); metric names are `<stage>/<split>/<metric>` plus an unprefixed run-level tail, all derived from `giant/training/trainers.py` `MetricSpec`s
- `--no-cache-setup` / `--rebuild-setup-cache` — control the setup-stage sidecar cache (vocab maps, event split, normalizer stats); `dwarf warm-cache` precomputes it
- `--stage1-init-from`/`--stage2-init-from` (checkpoint `.pt`) + `--stage1-freeze`/`--stage2-freeze` — load a stage's weights from another checkpoint and never update them, so the other stage can be retrained alone against a fixed, known-good one while still producing a complete, rollout-capable checkpoint
Config-file-only knobs (no CLI flag — use `--config config.toml`): `stage2_model.autoregressive.teacher_forcing`/`.history`, `stage2_model.particle_type.target`. v0.2 flat-schema configs and checkpoints load fine (auto-migrated).
+43 -1
View File
@@ -488,6 +488,36 @@ def train(
help="WGAN-GP (--mode wgan only): critic depth for stage 2 (default: same as generator's n_res_blocks)",
),
] = None,
stage1_init_from: Annotated[
Optional[Path],
typer.Option(
"--stage1-init-from",
help="Checkpoint .pt to load stage 1's weights from before training starts "
"(gitea #42) — combine with --stage1-freeze to retrain stage 2 alone "
"against a fixed, known-good stage 1",
),
] = None,
stage1_freeze: Annotated[
Optional[bool],
typer.Option(
"--stage1-freeze/--no-stage1-freeze",
help="Never update stage 1's weights (requires --stage1-init-from, or --resume)",
),
] = None,
stage2_init_from: Annotated[
Optional[Path],
typer.Option(
"--stage2-init-from",
help="Checkpoint .pt to load stage 2's weights from before training starts (gitea #42)",
),
] = None,
stage2_freeze: Annotated[
Optional[bool],
typer.Option(
"--stage2-freeze/--no-stage2-freeze",
help="Never update stage 2's weights (requires --stage2-init-from, or --resume)",
),
] = None,
val_fraction: Annotated[Optional[float], typer.Option("--val-fraction", "-f")] = None,
seed: Annotated[
Optional[int],
@@ -652,11 +682,15 @@ def train(
"stage1_critic_n_res_blocks": stage1_critic_n_res_blocks,
"stage2_critic_hidden_dim": stage2_critic_hidden_dim,
"stage2_critic_n_res_blocks": stage2_critic_n_res_blocks,
"stage1_init_from": str(stage1_init_from) if stage1_init_from is not None else None,
"stage1_freeze": stage1_freeze,
"stage2_init_from": str(stage2_init_from) if stage2_init_from is not None else None,
"stage2_freeze": stage2_freeze,
}
overrides = gconfig.overrides_from_flags(flag_values)
cfg = gconfig.merge_cli_overrides(gconfig.DEFAULT_CONFIG, config, overrides)
gconfig.validate_config(cfg)
gconfig.validate_config(cfg, resume=resume is not None)
t = cfg["train"]
_device = torch.device(device) if device else gconfig.auto_device()
@@ -736,6 +770,10 @@ def new_run(
stage2_k_max: Annotated[Optional[int], typer.Option("--stage2-k-max")] = None,
stage2_context_dim: Annotated[Optional[int], typer.Option("--stage2-context-dim")] = None,
stage2_stage1_context: Annotated[Optional[Stage1Context], typer.Option("--stage2-stage1-context")] = None,
stage1_init_from: Annotated[Optional[Path], typer.Option("--stage1-init-from")] = None,
stage1_freeze: Annotated[Optional[bool], typer.Option("--stage1-freeze/--no-stage1-freeze")] = None,
stage2_init_from: Annotated[Optional[Path], typer.Option("--stage2-init-from")] = None,
stage2_freeze: Annotated[Optional[bool], typer.Option("--stage2-freeze/--no-stage2-freeze")] = None,
conditioning: Annotated[Optional[Conditioning], typer.Option("--conditioning")] = None,
router: Annotated[Optional[bool], typer.Option("--router/--no-router")] = None,
router_type: Annotated[Optional[str], typer.Option("--router-type")] = None,
@@ -796,6 +834,10 @@ def new_run(
"stage2_k_max": stage2_k_max,
"stage2_context_dim": stage2_context_dim,
"stage2_stage1_context": stage2_stage1_context.value if stage2_stage1_context is not None else None,
"stage1_init_from": str(stage1_init_from) if stage1_init_from is not None else None,
"stage1_freeze": stage1_freeze,
"stage2_init_from": str(stage2_init_from) if stage2_init_from is not None else None,
"stage2_freeze": stage2_freeze,
"mode": mode.value if mode is not None else None,
"stage1_generator": stage1_generator.value if stage1_generator is not None else None,
"stage2_generator": stage2_generator.value if stage2_generator is not None else None,
+45 -1
View File
@@ -594,6 +594,19 @@ class Stage1ModelConfig:
# false skips building/training stage 1 entirely. The resulting
# checkpoint holds only stage 2 and cannot be rolled out.
active: bool = True
# Checkpoint .pt to load this stage's weights from before training starts
# (its own "model"/"sec_decoder" key, not this run's own resume state) —
# "" means start from a fresh init. See `freeze` below for the partial-
# retrain use case this exists for (gitea #42).
init_from: str = ""
# true keeps this stage's weights exactly as loaded from `init_from` —
# forward/backward still run every batch (so its loss/grad_norm metrics
# stay meaningful, and a WGAN stage's critic still gets a real signal to
# report), but its optimizer never steps. Lets a rollout-capable
# checkpoint retrain only the *other* stage against a fixed, known-good
# one (gitea #42) — `validate_config` requires `init_from` to be set
# whenever this is true, unless the run is a `--resume`.
freeze: bool = False
# "flow": conditional flow matching (~10 ODE steps at inference).
# "ddpm": cosine-schedule diffusion baseline.
# "wgan": WGAN-GP, single forward pass at inference.
@@ -620,6 +633,8 @@ class Stage1ModelConfig:
d = d or {}
return cls(
active=d.get("active", True),
init_from=d.get("init_from", ""),
freeze=d.get("freeze", False),
generator=d.get("generator", "flow"),
hidden_dim=d.get("hidden_dim", 256),
n_res_blocks=d.get("n_res_blocks", 6),
@@ -636,6 +651,8 @@ class Stage1ModelConfig:
def to_dict(self) -> dict:
return {
"active": self.active,
"init_from": self.init_from,
"freeze": self.freeze,
"generator": self.generator,
"hidden_dim": self.hidden_dim,
"n_res_blocks": self.n_res_blocks,
@@ -655,6 +672,9 @@ class Stage2ModelConfig:
# false trains stage 1 alone. giant rollout must then refuse the
# checkpoint; giant predict still works.
active: bool = True
# See Stage1ModelConfig.init_from/.freeze — same semantics, this stage.
init_from: str = ""
freeze: bool = False
# "one_shot": predict all k_max slots simultaneously with padded slots
# masked from the loss (v0.2 behaviour).
# "autoregressive": emit one secondary at a time in descending-energy
@@ -699,6 +719,8 @@ class Stage2ModelConfig:
d = d or {}
return cls(
active=d.get("active", True),
init_from=d.get("init_from", ""),
freeze=d.get("freeze", False),
decoder=d.get("decoder", "autoregressive"),
generator=d.get("generator", "wgan"),
hidden_dim=d.get("hidden_dim", 256),
@@ -724,6 +746,8 @@ class Stage2ModelConfig:
def to_dict(self) -> dict:
return {
"active": self.active,
"init_from": self.init_from,
"freeze": self.freeze,
"decoder": self.decoder,
"generator": self.generator,
"hidden_dim": self.hidden_dim,
@@ -1142,6 +1166,13 @@ FLAG_SPECS: tuple[FlagSpec, ...] = (
FlagSpec("stage1_critic_n_res_blocks", ("stage1_model.wgan.critic_n_res_blocks",)),
FlagSpec("stage2_critic_hidden_dim", ("stage2_model.wgan.critic_hidden_dim",)),
FlagSpec("stage2_critic_n_res_blocks", ("stage2_model.wgan.critic_n_res_blocks",)),
# Partial-retrain (gitea #42): stage-scoped only, no shared alias — a
# shared "freeze both stages from the same file" flag has no sensible
# meaning (a checkpoint has one set of weights per stage).
FlagSpec("stage1_init_from", ("stage1_model.init_from",)),
FlagSpec("stage1_freeze", ("stage1_model.freeze",)),
FlagSpec("stage2_init_from", ("stage2_model.init_from",)),
FlagSpec("stage2_freeze", ("stage2_model.freeze",)),
)
@@ -1374,7 +1405,7 @@ def merge_cli_overrides(
return cfg
def validate_config(cfg: dict) -> None:
def validate_config(cfg: dict, *, resume: bool = False) -> None:
"""Cross-block validation the per-block schema can't express on its own.
Raises ValueError with a clear message on the first violation found. Call
@@ -1382,6 +1413,10 @@ def validate_config(cfg: dict) -> None:
these checks need to see across blocks, so they don't belong in
`migrate_config` (which only ever sees one dict's own keys) or in any
single block's defaults.
`resume=True` (only `giant train --resume` passes this) relaxes the
`stage{1,2}_model.freeze` -> `.init_from` requirement below: a resumed
frozen stage's weights come from the resume checkpoint, not `init_from`.
"""
particle_type = _get_path(cfg, "conditioning.particle.type")
@@ -1395,6 +1430,13 @@ def validate_config(cfg: dict) -> None:
)
for stage_name in ("stage1_model", "stage2_model"):
if _get_path(cfg, f"{stage_name}.freeze") and not _get_path(cfg, f"{stage_name}.init_from") and not resume:
raise ValueError(
f"{stage_name}.freeze = true requires {stage_name}.init_from "
"to be set (or --resume) — freezing a randomly-initialized "
"model is almost certainly a mistake"
)
router = _get_path(cfg, f"{stage_name}.router") or {}
if router.get("enabled") and router.get("type") in ("pdg", "process") and particle_type == "physical":
raise ValueError(
@@ -1614,6 +1656,8 @@ _OUT_DIR_NAME_CANDIDATES = [
),
("particle_conditioning", _conditioning_candidate("particle", "c")),
("material_conditioning", _conditioning_candidate("material", "m")),
("stage1_freeze", _path_candidate("stage1_model.freeze", "s1frozen", formatter=lambda _: "")),
("stage2_freeze", _path_candidate("stage2_model.freeze", "s2frozen", formatter=lambda _: "")),
("stage1_hidden_dim", _path_candidate("stage1_model.hidden_dim", "h")),
("stage2_hidden_dim", _path_candidate("stage2_model.hidden_dim", "s2h")),
("stage1_n_res_blocks", _path_candidate("stage1_model.n_res_blocks", "b")),
+4
View File
@@ -65,6 +65,10 @@ _STRING_ALTERNATIVES: dict[str, tuple[str, ...]] = {
# and giant/rollout.py while implementing gitea #46 — not auto-derived, so a
# future reader touching these fields should re-check this table still holds.
_NOT_BUILD_TIME: dict[str, str] = {
"stage1_model.init_from": "training/checkpoint.py's init_stages_from_checkpoints, run before build_stage_trainers (gitea #42)",
"stage1_model.freeze": "trainers.py: StageSpec.freeze, gates StageTrainer._step_optimizer (gitea #42)",
"stage2_model.init_from": "training/checkpoint.py's init_stages_from_checkpoints, run before build_stage_trainers (gitea #42)",
"stage2_model.freeze": "trainers.py: StageSpec.freeze, gates StageTrainer._step_optimizer (gitea #42)",
"stage1_model.lambda": "trainers.py: StageSpec.lambda_weight, the total-loss mix weight",
"stage2_model.lambda": "trainers.py: StageSpec.lambda_weight, the total-loss mix weight",
"stage2_model.n_sec.lambda": "trainers.py: StageSpec.n_sec_lambda, the n_sec-head loss weight",
+1 -1
View File
@@ -359,7 +359,7 @@ def run_train_job(
"section)"
)
config.validate_config(cfg)
config.validate_config(cfg, resume=resume is not None)
particle_conditioning = cfg["conditioning"]["particle"]["type"]
material_conditioning = cfg["conditioning"]["material"]["type"]
k_max = cfg["stage2_model"]["k_max"]
+2 -1
View File
@@ -5,7 +5,7 @@ Split out of the former single-module `giant/train.py`. The public surface is
that tests and tooling construct directly.
"""
from giant.training.checkpoint import build_checkpoint, load_checkpoint
from giant.training.checkpoint import build_checkpoint, init_stages_from_checkpoints, load_checkpoint
from giant.training.metrics import MetricsCollector, MetricSpec
from giant.training.loop import train
from giant.training.trainers import (
@@ -25,6 +25,7 @@ __all__ = [
"WGANStageTrainer",
"build_checkpoint",
"build_stage_trainers",
"init_stages_from_checkpoints",
"load_checkpoint",
"train",
]
+33
View File
@@ -8,6 +8,8 @@ and per-stage `optimizer_<stage>` / `optimizer_d_<stage>` / `lr_sched_<stage>`
entries.
"""
import torch
from giant.training.trainers import StageTrainer
#: Stage name -> the checkpoint key its weights live under. Historical: stage
@@ -46,6 +48,37 @@ def build_checkpoint(
return ckpt
def init_stages_from_checkpoints(trainers: dict[str, StageTrainer]) -> list[str]:
"""Load each trainer's `spec.init_from` checkpoint (gitea #42) into its
model, before training starts the partial-retrain counterpart to
`load_checkpoint`'s full-run `--resume`. Only weights move: unlike
`load_checkpoint`, this never touches optimizer/lr_sched/epoch state, so
it composes cleanly with `--resume` (call this first; a resume's own
`load_checkpoint` then overwrites whatever this loaded with the resumed
run's own weights).
A stage with no `init_from` set (`""`, the default) is left alone. The
EMA companion (`<key>_ema`) is loaded too when both the source checkpoint
and this trainer have one, so `--weights ema` at inference still sees the
source's EMA shadow rather than a copy of its raw weights. Returns one
description string per stage actually initialized, for the caller to
echo.
"""
loaded = []
for name, trainer in trainers.items():
init_from = trainer.spec.init_from
if not init_from:
continue
key = _STAGE_KEY[name]
ckpt = torch.load(init_from, map_location="cpu", weights_only=False)
trainer.model.load_state_dict(ckpt[key])
ema_key = f"{key}_ema"
if trainer.ema_model is not None and ema_key in ckpt:
trainer.ema_model.load_state_dict(ckpt[ema_key])
loaded.append(f"{name}: loaded from {init_from}" + (" (frozen)" if trainer.frozen else ""))
return loaded
def load_checkpoint(trainers: dict[str, StageTrainer], ckpt: dict, lr: float) -> None:
"""Restore every active stage, then hand `lr`'s authority back to the
config `load_state_dict` would otherwise leave the checkpoint's own
+3 -1
View File
@@ -20,7 +20,7 @@ from tqdm import tqdm
from giant.data.loader import TopNMap
from giant.data.setup_cache import topnmap_to_json
from giant.training.checkpoint import build_checkpoint, load_checkpoint
from giant.training.checkpoint import build_checkpoint, init_stages_from_checkpoints, load_checkpoint
from giant.training.metrics import MetricsCollector
from giant.training.trainers import (
FlowDDPMStageTrainer,
@@ -135,6 +135,8 @@ def train(
trainers = build_stage_trainers(cfg, models, critics, device, total_train_batches)
if not trainers:
raise ValueError("no active stage — stage1_model.active and stage2_model.active are both false")
for line in init_stages_from_checkpoints(trainers):
print(line)
has_adversarial = any(not tr.supports_val_loss for tr in trainers.values())
checkpoint_extras = {
+23 -8
View File
@@ -89,6 +89,10 @@ class StageSpec:
generator: str
decoder: str = "one_shot"
# partial-retrain (gitea #42)
init_from: str = ""
freeze: bool = False
# loss weights
lambda_weight: float = 1.0
n_sec_lambda: float = 0.1
@@ -151,6 +155,8 @@ class StageSpec:
is_stage2=is_stage2,
generator=stage_spec.generator,
decoder=s2_spec.decoder if is_stage2 else "one_shot",
init_from=stage_spec.init_from,
freeze=stage_spec.freeze,
lambda_weight=stage_spec.lambda_weight,
n_sec_lambda=s2_spec.n_sec.lambda_weight,
n_sec_mode=s2_spec.n_sec.mode,
@@ -244,6 +250,7 @@ class StageTrainer:
self.is_stage2 = spec.is_stage2
self.generator = spec.generator
self.decoder = spec.decoder
self.frozen = spec.freeze
self.device = device
self.model = model.to(device)
self.router = _stage_router(self.model)
@@ -511,14 +518,20 @@ class StageTrainer:
stop_acc = (((logits >= 0).float() == target).float() * mask_f).sum() / denom
return l_stop, stop_acc
@staticmethod
def _step_optimizer(optimizer: optim.Optimizer, loss: torch.Tensor, params: list) -> float:
def _step_optimizer(self, optimizer: optim.Optimizer, loss: torch.Tensor, params: list) -> float:
"""`zero_grad -> backward -> clip_grad_norm_(1.0) -> step`, returning
the pre-clip grad norm. The one place the grad-clip constant lives."""
the pre-clip grad norm. The one place the grad-clip constant lives.
`self.frozen` (`stage{1,2}_model.freeze`, gitea #42) skips only the
final `optimizer.step()` backward/clip still run so loss/grad_norm
stay meaningful to watch, but the stage's weights (and, for a WGAN
stage, its critic's — this same method is both trainers' single
optimizer-step choke point) never move."""
optimizer.zero_grad()
loss.backward()
grad_norm = torch.nn.utils.clip_grad_norm_(params, 1.0)
optimizer.step()
if not self.frozen:
optimizer.step()
return grad_norm.item()
def _extra_state(self) -> dict:
@@ -769,8 +782,9 @@ class FlowDDPMStageTrainer(StageTrainer):
epoch = global_step // self.spec.steps_per_epoch
out = self._compute(batch, device, epoch=epoch)
grad_norm = self._step_optimizer(self.optimizer, out["loss"], self.params)
self.lr_sched.step()
if self.ema_model is not None:
if not self.frozen:
self.lr_sched.step()
if self.ema_model is not None and not self.frozen:
_update_ema(self.ema_model, self.model, self.ema_decay)
stats = {key: value.item() for key, value in out.items()}
stats["grad_norm"] = grad_norm
@@ -1000,8 +1014,9 @@ class WGANStageTrainer(StageTrainer):
grad_norm_g = self._step_optimizer(self.optimizer, g_loss, self.g_params)
if did_g_step:
self.lr_sched.step()
if self.ema_model is not None:
if not self.frozen:
self.lr_sched.step()
if self.ema_model is not None and not self.frozen:
_update_ema(self.ema_model, self.model, self.ema_decay)
return {
+1 -1
View File
@@ -1,6 +1,6 @@
[project]
name = "giant"
version = "0.3.2"
version = "0.3.3"
description = "Geant4 step-function surrogate via conditional flow matching"
readme = "README.md"
requires-python = ">=3.12"
+25
View File
@@ -91,6 +91,31 @@ def test_dry_run_writes_nothing(tmp_path: Path):
assert not out_dir.exists()
def test_stage1_init_from_and_freeze_flags_scaffold_a_partial_retrain_config(tmp_path: Path):
"""gitea #42."""
out_dir = tmp_path / "run5"
result = runner.invoke(
app,
[
"new-run",
"--out",
str(out_dir),
"--stage1-init-from",
"ckpt/stage1_good/best.pt",
"--stage1-freeze",
],
)
assert result.exit_code == 0, result.output
with open(out_dir / "config.toml", "rb") as f:
cfg = tomllib.load(f)
assert cfg["stage1_model"]["init_from"] == "ckpt/stage1_good/best.pt"
assert cfg["stage1_model"]["freeze"] is True
assert cfg["stage2_model"]["init_from"] == ""
assert cfg["stage2_model"]["freeze"] is False
def test_force_guard_refuses_to_clobber_existing_checkpoints(tmp_path: Path):
out_dir = tmp_path / "run5"
out_dir.mkdir()
+27
View File
@@ -97,6 +97,33 @@ def test_wgan_knobs_split_per_stage(monkeypatch, tmp_path):
assert cfg["stage2_model"]["wgan"]["gp_weight"] == 2.5
def test_stage1_init_from_and_freeze_flags_land_in_cfg_and_dont_touch_stage2(monkeypatch, tmp_path):
"""gitea #42: --stage{1,2}-init-from/--stage{1,2}-freeze are stage-scoped
only. --stage1-freeze alone would fail validate_config (freeze requires
init_from or --resume), so both flags are passed together here."""
cfg = _invoke_and_capture_cfg(
monkeypatch,
tmp_path,
["--stage1-init-from", "ckpt/stage1_good/best.pt", "--stage1-freeze"],
)
assert cfg["stage1_model"]["init_from"] == "ckpt/stage1_good/best.pt"
assert cfg["stage1_model"]["freeze"] is True
assert cfg["stage2_model"]["init_from"] == ""
assert cfg["stage2_model"]["freeze"] is False
def test_stage2_init_from_and_freeze_flags_land_in_cfg_and_dont_touch_stage1(monkeypatch, tmp_path):
cfg = _invoke_and_capture_cfg(
monkeypatch,
tmp_path,
["--stage2-init-from", "ckpt/stage2_good/best.pt", "--stage2-freeze"],
)
assert cfg["stage2_model"]["init_from"] == "ckpt/stage2_good/best.pt"
assert cfg["stage2_model"]["freeze"] is True
assert cfg["stage1_model"]["init_from"] == ""
assert cfg["stage1_model"]["freeze"] is False
def test_batch_size_invalid_string_errors(monkeypatch, tmp_path):
monkeypatch.setattr(cli, "run_train_job", lambda *a, **kw: None)
result = runner.invoke(
+51
View File
@@ -97,6 +97,19 @@ def test_trunk_config_defaults_block_conditioning_to_add_for_both_stages():
assert gconfig.DEFAULT_CONFIG["stage2_model"]["trunk"]["block_conditioning"] == "add"
def test_init_from_freeze_default_to_unset_for_both_stages():
"""gitea #42: a pre-existing config with no init_from/freeze key must
reproduce today's from-scratch, always-training behaviour exactly."""
assert gconfig.Stage1ModelConfig().init_from == ""
assert gconfig.Stage1ModelConfig().freeze is False
assert gconfig.Stage2ModelConfig().init_from == ""
assert gconfig.Stage2ModelConfig().freeze is False
assert gconfig.DEFAULT_CONFIG["stage1_model"]["init_from"] == ""
assert gconfig.DEFAULT_CONFIG["stage1_model"]["freeze"] is False
assert gconfig.DEFAULT_CONFIG["stage2_model"]["init_from"] == ""
assert gconfig.DEFAULT_CONFIG["stage2_model"]["freeze"] is False
def test_heads_config_defaults_reproduce_pre_gitea_36_hardcoded_shape():
"""gitea #36: a pre-existing config with no `heads` key must reproduce
today's hardcoded `hidden_dim // 2`, one-hidden-layer architecture
@@ -735,6 +748,29 @@ def test_validate_config_tie_to_stage1_requires_stage1_active():
assert "tie_to_stage1" in str(e)
@pytest.mark.parametrize("stage_name", ["stage1_model", "stage2_model"])
def test_validate_config_freeze_without_init_from_or_resume_rejected(stage_name):
cfg = _cfg_with(**{f"{stage_name}.freeze": True})
try:
gconfig.validate_config(cfg)
assert False, "expected ValueError"
except ValueError as e:
assert "init_from" in str(e)
assert "--resume" in str(e)
@pytest.mark.parametrize("stage_name", ["stage1_model", "stage2_model"])
def test_validate_config_freeze_with_init_from_passes(stage_name):
cfg = _cfg_with(**{f"{stage_name}.freeze": True, f"{stage_name}.init_from": "ckpt/best.pt"})
gconfig.validate_config(cfg) # must not raise
@pytest.mark.parametrize("stage_name", ["stage1_model", "stage2_model"])
def test_validate_config_freeze_without_init_from_passes_under_resume(stage_name):
cfg = _cfg_with(**{f"{stage_name}.freeze": True})
gconfig.validate_config(cfg, resume=True) # must not raise
def test_validate_config_stop_token_accepted_under_autoregressive():
"""DEFAULT_CONFIG's stage2_model.decoder is already "autoregressive"
(see test_stage2_model_config_defaults_match_documented_v030_intent), so
@@ -1259,6 +1295,21 @@ def test_overrides_from_flags_critic_sizing_is_stage_scoped_only(stage_flag, sta
assert overrides == {stage_model: {"wgan": {path_key: 32}}}
@pytest.mark.parametrize(
("init_from_flag", "freeze_flag", "stage_model"),
[
("stage1_init_from", "stage1_freeze", "stage1_model"),
("stage2_init_from", "stage2_freeze", "stage2_model"),
],
)
def test_overrides_from_flags_init_from_freeze_is_stage_scoped_only(init_from_flag, freeze_flag, stage_model):
"""gitea #42: no shared alias — a checkpoint has one set of weights per
stage, so "freeze both stages from the same file" has no sensible
meaning."""
overrides = gconfig.overrides_from_flags({init_from_flag: "ckpt/best.pt", freeze_flag: True})
assert overrides == {stage_model: {"init_from": "ckpt/best.pt", "freeze": True}}
# ---------------------------------------------------------------------------
# checkpoint config-mismatch warnings (unchanged surface, still exercised)
# ---------------------------------------------------------------------------
+169 -2
View File
@@ -7,6 +7,7 @@ import tempfile
from pathlib import Path
from unittest.mock import MagicMock, patch
import numpy as np
import pytest
import torch
@@ -19,14 +20,18 @@ from giant.constants import (
SEC_SLOT_DIM,
X_DIM,
)
from giant.checkpoint_io import load_for_inference
from giant.data.dataset import StepBatch
from giant.data.transforms import Normalizer
from giant.model.network import Stage2Autoregressive, build_critics, build_models
from giant.sample import sample_stage1 as trainers_sample_stage1
from giant.training import (
FlowDDPMStageTrainer,
StageSpec,
WGANStageTrainer,
build_checkpoint,
build_stage_trainers,
init_stages_from_checkpoints,
train,
)
from giant.training.metrics import _wandb_run_config
@@ -352,7 +357,7 @@ def _model_config(cfg):
}
def _run_train(cfg, out_dir, resume_path=None):
def _run_train(cfg, out_dir, resume_path=None, normalizer_dict=None):
model_config = _model_config(cfg)
models = build_models(model_config)
critics = build_critics(model_config)
@@ -366,7 +371,7 @@ def _run_train(cfg, out_dir, resume_path=None):
val_loader=val_loader,
device=torch.device("cpu"),
out_dir=out_dir,
normalizer_dict={"cond": {}, "target": {}, "sec_phys": {}},
normalizer_dict=normalizer_dict or {"cond": {}, "target": {}, "sec_phys": {}},
pdg_map={"22": 0},
mat_map={"G4_AIR": 0},
proc_map=None,
@@ -594,6 +599,168 @@ def test_flow_stage_trainer_ddpm_not_implemented_for_stage2():
FlowDDPMStageTrainer(spec, torch.nn.Linear(1, 1), torch.device("cpu"))
# --- gitea #42: freeze / init_from -------------------------------------------
def _state_dict_clone(module):
return {k: v.clone() for k, v in module.state_dict().items()}
def _assert_state_dicts_equal(before, after, label):
for key, value in before.items():
assert torch.equal(value, after[key]), f"{label}: {key} changed while frozen"
def test_frozen_flow_stage_trainer_step_does_not_update_model_or_ema():
cfg = _base_cfg()
model_config = _model_config(cfg)
model = build_models(model_config)["stage1"]
assert model is not None
spec = StageSpec(name="stage1", is_stage2=False, generator="flow", freeze=True, ema_decay=0.999, steps_per_epoch=4)
trainer = FlowDDPMStageTrainer(spec, model, torch.device("cpu"))
assert trainer.ema_model is not None
model_before = _state_dict_clone(trainer.model)
ema_before = _state_dict_clone(trainer.ema_model)
for batch in _fake_batches(4, 8):
trainer.step(batch, torch.device("cpu"), global_step=1)
_assert_state_dicts_equal(model_before, trainer.model.state_dict(), "frozen flow model")
_assert_state_dicts_equal(ema_before, trainer.ema_model.state_dict(), "frozen flow ema")
def test_frozen_wgan_stage_trainer_step_does_not_update_generator_or_critic():
cfg = _base_cfg()
cfg["stage1_model"]["generator"] = "wgan"
model_config = _model_config(cfg)
models = build_models(model_config)
critics = build_critics(model_config)
assert models["stage1"] is not None and critics["stage1"] is not None
spec = StageSpec(
name="stage1",
is_stage2=False,
generator="wgan",
freeze=True,
n_critic=1, # a generator step every batch, so a bug would surface immediately
ema_decay=0.999,
steps_per_epoch=4,
)
trainer = WGANStageTrainer(spec, models["stage1"], critics["stage1"], torch.device("cpu"))
assert trainer.ema_model is not None
model_before = _state_dict_clone(trainer.model)
critic_before = _state_dict_clone(trainer.critic)
ema_before = _state_dict_clone(trainer.ema_model)
for global_step, batch in enumerate(_fake_batches(4, 8)):
trainer.step(batch, torch.device("cpu"), global_step=global_step)
_assert_state_dicts_equal(model_before, trainer.model.state_dict(), "frozen wgan generator")
_assert_state_dicts_equal(critic_before, trainer.critic.state_dict(), "frozen wgan critic")
_assert_state_dicts_equal(ema_before, trainer.ema_model.state_dict(), "frozen wgan ema")
@pytest.mark.parametrize("stage1_generator", ["flow", "wgan"])
def test_train_end_to_end_frozen_stage1_unchanged_while_stage2_trains(stage1_generator):
cfg = _base_cfg()
cfg["stage1_model"]["generator"] = stage1_generator
cfg["stage1_model"]["freeze"] = True
model_config = _model_config(cfg)
models = build_models(model_config)
critics = build_critics(model_config)
assert models["stage1"] is not None and models["stage2"] is not None
stage1_before = _state_dict_clone(models["stage1"])
stage2_before = _state_dict_clone(models["stage2"])
with tempfile.TemporaryDirectory() as tmp:
train(
cfg=cfg,
models=models,
critics=critics,
train_loader=_fake_batches(4, cfg["train"]["batch_size"]),
val_loader=_fake_batches(2, cfg["train"]["batch_size"], seed=1),
device=torch.device("cpu"),
out_dir=Path(tmp) / "run",
normalizer_dict={"cond": {}, "target": {}, "sec_phys": {}},
pdg_map={"22": 0},
mat_map={"G4_AIR": 0},
proc_map=None,
model_config=model_config,
total_train_batches=4,
)
_assert_state_dicts_equal(stage1_before, models["stage1"].state_dict(), "frozen stage1")
stage2_after = models["stage2"].state_dict()
assert any(not torch.equal(v, stage2_after[k]) for k, v in stage2_before.items()), (
"unfrozen stage2 should have trained"
)
def test_init_stages_from_checkpoints_loads_matching_stage_and_ema_weights(tmp_path):
cfg = _base_cfg()
model_config = _model_config(cfg)
source_models = build_models(model_config)
source_critics = build_critics(model_config)
source_trainers = build_stage_trainers(cfg, source_models, source_critics, torch.device("cpu"), 4)
source_stage1_ema = source_trainers["stage1"].ema_model
assert source_stage1_ema is not None
# Diverge the source's EMA from its raw weights so a same-vs-different
# check below actually distinguishes the two copy paths.
for p in source_stage1_ema.parameters():
p.data.add_(1.0)
ckpt_path = tmp_path / "source.pt"
ckpt = build_checkpoint(source_trainers, epoch=1, global_step=1, best_val_loss=0.0, extras={})
torch.save(ckpt, ckpt_path)
cfg2 = copy.deepcopy(cfg)
cfg2["stage1_model"]["init_from"] = str(ckpt_path)
dest_models = build_models(_model_config(cfg2))
dest_critics = build_critics(_model_config(cfg2))
dest_trainers = build_stage_trainers(cfg2, dest_models, dest_critics, torch.device("cpu"), 4)
loaded = init_stages_from_checkpoints(dest_trainers)
assert len(loaded) == 1 and "stage1" in loaded[0]
dest_stage1_ema = dest_trainers["stage1"].ema_model
assert dest_stage1_ema is not None
_assert_state_dicts_equal(
source_trainers["stage1"].model.state_dict(), dest_trainers["stage1"].model.state_dict(), "init_from raw"
)
_assert_state_dicts_equal(
source_stage1_ema.state_dict(),
dest_stage1_ema.state_dict(),
"init_from ema",
)
# stage2 has no init_from set -- untouched fresh init, not the source's.
stage2_matches_source = all(
torch.equal(v, dest_trainers["stage2"].model.state_dict()[k])
for k, v in source_trainers["stage2"].model.state_dict().items()
)
assert not stage2_matches_source
def test_run_train_job_stage1_init_from_freeze_produces_rollout_capable_checkpoint(tmp_path):
"""The exact scenario gitea #42 exists for: retrain stage 2 alone against
a fixed, known-good stage 1, and still get a checkpoint giant rollout can
load (checkpoint_io.load_for_inference with require_stage2=True)."""
normalizer_dict = {
"cond": Normalizer().fit(np.zeros((1, COND_DIM), dtype=np.float32)).to_dict(),
"target": Normalizer().fit(np.zeros((1, X_DIM), dtype=np.float32)).to_dict(),
"sec_phys": Normalizer().fit(np.zeros((1, 2), dtype=np.float32)).to_dict(),
}
cfg = _base_cfg()
source_out = tmp_path / "source"
_run_train(cfg, source_out, normalizer_dict=normalizer_dict)
source_ckpt = torch.load(source_out / "best.pt", weights_only=False)
cfg2 = copy.deepcopy(cfg)
cfg2["stage1_model"]["init_from"] = str(source_out / "best.pt")
cfg2["stage1_model"]["freeze"] = True
retrain_out = tmp_path / "retrain"
_run_train(cfg2, retrain_out, normalizer_dict=normalizer_dict)
ctx = load_for_inference(retrain_out / "best.pt", torch.device("cpu"), "rollout", require_stage2=True)
assert ctx.stage1 is not None and ctx.stage2 is not None
retrain_ckpt = torch.load(retrain_out / "best.pt", weights_only=False)
for key, value in source_ckpt["model"].items():
assert torch.equal(value, retrain_ckpt["model"][key]), f"frozen stage1 {key} drifted across the retrain"
def test_stage_spec_from_config_omitted_decoder_and_particle_type_match_default_config():
"""Regression for issues.md Issue 1: StageSpec.from_config's own fallback
defaults for stage2_model.decoder/particle_type must equal
Generated
+1 -1
View File
@@ -633,7 +633,7 @@ wheels = [
[[package]]
name = "giant"
version = "0.3.2"
version = "0.3.3"
source = { editable = "." }
dependencies = [
{ name = "numpy" },