diff --git a/CLAUDE.md b/CLAUDE.md index dc1a500..ccab334 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -10,6 +10,7 @@ uv sync --extra cuda # install dependencies with CUDA 11.8 torch uv sync --extra cpu --extra dev # add dev extras (pytest, etc.) uv sync --extra cpu --extra geometry # add scikit-learn for the geometry oracle (giant rollout) pytest # run tests +giant new-run --hidden-dim 512 --lr 3e-4 # scaffold a config.toml + run dir ahead of training giant train path/to/steps.parquet --mode flow # train (flow matching) giant train path/to/steps.parquet --mode ddpm # train (DDPM baseline) giant train path/to/steps.parquet --mode wgan # train (WGAN-GP, single-pass eval; implemented, not yet tested) diff --git a/README.md b/README.md index 37beb00..02c395f 100644 --- a/README.md +++ b/README.md @@ -2,15 +2,19 @@ **G**eant4 **I**nference via **A**utoregressive **N**eural s**T**ep surrogate — a play on *Geant4* and the step function being the computationally heaviest part of the simulation. -Proof-of-concept surrogate model for the Geant4 step function. Given a pre-step particle state, the model samples a physically plausible post-step outcome — the primary's continuation plus the secondary particles it produces — replacing the stochastic Geant4 physics engine with a trained conditional generative model. +Conditional generative surrogate for the Geant4 step function. Given a pre-step particle state, the model samples a physically plausible post-step outcome — the primary's continuation plus the variable-length list of secondary particles it produces — replacing the stochastic Geant4 physics engine with a trained generative model. A trained checkpoint autoregressively rolls out full showers, stepping each primary and pushing secondaries as new tracks. Training is driven entirely from parquet files of the miniCaloSim steps tree. No Geant4 runtime dependency. ## Architecture -A **two-stage conditional flow matching** model (Lipman et al. 2022): a small MLP learns a vector field mapping noise → step outcomes in ~10 ODE steps per sample. Falls back to DDPM for comparison. +A **two-stage model**, both stages checkpointed together, with a choice of generative mode per stage (`--mode`): -**Stage 1 — primary (9D, diffused):** +- **`flow`** (default) — conditional flow matching (Lipman et al. 2022): an MLP learns a vector field mapping noise → step outcomes, sampled via ODE integration in ~10 steps. +- **`ddpm`** — a standard denoising diffusion baseline for comparison (`giant/model/schedule.py:CosineSchedule`). +- **`wgan`** — a single-pass Wasserstein-GAN-GP generator/critic (`giant/model/wgan.py`), trading iterative sampling for one forward pass; implemented, not yet validated against the flow-matching baseline. + +**Stage 1 — primary (9D, `giant/constants.py:LOCAL_TARGET_NAMES`):** | Index | Variable | Encoding | |-------|----------|----------| @@ -19,13 +23,20 @@ A **two-stage conditional flow matching** model (Lipman et al. 2022): a small ML | 3–5 | `post_dir` in local frame | unit vector | | 6–8 | `travel_dir` (`post_pos − pre_pos`) in local frame | unit vector | -The two energy logits decode via softmax over `[edep_logit, sec_logit, 0]` × `pre_E`, so `edep + e_sec + post_E == pre_E` exactly — **energy conservation is built into the parametrization**, not left to the loss. Stage 1 also has a classifier head predicting the number of secondaries `n_sec ∈ {0..15}` from the conditioning alone. +The two energy logits decode via softmax over `[edep_logit, sec_logit, 0]` × `pre_E`, so `edep + e_sec + post_E == pre_E` exactly — **energy conservation is built into the parametrization**, not left to the loss (`energy_simplex_decode`). Stage 1 also has a classifier head (`predict_n_sec`) predicting the number of secondaries `n_sec ∈ {0..K_MAX}` (`K_MAX = 15`) from the conditioning alone, no diffusion noise involved. Both `post_dir` and `travel_dir` are expressed in the coordinate frame where `pre_dir = ẑ`, making the scattering distribution nearly azimuthally symmetric. `post_pos` itself is not a raw target — it's reconstructed at inference as `pre_pos + step_length * world_frame(travel_dir)`, so the two stay consistent by construction instead of being learned (and potentially diverging) independently. -**Stage 2 — secondaries (`SecondaryDecoder`):** conditioned on the pre-step state *and* the Stage-1 outcome, a second flow net generates all `K_MAX = 15` secondary slots at once. Each slot carries a stick-breaking energy fraction, a local-frame direction, and a continuous particle-type embedding (snapped to the nearest PDG at inference), ordered by descending energy; slots beyond the predicted `n_sec` are masked. The secondary energies are a stick-breaking partition of the `e_sec` budget from Stage 1, so the full chain conserves energy. Each secondary's momentum is reconstructed afterward from `(energy, direction, species)` rather than predicted. +**Stage 2 — secondaries (`SecondaryDecoder`):** conditioned on the pre-step state *and* the Stage-1 outcome, a second net generates all `K_MAX` secondary slots at once — `(stick-breaking energy logit, local-frame direction, log-mass, charge)` per slot, ordered by descending energy; slots beyond the predicted `n_sec` are masked. Secondary energies are a stick-breaking partition of the `e_sec` budget from Stage 1, so the whole chain conserves energy. A secondary's mass/charge are regressed directly against its ground-truth PDG code's physical values (`giant.particles.particle_mass_charge`) and used as-is at inference — including for its own conditioning if it takes further steps in a rollout. No snapping to a known PDG code happens in the model path; `giant.particles.nearest_known_pdg` is a reporting-only lookup used to populate a nominal `pdg` label on output rows. -**Conditioning:** PDG code (embedding), pre-step position, log(pre-energy), pre-step direction, material (embedding), layer ID. (`n_sec` / `e_sec` are outputs now, not inputs.) +**Conditioning (`--conditioning`, per-checkpoint):** pre-step position, log(pre-energy), pre-step direction, layer ID, plus particle/material physical properties — mass/charge (`giant/particles.py`) and Z_eff/A_eff/density/X0/λ_int (`giant/materials.py`). Two mutually exclusive modes: + +- **`physical`** (default) — the physical-property columns are routed through small MLPs, computable for any PDG code / material, letting the surrogate generalize to species/materials outside the training menu. +- **`embedding`** — the original design: a learned `nn.Embedding` per PDG code / material, kept as a generalization-comparison baseline (memorizes the training menu). + +`n_sec` and `e_sec` are model outputs, not conditioning inputs — a rollout is self-contained and never injects ground truth. + +**Mixture-of-experts routing (`--router`, opt-in):** `giant/model/network.py` also implements a pluggable `Router` contract (`ROUTER_REGISTRY`: `energy`, `pdg`, `process`, plus a `composed` router combining several axes) that splits `DenoisingMLP`/`SecondaryDecoder` into per-expert trunks, soft-gated in training and top-1 dispatched at eval. Implemented; first rollout benchmark needs a retrain with a load-balancing loss and better-seeded router centers (see Roadmap). See `--router-type`/`--n-experts`/`--router-axis` on `giant train`/`giant new-run`. ## Roadmap @@ -33,11 +44,15 @@ Both `post_dir` and `travel_dir` are expressed in the coordinate frame where `pr **Phase 2 (implemented — baseline):** the two-stage model above predicts `n_sec` and each secondary's energy, direction, and species jointly with the primary, so a shower rollout is fully self-contained. -**Next directions:** faster-eval architectures against a ~10× native-Geant4 budget (Wasserstein-GAN, mixture-of-experts routing tree), a multi-material sampling-calorimeter dataset, and physical-property conditioning over learned embeddings. +**Physical-property conditioning (implemented):** replaces learned PDG/material embeddings with physical-property MLPs (see above); Stage 2 predicts a secondary's mass/charge directly instead of a snapped species embedding. Not yet done: the held-out-material/species generalization comparison against the `embedding` baseline — the natural dataset for that is the 34GB multi-material dataset at the repo root (6 materials, 237 PDG codes). + +**Faster-eval architectures (implemented, validation in progress):** both target a ~10× native-Geant4 eval budget. WGAN-GP (`--mode wgan`) has no rollout-vs-reference analysis run against it yet. The MoE router (`--router`) had its first rollout benchmark diverge from Geant4 despite matching bulk deposited energy — the experts weren't specializing (near-uniform gating), traced to a missing load-balance loss and a center-init that didn't match the real energy distribution; both are now fixable via `lambda_balance > 0` and quantile-seeded router centers, but a re-run to confirm hasn't happened yet. + +A multi-material sampling-calorimeter dataset is a planned future direction, not yet built. ## Data -Input: parquet files produced by [miniCaloSim](https://gitlab.etp.kit.edu/lbogner/minicalosim), or converted from a ROOT file via `uv run dwarf convert`. Each row is one Geant4 step. Train/val split is by `event_id` (not row shuffle) to avoid leaking correlated steps from the same shower. +Input: parquet files produced by [miniCaloSim](https://gitlab.etp.kit.edu/lbogner/minicalosim), or converted from a ROOT file via `dwarf convert`. Each row is one Geant4 step. Train/val split is by `event_id` (not row shuffle, and `--seed`-controlled) to avoid leaking correlated steps from the same shower. ## Project structure @@ -49,21 +64,32 @@ giant/ │ │ ├── transforms.py # log transforms, local-frame rotation, energy simplex, secondary encode/decode │ │ └── dataset.py # StepsDataset / StreamingStepsDataset (PyTorch) │ ├── model/ -│ │ ├── network.py # SinusoidalEmbedding, ConditionEncoder, DenoisingMLP, SecondaryDecoder -│ │ └── schedule.py # CosineSchedule (DDPM) and flow matching utilities +│ │ ├── network.py # ConditionEncoder, DenoisingMLP, SecondaryDecoder, Router/MoE, WGAN generator/critic +│ │ ├── schedule.py # CosineSchedule (DDPM) and flow matching utilities +│ │ └── wgan.py # WGAN-GP gradient penalty / critic / generator losses │ ├── constants.py # output/conditioning dims, K_MAX, secondary slot layout, schema keys +│ ├── particles.py # PDG → (mass, charge) decode, incl. nuclear/ion codes; nearest-known-PDG lookup +│ ├── materials.py # material name → (Z_eff, A_eff, density, X0, λ_int) │ ├── config.py # default hyperparameters, TOML config merging, device autodetect -│ ├── pipeline.py # builds datasets/normalizers and kicks off a training run -│ ├── train.py # two-stage training loop, checkpointing, graceful shutdown -│ ├── sample.py # DDPM / DDIM / flow matching samplers + secondary sampling +│ ├── pipeline.py # builds datasets/normalizers and kicks off a training run (with setup-stage caching) +│ ├── train.py # two-stage training loop, checkpointing, graceful shutdown, W&B logging +│ ├── sample.py # DDPM / DDIM / flow matching / WGAN samplers + secondary sampling │ ├── geometry.py # GeometryOracle: position → (material, layer_id, escaped) for rollout │ ├── rollout.py # autoregressive shower rollout driver │ ├── validate.py # step-level marginal + KL-divergence validation -│ ├── analysis.py # step- and shower-level diagnostics: marginals, correlations, rollout observables -│ └── cli.py # `giant train` / `predict` / `rollout` Typer app -├── scripts/ # dataset/tooling logic, unified under the `dwarf` CLI (`uv run dwarf --help`) +│ ├── analysis/ # rollout-vs-reference analysis pipeline (see `giant analyze` below) +│ │ ├── sources.py # canonical LazyFrames + secondary view +│ │ ├── reduce.py # streaming reduction primitives (hist1d, per-event scalars, profiles, ...) +│ │ ├── grouping.py # fixed bin edges + energy/pdg/material group sets +│ │ ├── context.py # resolves grouping into `shared.json` once per run +│ │ ├── catalog.py # declarative PlotSpec registry +│ │ ├── condor.py # prep / compute-one / submit-description plumbing +│ │ └── render.py # PDFs + HTML gallery (only module importing plotstyle/LaTeX) +│ └── cli.py # `giant train` / `new-run` / `predict` / `rollout` / `analyze` Typer app +├── scripts/ # dataset/tooling logic, unified under the `dwarf` CLI (`dwarf --help`) │ ├── dwarf.py # Typer app: convert, migrate, bump-gen, bump-schema, status, -│ │ # update-manifest, create-manifest, make-root, hparam-scan +│ │ # update-manifest, create-manifest, make-root, +│ │ # build-geometry-oracle, warm-cache, hparam-scan │ ├── steps_to_parquet.py # ROOT → parquet conversion (uproot/awkward/polars) — `dwarf convert` │ ├── steps_to_parquet_parallel.py # fan out conversion over several ROOT files — `dwarf convert --jobs N` │ ├── migrate_geant_steps.py # one-time move into the raw/processed/pools/derived layout — `dwarf migrate` @@ -71,7 +97,9 @@ giant/ │ │ # `dwarf bump-gen` / `bump-schema` / `status` / `update-manifest` / `create-manifest` │ ├── create_root_files.py # generate new ROOT shards via a minicalosim executable — `dwarf make-root` │ ├── geometry_oracle.py # fit a position → (material, layer_id) oracle — `dwarf build-geometry-oracle` -│ └── hparam_scan.py # hyperparameter grid scan over `giant train` runs — `dwarf hparam-scan` +│ ├── warm_setup_cache.py # precompute `giant train`'s setup-stage sidecar — `dwarf warm-cache` +│ ├── hparam_scan.py # hyperparameter grid scan over `giant train` runs — `dwarf hparam-scan` +│ └── profile_analysis_costs.py # profiling helper for the `giant analyze` reduction pipeline └── tests/ ``` @@ -80,27 +108,37 @@ giant/ ```bash uv sync --extra cpu # CPU-only torch (use --extra cuda for CUDA 11.8 instead) uv sync --extra cpu --extra dev # add dev tools (pytest, ruff, ty) +uv sync --extra cpu --extra geometry # add scikit-learn, for `dwarf build-geometry-oracle` / rollout ``` -`cpu` and `cuda` are mutually exclusive extras selecting the torch build; plain `uv sync` installs no torch at all. See `CLAUDE.md` for details. +`cpu` and `cuda` are mutually exclusive extras selecting the torch build (pinned to 2.3.x); plain `uv sync` installs no torch at all. See `CLAUDE.md` for details. ## Training, prediction, and rollout ```bash -giant train path/to/steps.parquet --mode flow +giant new-run --hidden-dim 512 --lr 3e-4 --comment "..." # scaffold a config.toml + run dir for a new run +giant train path/to/steps.parquet --mode flow # train (flow matching; also --mode ddpm / wgan) giant predict path/to/steps.parquet --checkpoint checkpoints/.../best.pt # Full-shower rollout needs a geometry oracle (position → material/layer_id): -uv sync --extra cpu --extra geometry dwarf build-geometry-oracle path/to/steps.parquet --out oracle.pkl giant rollout path/to/steps.parquet --checkpoint checkpoints/.../best.pt --geometry oracle.pkl ``` -`train`/`predict` accept a TOML config file (`--config`) and CLI overrides for hyperparameters; see `--help` on any command for the full option list. `giant train --wandb` logs per-epoch metrics (the same ones written to `metrics.csv`) to Weights & Biases; requires `uv sync --extra wandb`. `giant rollout` seeds showers from the highest-energy entry step per event, then autoregressively steps the two-stage model to completion — pushing secondaries as new tracks and looking up `material`/`layer_id` from the oracle at each step. Tracks terminate on energy cutoff, per-track max steps, detector escape, or natural end; energy is deposited locally on every stop except escape (leakage), so showers conserve energy by construction. +`train`/`predict` accept a TOML config file (`--config`) and CLI overrides for hyperparameters; see `--help` on any command for the full option list. `giant train --wandb` logs per-epoch metrics (the same ones written to `metrics.csv`) to Weights & Biases; requires `uv sync --extra wandb`. A repeat `giant train` against the same dataset (e.g. a hyperparameter sweep) reuses a cached setup-stage sidecar (vocab maps, event split, normalizer stats) unless `--no-cache-setup`/`--rebuild-setup-cache`; `dwarf warm-cache` precomputes it ahead of time. `giant rollout` seeds showers from the highest-energy entry step per event, then autoregressively steps the two-stage model to completion — pushing secondaries as new tracks and looking up `material`/`layer_id` from the oracle at each step. Tracks terminate on energy cutoff, per-track max steps, detector escape, or natural end; energy is deposited locally on every stop except escape (leakage), so showers conserve energy by construction. -## Validation +## Validation and analysis -`giant.validate.validate_marginals` runs step-level marginal and KL-divergence checks during training (`--validate-every`). For deeper diagnostics on a trained checkpoint — stratified marginals, correlation structure, physical-constraint violations, and shower-level rollout observables (longitudinal/transverse profiles, PDG energy shares) — see `giant.analysis`. +`giant.validate.validate_marginals` runs step-level marginal and KL-divergence checks during training (`--validate-every`). + +For deeper rollout-vs-reference diagnostics — marginals stratified by energy/pdg/material, per-event totals, shower profiles, species share, leakage, and secondaries — `giant analyze` runs a streaming compute/render pipeline against a `giant rollout` YAML sidecar: + +```bash +giant analyze submit rollout.yaml --accounting-group cms # prep + one HTCondor job per plot (compute only) +giant analyze render --gallery # local: styled PDFs + HTML gallery (needs LaTeX) +``` + +`` is derived next to the rollout parquet (`analyze prep`/`submit` print it). Compute jobs are polars/numpy only; only `render` imports plotstyle/LaTeX, so it always runs locally. ## Development diff --git a/giant/cli.py b/giant/cli.py index a759861..ea40946 100644 --- a/giant/cli.py +++ b/giant/cli.py @@ -134,6 +134,30 @@ def _parse_router_axis_flags(specs: list[str]) -> dict[str, object]: return out +def _router_cli_overrides( + router: bool | None, + router_type: str | None, + n_experts: int | None, + router_axis: list[str] | None, +) -> dict[str, object]: + """Build the `model.router` override dict from `--router`/`--router-type`/ + `--n-experts`/`--router-axis` flags (empty if none were given). Shared by + `train` and `new-run` so both resolve router overrides identically. + """ + cli_router: dict[str, object] = { + k: v + for k, v in { + "enabled": router, + "type": router_type, + "n_experts": n_experts, + }.items() + if v is not None + } + if router_axis: + cli_router.update(_parse_router_axis_flags(router_axis)) + return cli_router + + _CEPH_PREDICTIONS = Path("/ceph/lbogner/geant_steps/predictions") @@ -504,17 +528,7 @@ def train( }.items() if v is not None } - cli_router: dict[str, object] = { - k: v - for k, v in { - "enabled": router, - "type": router_type, - "n_experts": n_experts, - }.items() - if v is not None - } - if router_axis: - cli_router.update(_parse_router_axis_flags(router_axis)) + cli_router = _router_cli_overrides(router, router_type, n_experts, router_axis) if cli_router: cli_model["router"] = cli_router cfg = gconfig.merge_cli_overrides( @@ -552,13 +566,9 @@ def train( # Name only encodes what's non-default (see default_out_dir_name), so # two runs with identical hyperparams in the same to-the-minute # timestamp would otherwise collide on this name — which also - # doubles as the W&B run id (giant.train) — hence the suffix loop. - base_name = gconfig.default_out_dir_name(cfg) - out_dir = Path("checkpoints") / base_name - suffix = 2 - while out_dir.exists(): - out_dir = Path("checkpoints") / f"{base_name}_{suffix}" - suffix += 1 + # doubles as the W&B run id (giant.train) — hence the suffix loop in + # resolve_default_out_dir. + out_dir = gconfig.resolve_default_out_dir(cfg) typer.echo(f"device: {_device}") typer.echo(f"out_dir: {out_dir}") @@ -577,6 +587,144 @@ def train( ) +@app.command("new-run") +def new_run( + config: Annotated[ + Optional[Path], + typer.Option( + "--config", + "-c", + help="Base TOML to start from (default: built-in defaults)", + ), + ] = None, + mode: Annotated[Optional[Mode], typer.Option("--mode", "-m")] = None, + epochs: Annotated[Optional[int], typer.Option("--epochs", "-e")] = None, + batch_size: Annotated[Optional[int], typer.Option("--batch-size", "-b")] = None, + lr: Annotated[Optional[float], typer.Option("--lr", "-l")] = None, + hidden_dim: Annotated[Optional[int], typer.Option("--hidden-dim", "-H")] = None, + n_blocks: Annotated[Optional[int], typer.Option("--n-blocks", "-n")] = None, + emb_dim: Annotated[Optional[int], typer.Option("--emb-dim", "-E")] = None, + dropout: Annotated[Optional[float], typer.Option("--dropout", "-d")] = 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, + n_experts: Annotated[Optional[int], typer.Option("--n-experts")] = None, + router_axis: Annotated[Optional[list[str]], typer.Option("--router-axis")] = None, + out: Annotated[ + Optional[Path], + typer.Option("--out", "-o", help="Run dir (default: auto from hyperparams)"), + ] = None, + comment: Annotated[ + Optional[str], + typer.Option( + "--comment", help="Free-text note recorded in config.toml's meta section" + ), + ] = None, + data: Annotated[ + Optional[Path], + typer.Option( + "--data", + help="Dataset path to fill in the printed next-step command " + "(not stored in the config)", + ), + ] = None, + force: Annotated[ + bool, + typer.Option( + "--force", + help="Overwrite config.toml even if --out already has checkpoints", + ), + ] = False, + dry_run: Annotated[ + bool, + typer.Option( + "--dry-run", help="Print the resolved config without writing anything" + ), + ] = False, +) -> None: + """Scaffold a new training run: resolve hyperparams to a config.toml and lay out its run dir. + + This is the config-file-first counterpart to hand-editing a TOML: start + from a base --config (or built-in defaults), override a few hyperparams + inline, and this resolves+writes the full `config.toml` into a fresh (or + explicit --out) run dir — the same file `giant train --config ...` reads. + `giant train` itself overwrites this file in place once it actually runs + (with the full dataset-derived meta section), so this scaffold's meta + section is just a placeholder recording what was asked for and when. + """ + cli_train = { + k: v + for k, v in { + "mode": mode.value if mode is not None else None, + "epochs": epochs, + "batch_size": batch_size, + "lr": lr, + }.items() + if v is not None + } + cli_model: dict[str, object] = { + k: v + for k, v in { + "hidden_dim": hidden_dim, + "n_blocks": n_blocks, + "emb_dim": emb_dim, + "dropout": dropout, + "conditioning": conditioning.value if conditioning is not None else None, + }.items() + if v is not None + } + cli_router = _router_cli_overrides(router, router_type, n_experts, router_axis) + if cli_router: + cli_model["router"] = cli_router + + cfg = gconfig.merge_cli_overrides( + gconfig.DEFAULT_CONFIG, config, cli_train, cli_model + ) + run_dir = (out or gconfig.resolve_default_out_dir(cfg)).resolve() + + if not force: + existing = [n for n in ("last.pt", "best.pt") if (run_dir / n).exists()] + if existing: + typer.echo( + f"error: {run_dir} already has {', '.join(existing)} — pass " + "--force to overwrite its config.toml anyway", + err=True, + ) + raise typer.Exit(1) + + typer.echo(f"run dir: {run_dir}") + + if dry_run: + typer.echo("dry-run: not writing anything. Resolved config:") + for section in ("train", "model"): + typer.echo(f"[{section}]") + for k, v in cfg[section].items(): + if k == "router": + continue + typer.echo(f" {k} = {v}") + return + + meta = { + "git_hash": gconfig.git_hash(), + "created_at": datetime.now(timezone.utc).isoformat(timespec="seconds"), + "created_by": "giant new-run", + } + if comment: + meta["comment"] = comment + + run_dir.mkdir(parents=True, exist_ok=True) + gconfig.save_config(cfg, run_dir, meta) + config_path = run_dir / "config.toml" + typer.echo(f"wrote {config_path}") + + data_arg = str(data) if data is not None else "" + typer.echo("") + typer.echo("next:") + typer.echo(f" giant train {data_arg} --config {config_path} --out {run_dir}") + + @app.command() def predict( data: Annotated[ diff --git a/giant/config.py b/giant/config.py index 38d2f59..d9e4265 100644 --- a/giant/config.py +++ b/giant/config.py @@ -372,6 +372,21 @@ def default_out_dir_name(cfg: dict, now: datetime | None = None) -> str: return name +def resolve_default_out_dir(cfg: dict, base: Path = Path("checkpoints")) -> Path: + """Auto-derived out dir from cfg's hyperparams (see `default_out_dir_name`), + with a numeric suffix loop so two runs whose name collides (same + non-default hyperparams, same to-the-minute timestamp) don't clobber each + other's directory. Shared by `giant train` and `giant new-run`. + """ + base_name = default_out_dir_name(cfg) + out_dir = base / base_name + suffix = 2 + while out_dir.exists(): + out_dir = base / f"{base_name}_{suffix}" + suffix += 1 + return out_dir + + def seed_everything(seed: int) -> None: random.seed(seed) np.random.seed(seed) diff --git a/tests/test_cli_new_run.py b/tests/test_cli_new_run.py new file mode 100644 index 0000000..54788c4 --- /dev/null +++ b/tests/test_cli_new_run.py @@ -0,0 +1,118 @@ +"""Tests for `giant new-run` (config.toml + run-dir scaffolding).""" + +from __future__ import annotations + +import tomllib +from pathlib import Path + +from typer.testing import CliRunner + +from giant.cli import app + +runner = CliRunner() + + +def test_writes_config_with_overrides_applied(tmp_path: Path): + out_dir = tmp_path / "run1" + result = runner.invoke( + app, + [ + "new-run", + "--out", + str(out_dir), + "--mode", + "ddpm", + "--hidden-dim", + "128", + "--n-blocks", + "4", + "--lr", + "0.0005", + ], + ) + assert result.exit_code == 0, result.output + + config_path = out_dir / "config.toml" + assert config_path.exists() + with open(config_path, "rb") as f: + cfg = tomllib.load(f) + + assert cfg["train"]["mode"] == "ddpm" + assert cfg["train"]["lr"] == 0.0005 + assert cfg["model"]["hidden_dim"] == 128 + assert cfg["model"]["n_blocks"] == 4 + # untouched defaults still present + assert cfg["train"]["epochs"] == 100 + assert "router" in cfg["model"] + + assert str(out_dir) in result.output + assert "" in result.output + assert "giant train" in result.output + + +def test_comment_and_provenance_recorded_in_meta(tmp_path: Path): + out_dir = tmp_path / "run2" + result = runner.invoke( + app, + ["new-run", "--out", str(out_dir), "--comment", "quick test"], + ) + assert result.exit_code == 0, result.output + + with open(out_dir / "config.toml", "rb") as f: + cfg = tomllib.load(f) + + assert cfg["meta"]["comment"] == "quick test" + assert cfg["meta"]["created_by"] == "giant new-run" + assert "created_at" in cfg["meta"] + assert "git_hash" in cfg["meta"] + + +def test_data_flag_fills_printed_next_step_commands(tmp_path: Path): + out_dir = tmp_path / "run3" + result = runner.invoke( + app, + ["new-run", "--out", str(out_dir), "--data", "/ceph/lbogner/train.parquet"], + ) + assert result.exit_code == 0, result.output + assert "/ceph/lbogner/train.parquet" in result.output + assert "" not in result.output + + +def test_dry_run_writes_nothing(tmp_path: Path): + out_dir = tmp_path / "run4" + result = runner.invoke( + app, + ["new-run", "--out", str(out_dir), "--hidden-dim", "512", "--dry-run"], + ) + assert result.exit_code == 0, result.output + assert "dry-run" in result.output + assert "hidden_dim = 512" in result.output + assert not out_dir.exists() + + +def test_force_guard_refuses_to_clobber_existing_checkpoints(tmp_path: Path): + out_dir = tmp_path / "run5" + out_dir.mkdir() + (out_dir / "last.pt").touch() + + result = runner.invoke(app, ["new-run", "--out", str(out_dir), "--mode", "ddpm"]) + assert result.exit_code != 0 + assert "already has last.pt" in result.output + assert not (out_dir / "config.toml").exists() + + result = runner.invoke( + app, ["new-run", "--out", str(out_dir), "--mode", "ddpm", "--force"] + ) + assert result.exit_code == 0, result.output + assert (out_dir / "config.toml").exists() + + +def test_default_out_dir_used_when_out_omitted(tmp_path: Path, monkeypatch): + monkeypatch.chdir(tmp_path) + result = runner.invoke(app, ["new-run", "--hidden-dim", "64"]) + assert result.exit_code == 0, result.output + + checkpoints_dir = tmp_path / "checkpoints" + run_dirs = list(checkpoints_dir.iterdir()) if checkpoints_dir.exists() else [] + assert len(run_dirs) == 1 + assert (run_dirs[0] / "config.toml").exists()