cc9646f279
CI / Lint (ruff check) (push) Successful in 37s
CI / Format (ruff format) (push) Successful in 38s
CI / Sync project version with tag (push) Has been skipped
CI / Type check (ty) (push) Successful in 27s
CI / Tests (push) Successful in 2m34s
CI / Lint (ruff check) (pull_request) Successful in 30s
CI / Format (ruff format) (pull_request) Successful in 30s
CI / Sync project version with tag (pull_request) Has been skipped
CI / Type check (ty) (pull_request) Successful in 29s
CI / Tests (pull_request) Successful in 2m35s
giant model summary --config config.toml builds the resolved Stage1/Stage2 graph from a config with no dataset attached (pdg_vocab/mat_vocab are supplied as placeholders via --pdg-vocab/--mat-vocab, since the real training vocab is dataset-derived) and prints per-module parameter counts, trunk in/out widths, which heads exist, and which conditioning/stage1_model/stage2_model config keys actually shaped the build. The consumed-keys half uses differential probing rather than static identifier matching: build once for a fingerprint (submodule presence, every parameter's/buffer's shape+dtype, every plain scalar attribute a module stores on itself), then perturb one leaf at a time, rebuild, and compare. A changed fingerprint (or a raise) means the key is consumed; no change means it's inert *under this particular config* -- e.g. any stage1_model.router.* key when router.enabled=false. A curated _NOT_BUILD_TIME table separates keys legitimately owned by the trainer/sampler/rollout (loss weights, WGAN-GP hyperparameters, teacher-forcing schedules) from genuinely-inert ones, verified against those call sites. A few config keys branch on equality against one specific string literal (n_sec.owner=="stage1", n_sec.mode=="stop_token", particle_type.target=="physical"); a single generic sentinel probe missed all three since the config's current value and the sentinel landed in the same branch, so those three leaves get their real alternative value tried too (_STRING_ALTERNATIVES). giant.config.leaf_paths is promoted out of tests/test_config_consumed_keys.py (previously a private test-local duplicate) so both audits -- the static per-identifier one and this new runtime per-config one -- walk the exact same DEFAULT_CONFIG tree. ExpertTrunk/RoutedTrunk now also expose in_dim (out_dim already existed), needed to report trunk widths generically. Decisions made during planning: --pdg-vocab/--mat-vocab default to 300 and len(MATERIAL_PROPERTIES); the consumed-keys report is scoped to conditioning/stage1_model/stage2_model only (train/meta are out of scope for a model-only build); the module tree prints every submodule at any depth. Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
1659 lines
64 KiB
Python
1659 lines
64 KiB
Python
from collections import Counter
|
|
from datetime import datetime, timezone
|
|
from enum import Enum
|
|
import math
|
|
from pathlib import Path
|
|
import re
|
|
from typing import Optional
|
|
import uuid as uuid_mod
|
|
|
|
import numpy as np
|
|
import yaml
|
|
import torch
|
|
import typer
|
|
from typing_extensions import Annotated
|
|
|
|
import pyarrow as pa
|
|
import pyarrow.parquet as pq
|
|
from tqdm import tqdm
|
|
|
|
from giant import config as gconfig
|
|
from giant.constants import (
|
|
LOCAL_TARGET_NAMES,
|
|
PREDICT_COORD_METADATA_KEY,
|
|
PREDICT_SCHEMA_VERSION,
|
|
PREDICT_SCHEMA_VERSION_KEY,
|
|
ROLLOUT_COORD_VALUE,
|
|
)
|
|
from giant.data.loader import (
|
|
event_id_offset,
|
|
find_parquet_files,
|
|
iter_file_chunks,
|
|
iter_cond_chunks,
|
|
)
|
|
from giant.data.transforms import (
|
|
build_features,
|
|
build_cond_features,
|
|
energy_simplex_decode,
|
|
inv_local_frame_rotation,
|
|
inv_log_transform,
|
|
reconstruct_post_pos,
|
|
)
|
|
from giant.checkpoint_io import CheckpointCompatibilityError, load_for_inference
|
|
from giant.geometry import GeometryOracle
|
|
from giant.materials import MATERIAL_PROPERTIES
|
|
from giant.pipeline import run_train_job
|
|
from giant.rollout import (
|
|
L1DistCollector,
|
|
decode_secondary_identity,
|
|
rollout as run_rollout,
|
|
)
|
|
from giant.sample import resolve_n_sec, sample_stage1, sample_stage2
|
|
|
|
app = typer.Typer(no_args_is_help=True)
|
|
|
|
|
|
def _router_total_experts(router_cfg: dict) -> int:
|
|
"""Total expert count for a router config, single-axis or composed.
|
|
|
|
A composed router runs one expert per *joint* cell, so its count is the
|
|
product of the per-axis `axis{i}_n_experts` (mirrors
|
|
`ComposedRouter.__init__` in giant.model.network); a single-axis router
|
|
just reports its own `n_experts`.
|
|
"""
|
|
if router_cfg.get("type") == "composed":
|
|
axis_counts = {
|
|
m.group(1): int(v) for k, v in router_cfg.items() if (m := re.match(r"^axis(\d+)_n_experts$", k))
|
|
}
|
|
return math.prod(axis_counts.values()) if axis_counts else 1
|
|
return int(router_cfg.get("n_experts", 1))
|
|
|
|
|
|
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.
|
|
|
|
`model_cfg` is either the new nested shape (has a `f"{stage}_model"` key
|
|
— the merged training `cfg`, or a checkpoint's new-format `model_config`)
|
|
or a v0.2 checkpoint's flat `model_config`. v0.3.0 dropped per-expert
|
|
sizing (giant.model.network's routed trunks always inherit the stage's
|
|
own hidden_dim/n_res_blocks — no more `resolve_expert_dims`), so the new
|
|
shape needs no special-casing there; the legacy flat shape may still
|
|
carry a v0.2 `expert_hidden_dim`/`expert_n_blocks` override, honoured
|
|
only when that checkpoint's router was actually enabled.
|
|
|
|
Routed models spend their FLOPs in the (smaller) expert trunks. Training
|
|
runs the full soft mixture (every expert on the whole batch), so its
|
|
activation memory scales with the expert count; inference does top-1
|
|
dispatch (each row hits one expert), so the batch just partitions across
|
|
experts and one expert's dims already bound it. estimate_batch_size
|
|
scales memory linearly with hidden_dim * n_blocks, so the training
|
|
multiplier folds into n_blocks.
|
|
"""
|
|
if f"{stage}_model" in model_cfg:
|
|
stage_cfg = model_cfg[f"{stage}_model"]
|
|
hidden_dim, n_blocks = stage_cfg["hidden_dim"], stage_cfg["n_res_blocks"]
|
|
router_cfg = stage_cfg.get("router")
|
|
else:
|
|
hidden_dim, n_blocks = model_cfg["hidden_dim"], model_cfg["n_blocks"]
|
|
router_cfg = model_cfg.get("router")
|
|
if router_cfg and router_cfg.get("enabled"):
|
|
hidden_dim = model_cfg.get("expert_hidden_dim") or hidden_dim
|
|
n_blocks = model_cfg.get("expert_n_blocks") or n_blocks
|
|
if router_cfg and router_cfg.get("enabled") and training:
|
|
n_blocks = n_blocks * _router_total_experts(router_cfg)
|
|
return hidden_dim, n_blocks
|
|
|
|
|
|
def _coerce_scalar(value: str) -> object:
|
|
"""Best-effort str -> bool/int/float, else leave as str.
|
|
|
|
CLI flag values always arrive as strings; router kwargs like
|
|
`n_experts` (int) or `temperature` (float) need to come out typed the
|
|
same way a TOML file's native types would, since they're merged into
|
|
the same `model.router` dict as file-sourced config.
|
|
"""
|
|
if value.lower() in ("true", "false"):
|
|
return value.lower() == "true"
|
|
try:
|
|
return int(value)
|
|
except ValueError:
|
|
pass
|
|
try:
|
|
return float(value)
|
|
except ValueError:
|
|
pass
|
|
return value
|
|
|
|
|
|
def _parse_router_axis_flags(specs: list[str]) -> dict[str, object]:
|
|
"""Parse repeated `--router-axis "type:key=val,key=val"` flags into
|
|
`axis{i}_{field}` flat keys (see `_parse_composed_axes` in
|
|
giant.model.network), indexed by flag order — the Nth `--router-axis`
|
|
becomes axis N.
|
|
"""
|
|
out: dict[str, object] = {}
|
|
for i, spec in enumerate(specs):
|
|
axis_type, _, rest = spec.partition(":")
|
|
out[f"axis{i}_type"] = axis_type
|
|
for pair in filter(None, rest.split(",")):
|
|
key, _, val = pair.partition("=")
|
|
out[f"axis{i}_{key}"] = _coerce_scalar(val)
|
|
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")
|
|
|
|
|
|
def _resolve_prediction_output(data: Path, out: Path | None) -> tuple[Path, Path, str]:
|
|
"""Return (out_path, resolved_dataset_path, pred_uuid).
|
|
|
|
When *out* is None the output path is derived from *data*:
|
|
- under /ceph/ → fixed central store with a UUID filename
|
|
- elsewhere → sibling of *data* with a UUID filename
|
|
"""
|
|
dataset_path = data.resolve()
|
|
pred_uuid = str(uuid_mod.uuid4())
|
|
if out is None:
|
|
if str(dataset_path).startswith("/ceph/"):
|
|
out = _CEPH_PREDICTIONS / f"{pred_uuid}.parquet"
|
|
else:
|
|
out = data.parent / f"{pred_uuid}.parquet"
|
|
return out, dataset_path, pred_uuid
|
|
|
|
|
|
def _write_prediction_ref(
|
|
checkpoint: Path,
|
|
pred_uuid: str,
|
|
out: Path,
|
|
dataset_path: Path,
|
|
comment: str | None = None,
|
|
) -> Path:
|
|
"""Write a YAML sidecar in the checkpoint directory and return its path."""
|
|
ref = {
|
|
"prediction_id": pred_uuid,
|
|
"output": str(out),
|
|
"dataset": str(dataset_path),
|
|
"checkpoint": str(checkpoint.resolve()),
|
|
"timestamp": datetime.now(timezone.utc).isoformat(),
|
|
}
|
|
if comment is not None:
|
|
ref["comment"] = comment
|
|
ref_path = checkpoint.parent / f"{pred_uuid}.yaml"
|
|
ref_path.write_text(yaml.dump(ref, default_flow_style=False, sort_keys=False))
|
|
return ref_path
|
|
|
|
|
|
@app.callback()
|
|
def _main() -> None:
|
|
"""GIANT — Geant4 step-function surrogate."""
|
|
|
|
|
|
class Mode(str, Enum):
|
|
flow = "flow"
|
|
ddpm = "ddpm"
|
|
wgan = "wgan"
|
|
|
|
|
|
class Decoder(str, Enum):
|
|
one_shot = "one_shot"
|
|
autoregressive = "autoregressive"
|
|
|
|
|
|
class Stage1Context(str, Enum):
|
|
truth = "truth"
|
|
sampled = "sampled"
|
|
|
|
|
|
# Conditioning itself lives in giant.config (imported below as gconfig) —
|
|
# shared with giant/tools/dwarf.py's Typer commands so the two CLIs can't
|
|
# silently drift apart on the option's valid values.
|
|
Conditioning = gconfig.Conditioning
|
|
|
|
|
|
class Coord(str, Enum):
|
|
global_ = "global"
|
|
local = "local"
|
|
|
|
|
|
class Weights(str, Enum):
|
|
raw = "raw"
|
|
ema = "ema"
|
|
|
|
|
|
@app.command()
|
|
def train(
|
|
data: Annotated[Path, typer.Argument(help="Parquet file or directory of parquet files")],
|
|
config: Annotated[
|
|
Optional[Path],
|
|
typer.Option("--config", "-c", help="TOML config file (overridden by explicit flags)"),
|
|
] = None,
|
|
mode: Annotated[
|
|
Optional[Mode],
|
|
typer.Option("--mode", "-m", help="Generative model: flow matching or DDPM"),
|
|
] = None,
|
|
epochs: Annotated[Optional[int], typer.Option("--epochs", "-e")] = None,
|
|
batch_size: Annotated[
|
|
Optional[str],
|
|
typer.Option(
|
|
"--batch-size",
|
|
"-b",
|
|
help="Integer, or 'auto' to estimate from free GPU memory (cuda devices only)",
|
|
),
|
|
] = None,
|
|
lr: Annotated[Optional[float], typer.Option("--lr", "-l")] = None,
|
|
weight_decay: Annotated[
|
|
Optional[float],
|
|
typer.Option("--weight-decay", "-W", help="AdamW weight decay (default: 0.01)"),
|
|
] = None,
|
|
ema_decay: Annotated[
|
|
Optional[float],
|
|
typer.Option(
|
|
"--ema-decay",
|
|
help="EMA decay for a shadow copy of the model weights, saved "
|
|
"alongside the raw weights in checkpoints (0 disables; default: 0.9999)",
|
|
),
|
|
] = None,
|
|
warmup_epochs: Annotated[Optional[int], typer.Option("--warmup-epochs", "-w")] = 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", help="Dropout probability in ResBlocks (default: 0.1)"),
|
|
] = None,
|
|
stage1_generator: Annotated[
|
|
Optional[Mode],
|
|
typer.Option(
|
|
"--stage1-generator",
|
|
help="Stage 1's generative objective — overrides --mode for stage 1 only",
|
|
),
|
|
] = None,
|
|
stage1_hidden_dim: Annotated[
|
|
Optional[int],
|
|
typer.Option(
|
|
"--stage1-hidden-dim",
|
|
help="Overrides --hidden-dim for stage 1 only (same effect today; "
|
|
"--hidden-dim is kept as a shorthand since stage 1 was the only "
|
|
"target before stage2_model got its own flags)",
|
|
),
|
|
] = None,
|
|
stage1_n_res_blocks: Annotated[
|
|
Optional[int],
|
|
typer.Option("--stage1-n-res-blocks", help="Overrides --n-blocks for stage 1 only"),
|
|
] = None,
|
|
stage1_dropout: Annotated[
|
|
Optional[float],
|
|
typer.Option("--stage1-dropout", help="Overrides --dropout for stage 1 only"),
|
|
] = None,
|
|
stage2_generator: Annotated[
|
|
Optional[Mode],
|
|
typer.Option(
|
|
"--stage2-generator",
|
|
help="Stage 2's generative objective — overrides --mode for stage 2 "
|
|
"only, e.g. combine with --stage1-generator flow for a mixed "
|
|
"flow/wgan run",
|
|
),
|
|
] = None,
|
|
stage2_hidden_dim: Annotated[
|
|
Optional[int],
|
|
typer.Option("--stage2-hidden-dim", help="Stage 2 trunk width"),
|
|
] = None,
|
|
stage2_n_res_blocks: Annotated[
|
|
Optional[int],
|
|
typer.Option("--stage2-n-res-blocks", help="Stage 2 trunk depth"),
|
|
] = None,
|
|
stage2_dropout: Annotated[
|
|
Optional[float],
|
|
typer.Option("--stage2-dropout", help="Dropout inside stage 2's ResBlocks"),
|
|
] = None,
|
|
stage2_decoder: Annotated[
|
|
Optional[Decoder],
|
|
typer.Option(
|
|
"--stage2-decoder",
|
|
help="one_shot: predict all k_max secondary slots at once (v0.2 "
|
|
"behaviour). autoregressive: emit one secondary at a time in "
|
|
"descending-energy order (default)",
|
|
),
|
|
] = None,
|
|
stage2_k_max: Annotated[
|
|
Optional[int],
|
|
typer.Option(
|
|
"--stage2-k-max",
|
|
help="Maximum secondary slots (fixed width under one_shot, a "
|
|
"generation-loop safety cap under autoregressive; default: 15)",
|
|
),
|
|
] = None,
|
|
stage2_context_dim: Annotated[
|
|
Optional[int],
|
|
typer.Option(
|
|
"--stage2-context-dim",
|
|
help="Width of the projected stage-1 outcome fed into stage 2's conditioning (default: 64)",
|
|
),
|
|
] = None,
|
|
stage2_stage1_context: Annotated[
|
|
Optional[Stage1Context],
|
|
typer.Option(
|
|
"--stage2-stage1-context",
|
|
help="What stage 2 conditions on during training: 'truth' (the "
|
|
"ground-truth stage-1 target, detached — default) or 'sampled' "
|
|
"(stage 1's own sampled output, closing the train/inference gap "
|
|
"at the cost of an extra sampling pass per batch)",
|
|
),
|
|
] = None,
|
|
conditioning: Annotated[
|
|
Optional[Conditioning],
|
|
typer.Option(
|
|
"--conditioning",
|
|
help="Input conditioning: continuous physical properties "
|
|
"(mass/charge/Z_eff/A_eff/density/X0/lambda_int, default) or the "
|
|
"original learned PDG/material embeddings",
|
|
),
|
|
] = None,
|
|
router: Annotated[
|
|
Optional[bool],
|
|
typer.Option(
|
|
"--router/--no-router",
|
|
help="Route both stages through a mixture of small experts "
|
|
"instead of one monolithic trunk (see model.router in config.toml)",
|
|
),
|
|
] = None,
|
|
router_type: Annotated[
|
|
Optional[str],
|
|
typer.Option("--router-type", help="Router implementation name (see ROUTER_REGISTRY)"),
|
|
] = None,
|
|
n_experts: Annotated[Optional[int], typer.Option("--n-experts", help="Number of routed experts")] = None,
|
|
router_axis: Annotated[
|
|
Optional[list[str]],
|
|
typer.Option(
|
|
"--router-axis",
|
|
help="Composed-router axis spec 'type:key=val,key=val' (repeatable; "
|
|
"Nth flag = axis N). Use with --router-type composed instead of "
|
|
"--n-experts, e.g. --router-axis 'energy:n_experts=4' "
|
|
"--router-axis 'pdg:n_experts=3,emb_dim=8'",
|
|
),
|
|
] = None,
|
|
n_critic: Annotated[
|
|
Optional[int],
|
|
typer.Option(
|
|
"--n-critic",
|
|
help="WGAN-GP (--mode wgan only): critic updates per generator update (default: 5)",
|
|
),
|
|
] = None,
|
|
gp_weight: Annotated[
|
|
Optional[float],
|
|
typer.Option(
|
|
"--gp-weight",
|
|
help="WGAN-GP (--mode wgan only): gradient-penalty coefficient (default: 10.0)",
|
|
),
|
|
] = None,
|
|
noise_dim: Annotated[
|
|
Optional[int],
|
|
typer.Option(
|
|
"--noise-dim",
|
|
help="WGAN (--mode wgan only): generator input noise-vector width (default: 64)",
|
|
),
|
|
] = None,
|
|
critic_lr: Annotated[
|
|
Optional[float],
|
|
typer.Option(
|
|
"--critic-lr",
|
|
help="WGAN-GP (--mode wgan only): critic learning rate (default: same as --lr)",
|
|
),
|
|
] = None,
|
|
stage1_n_critic: Annotated[
|
|
Optional[int],
|
|
typer.Option("--stage1-n-critic", help="Overrides --n-critic for stage 1 only"),
|
|
] = None,
|
|
stage1_gp_weight: Annotated[
|
|
Optional[float],
|
|
typer.Option("--stage1-gp-weight", help="Overrides --gp-weight for stage 1 only"),
|
|
] = None,
|
|
stage1_noise_dim: Annotated[
|
|
Optional[int],
|
|
typer.Option("--stage1-noise-dim", help="Overrides --noise-dim for stage 1 only"),
|
|
] = None,
|
|
stage1_critic_lr: Annotated[
|
|
Optional[float],
|
|
typer.Option("--stage1-critic-lr", help="Overrides --critic-lr for stage 1 only"),
|
|
] = None,
|
|
stage2_n_critic: Annotated[
|
|
Optional[int],
|
|
typer.Option("--stage2-n-critic", help="Overrides --n-critic for stage 2 only"),
|
|
] = None,
|
|
stage2_gp_weight: Annotated[
|
|
Optional[float],
|
|
typer.Option("--stage2-gp-weight", help="Overrides --gp-weight for stage 2 only"),
|
|
] = None,
|
|
stage2_noise_dim: Annotated[
|
|
Optional[int],
|
|
typer.Option(
|
|
"--stage2-noise-dim",
|
|
help="Overrides --noise-dim for stage 2 only; under "
|
|
"--stage2-decoder autoregressive a fresh draw is made per token",
|
|
),
|
|
] = None,
|
|
stage2_critic_lr: Annotated[
|
|
Optional[float],
|
|
typer.Option("--stage2-critic-lr", help="Overrides --critic-lr for stage 2 only"),
|
|
] = None,
|
|
stage1_critic_hidden_dim: Annotated[
|
|
Optional[int],
|
|
typer.Option(
|
|
"--stage1-critic-hidden-dim",
|
|
help="WGAN-GP (--mode wgan only): critic width for stage 1 (default: same as generator's hidden_dim)",
|
|
),
|
|
] = None,
|
|
stage1_critic_n_res_blocks: Annotated[
|
|
Optional[int],
|
|
typer.Option(
|
|
"--stage1-critic-n-res-blocks",
|
|
help="WGAN-GP (--mode wgan only): critic depth for stage 1 (default: same as generator's n_res_blocks)",
|
|
),
|
|
] = None,
|
|
stage2_critic_hidden_dim: Annotated[
|
|
Optional[int],
|
|
typer.Option(
|
|
"--stage2-critic-hidden-dim",
|
|
help="WGAN-GP (--mode wgan only): critic width for stage 2 (default: same as generator's hidden_dim)",
|
|
),
|
|
] = None,
|
|
stage2_critic_n_res_blocks: Annotated[
|
|
Optional[int],
|
|
typer.Option(
|
|
"--stage2-critic-n-res-blocks",
|
|
help="WGAN-GP (--mode wgan only): critic depth for stage 2 (default: same as generator's n_res_blocks)",
|
|
),
|
|
] = None,
|
|
val_fraction: Annotated[Optional[float], typer.Option("--val-fraction", "-f")] = None,
|
|
seed: Annotated[
|
|
Optional[int],
|
|
typer.Option("--seed", "-s", help="Random seed for reproducibility"),
|
|
] = None,
|
|
validate_every: Annotated[
|
|
Optional[int],
|
|
typer.Option(
|
|
"--validate-every",
|
|
"-v",
|
|
help="Run marginal+KL validation every N epochs (0 disables)",
|
|
),
|
|
] = None,
|
|
validate_steps: Annotated[
|
|
Optional[int],
|
|
typer.Option(
|
|
"--validate-steps",
|
|
"-t",
|
|
help="Flow matching ODE steps used during marginal validation "
|
|
"(ignored in ddpm mode, which always runs the full schedule)",
|
|
),
|
|
] = None,
|
|
max_val_batches: Annotated[
|
|
Optional[int],
|
|
typer.Option(
|
|
"--max-val-batches",
|
|
help="Cap the per-epoch val-loss pass to N batches (0 = full val set every epoch; default: 200)",
|
|
),
|
|
] = None,
|
|
shuffle_buffer: Annotated[
|
|
int,
|
|
typer.Option("--shuffle-buffer", "-B", help="Rows held in RAM per worker for shuffling"),
|
|
] = 65536,
|
|
cache_setup: Annotated[
|
|
bool,
|
|
typer.Option(
|
|
"--cache-setup/--no-cache-setup",
|
|
help="Cache the training setup stage's expensive per-file "
|
|
"precomputation (vocab maps, event split index, normalizer stats) "
|
|
"in a JSON sidecar next to the data, so a repeat `giant train` "
|
|
"against the same dataset (e.g. a hyperparameter sweep) can skip "
|
|
"re-deriving it",
|
|
),
|
|
] = True,
|
|
rebuild_setup_cache: Annotated[
|
|
bool,
|
|
typer.Option(
|
|
"--rebuild-setup-cache/--no-rebuild-setup-cache",
|
|
help="Ignore any existing setup cache sidecar and recompute every "
|
|
"section fresh for this run (still writes the refreshed sections "
|
|
"back to the sidecar for later runs; no effect if --no-cache-setup)",
|
|
),
|
|
] = False,
|
|
out: Annotated[
|
|
Optional[Path],
|
|
typer.Option(
|
|
"--out",
|
|
"-o",
|
|
help="Checkpoint dir (default: timestamped dir from hyperparams, "
|
|
"or the --resume checkpoint's own dir when resuming)",
|
|
),
|
|
] = None,
|
|
device: Annotated[
|
|
Optional[str],
|
|
typer.Option("--device", "-D", help="cpu | cuda | mps (default: auto)"),
|
|
] = None,
|
|
num_workers: Annotated[Optional[int], typer.Option("--num-workers", "-j")] = None,
|
|
resume: Annotated[
|
|
Optional[Path],
|
|
typer.Option("--resume", "-r", help="Checkpoint .pt to resume training from"),
|
|
] = None,
|
|
wandb: Annotated[
|
|
Optional[bool],
|
|
typer.Option(
|
|
"--wandb/--no-wandb",
|
|
help="Log per-epoch training metrics to Weights & Biases (requires `uv sync --extra wandb`)",
|
|
),
|
|
] = None,
|
|
wandb_project: Annotated[
|
|
Optional[str],
|
|
typer.Option("--wandb-project", help="W&B project name (default: giant)"),
|
|
] = None,
|
|
wandb_run_name: Annotated[
|
|
Optional[str],
|
|
typer.Option("--wandb-run-name", help="W&B run name (default: out_dir name)"),
|
|
] = None,
|
|
wandb_log_every: Annotated[
|
|
Optional[int],
|
|
typer.Option(
|
|
"--wandb-log-every",
|
|
help="Log batch-level loss/grad_norm/lr to W&B every N optimizer "
|
|
"steps (default: 50); per-epoch metrics always log in full",
|
|
),
|
|
] = None,
|
|
) -> None:
|
|
"""Train the GIANT surrogate model."""
|
|
batch_size_auto = False
|
|
batch_size_value: Optional[int] = None
|
|
if batch_size is not None:
|
|
if batch_size.strip().lower() == "auto":
|
|
batch_size_auto = True
|
|
else:
|
|
try:
|
|
batch_size_value = int(batch_size)
|
|
except ValueError:
|
|
typer.echo(
|
|
f"error: --batch-size must be an integer or 'auto', got {batch_size!r}",
|
|
err=True,
|
|
)
|
|
raise typer.Exit(1)
|
|
|
|
cli_router = _router_cli_overrides(router, router_type, n_experts, router_axis)
|
|
flag_values: dict[str, object] = {
|
|
"epochs": epochs,
|
|
"batch_size": batch_size_value,
|
|
"lr": lr,
|
|
"weight_decay": weight_decay,
|
|
"ema_decay": ema_decay,
|
|
"warmup_epochs": warmup_epochs,
|
|
"val_fraction": val_fraction,
|
|
"num_workers": num_workers,
|
|
"seed": seed,
|
|
"validate_every": validate_every,
|
|
"validate_steps": validate_steps,
|
|
"max_val_batches": max_val_batches,
|
|
"wandb": wandb,
|
|
"wandb_project": wandb_project,
|
|
"wandb_run_name": wandb_run_name,
|
|
"wandb_log_every": wandb_log_every,
|
|
"hidden_dim": hidden_dim,
|
|
"n_blocks": n_blocks,
|
|
"dropout": dropout,
|
|
"stage1_hidden_dim": stage1_hidden_dim,
|
|
"stage1_n_res_blocks": stage1_n_res_blocks,
|
|
"stage1_dropout": stage1_dropout,
|
|
"stage2_hidden_dim": stage2_hidden_dim,
|
|
"stage2_n_res_blocks": stage2_n_res_blocks,
|
|
"stage2_dropout": stage2_dropout,
|
|
"stage2_decoder": stage2_decoder.value if stage2_decoder is not None else None,
|
|
"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,
|
|
"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,
|
|
"conditioning": conditioning.value if conditioning is not None else None,
|
|
"emb_dim": emb_dim,
|
|
"router_config": cli_router or None,
|
|
"n_critic": n_critic,
|
|
"gp_weight": gp_weight,
|
|
"noise_dim": noise_dim,
|
|
"critic_lr": critic_lr,
|
|
"stage1_n_critic": stage1_n_critic,
|
|
"stage1_gp_weight": stage1_gp_weight,
|
|
"stage1_noise_dim": stage1_noise_dim,
|
|
"stage1_critic_lr": stage1_critic_lr,
|
|
"stage2_n_critic": stage2_n_critic,
|
|
"stage2_gp_weight": stage2_gp_weight,
|
|
"stage2_noise_dim": stage2_noise_dim,
|
|
"stage2_critic_lr": stage2_critic_lr,
|
|
"stage1_critic_hidden_dim": stage1_critic_hidden_dim,
|
|
"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,
|
|
}
|
|
overrides = gconfig.overrides_from_flags(flag_values)
|
|
|
|
cfg = gconfig.merge_cli_overrides(gconfig.DEFAULT_CONFIG, config, overrides)
|
|
gconfig.validate_config(cfg)
|
|
t = cfg["train"]
|
|
|
|
_device = torch.device(device) if device else gconfig.auto_device()
|
|
|
|
if batch_size_auto:
|
|
est_hidden_dim, est_n_blocks = _batch_size_estimate_dims(cfg, training=True, stage="stage1")
|
|
try:
|
|
t["batch_size"] = gconfig.estimate_batch_size(est_hidden_dim, est_n_blocks, _device)
|
|
except ValueError as exc:
|
|
typer.echo(f"error: {exc}", err=True)
|
|
raise typer.Exit(1)
|
|
typer.echo(f"batch_size: {t['batch_size']} (auto-estimated from free GPU memory)")
|
|
|
|
if out is not None:
|
|
out_dir = out
|
|
elif resume is not None:
|
|
# Continue writing into the resumed checkpoint's own directory
|
|
# rather than recomputing a hyperparam-derived name — the latter
|
|
# would (a) collide with the original run's dir only by accident
|
|
# (same day, unchanged hyperparams) and now never collides at all
|
|
# since the fresh-run name below is timestamped to the second, and
|
|
# (b) silently start a fresh directory if a resumed run tweaks any
|
|
# hyperparam baked into the name (e.g. --lr for a fine-tune).
|
|
out_dir = resume.parent
|
|
else:
|
|
# 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.training) — 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}")
|
|
|
|
run_train_job(
|
|
data=data,
|
|
cfg=cfg,
|
|
out_dir=out_dir,
|
|
device=_device,
|
|
shuffle_buffer=shuffle_buffer,
|
|
num_workers=t["num_workers"],
|
|
resume=resume,
|
|
cache_setup=cache_setup,
|
|
rebuild_setup_cache=rebuild_setup_cache,
|
|
echo=typer.echo,
|
|
)
|
|
|
|
|
|
@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,
|
|
stage1_generator: Annotated[Optional[Mode], typer.Option("--stage1-generator")] = None,
|
|
stage1_hidden_dim: Annotated[Optional[int], typer.Option("--stage1-hidden-dim")] = None,
|
|
stage1_n_res_blocks: Annotated[Optional[int], typer.Option("--stage1-n-res-blocks")] = None,
|
|
stage1_dropout: Annotated[Optional[float], typer.Option("--stage1-dropout")] = None,
|
|
stage2_generator: Annotated[Optional[Mode], typer.Option("--stage2-generator")] = None,
|
|
stage2_hidden_dim: Annotated[Optional[int], typer.Option("--stage2-hidden-dim")] = None,
|
|
stage2_n_res_blocks: Annotated[Optional[int], typer.Option("--stage2-n-res-blocks")] = None,
|
|
stage2_dropout: Annotated[Optional[float], typer.Option("--stage2-dropout")] = None,
|
|
stage2_decoder: Annotated[Optional[Decoder], typer.Option("--stage2-decoder")] = None,
|
|
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,
|
|
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_router = _router_cli_overrides(router, router_type, n_experts, router_axis)
|
|
flag_values: dict[str, object] = {
|
|
"epochs": epochs,
|
|
"batch_size": batch_size,
|
|
"lr": lr,
|
|
"hidden_dim": hidden_dim,
|
|
"n_blocks": n_blocks,
|
|
"dropout": dropout,
|
|
"stage1_hidden_dim": stage1_hidden_dim,
|
|
"stage1_n_res_blocks": stage1_n_res_blocks,
|
|
"stage1_dropout": stage1_dropout,
|
|
"stage2_hidden_dim": stage2_hidden_dim,
|
|
"stage2_n_res_blocks": stage2_n_res_blocks,
|
|
"stage2_dropout": stage2_dropout,
|
|
"stage2_decoder": stage2_decoder.value if stage2_decoder is not None else None,
|
|
"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,
|
|
"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,
|
|
"conditioning": conditioning.value if conditioning is not None else None,
|
|
"emb_dim": emb_dim,
|
|
"router_config": cli_router or None,
|
|
}
|
|
overrides = gconfig.overrides_from_flags(flag_values)
|
|
|
|
cfg = gconfig.merge_cli_overrides(gconfig.DEFAULT_CONFIG, config, overrides)
|
|
gconfig.validate_config(cfg)
|
|
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", "conditioning", "stage1_model", "stage2_model"):
|
|
typer.echo(f"[{section}]")
|
|
for k, v in cfg[section].items():
|
|
if k == "router":
|
|
continue
|
|
typer.echo(f" {k} = {v}")
|
|
return
|
|
|
|
meta = {
|
|
# Tags the written config.toml as v0.3-shaped so a later
|
|
# `migrate_config` load (e.g. `giant train --config ...`) treats it
|
|
# as already-migrated instead of misreading it as v0.2 and dropping
|
|
# its stage1_model/stage2_model/conditioning content.
|
|
"config_version": gconfig.CONFIG_VERSION,
|
|
"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 "<data.parquet>"
|
|
typer.echo("")
|
|
typer.echo("next:")
|
|
typer.echo(f" giant train {data_arg} --config {config_path} --out {run_dir}")
|
|
|
|
|
|
model_app = typer.Typer(
|
|
no_args_is_help=True,
|
|
help="Inspect a resolved model architecture without training.",
|
|
)
|
|
app.add_typer(model_app, name="model")
|
|
|
|
|
|
@model_app.command("summary")
|
|
def model_summary(
|
|
config: Annotated[
|
|
Optional[Path],
|
|
typer.Option("--config", "-c", help="TOML config file (default: built-in defaults)"),
|
|
] = None,
|
|
pdg_vocab: Annotated[
|
|
int,
|
|
typer.Option(
|
|
"--pdg-vocab",
|
|
help="Placeholder PDG vocab size for conditioning.particle.type='embedding' "
|
|
"or a pdg/process router (no dataset attached to derive the real training vocab)",
|
|
),
|
|
] = 300,
|
|
mat_vocab: Annotated[
|
|
int,
|
|
typer.Option(
|
|
"--mat-vocab",
|
|
help="Placeholder material vocab size for conditioning.material.type='embedding' "
|
|
"or a process router (default: the number of known materials in giant.materials)",
|
|
),
|
|
] = len(MATERIAL_PROPERTIES),
|
|
) -> None:
|
|
"""Build the resolved model graph from a config with no dataset attached, and
|
|
print per-module parameter counts, trunk widths, which heads exist, and which
|
|
conditioning/stage1_model/stage2_model config keys actually shaped it."""
|
|
from giant.model.summary import render_summary, summarize_model
|
|
|
|
cfg = gconfig.merge_cli_overrides(gconfig.DEFAULT_CONFIG, config, {})
|
|
try:
|
|
gconfig.validate_config(cfg)
|
|
except ValueError as exc:
|
|
typer.echo(f"error: {exc}", err=True)
|
|
raise typer.Exit(1)
|
|
|
|
summary = summarize_model(cfg, pdg_vocab=pdg_vocab, mat_vocab=mat_vocab)
|
|
typer.echo(render_summary(summary))
|
|
|
|
|
|
@app.command()
|
|
def predict(
|
|
data: Annotated[Path, typer.Argument(help="Parquet file or directory of parquet files")],
|
|
checkpoint: Annotated[
|
|
Path,
|
|
typer.Option(
|
|
"--checkpoint",
|
|
"-c",
|
|
help="Path to checkpoint .pt file (best.pt or last.pt)",
|
|
),
|
|
],
|
|
coord: Annotated[
|
|
Coord,
|
|
typer.Option(
|
|
"--coord",
|
|
"-C",
|
|
help="global: full physical units, world frame (default). "
|
|
"local: raw 9D model output (denormalised only, local frame, "
|
|
"log-scaled scalars) alongside the matching ground-truth target "
|
|
"for the same input file — requires post-step columns.",
|
|
),
|
|
] = Coord.global_,
|
|
out: Annotated[
|
|
Optional[Path],
|
|
typer.Option(
|
|
"--out",
|
|
"-o",
|
|
help="Output parquet path (default: <data>_predicted[_local].parquet)",
|
|
),
|
|
] = None,
|
|
batch_size: Annotated[
|
|
str,
|
|
typer.Option(
|
|
"--batch-size",
|
|
"-b",
|
|
help="Inference batch size, or 'auto' to estimate from free GPU memory (cuda devices only)",
|
|
),
|
|
] = "4096",
|
|
steps: Annotated[
|
|
int,
|
|
typer.Option(
|
|
"--steps",
|
|
"-s",
|
|
help="Flow matching ODE steps (ignored for a wgan checkpoint)",
|
|
),
|
|
] = 10,
|
|
weights: Annotated[
|
|
Weights,
|
|
typer.Option(
|
|
"--weights",
|
|
help="raw: the live training weights. ema: the EMA shadow copy "
|
|
"(see --ema-decay in `giant train`) — usually cleaner samples, "
|
|
"requires a checkpoint trained with EMA enabled.",
|
|
),
|
|
] = Weights.raw,
|
|
device: Annotated[
|
|
Optional[str],
|
|
typer.Option("--device", "-d", help="cpu | cuda | mps (default: auto)"),
|
|
] = None,
|
|
comment: Annotated[
|
|
Optional[str],
|
|
typer.Option(
|
|
"--comment",
|
|
"-m",
|
|
help="Free-text note recorded in the prediction's YAML sidecar",
|
|
),
|
|
] = None,
|
|
) -> None:
|
|
"""Run trained model on a parquet file and save predictions."""
|
|
batch_size_auto = False
|
|
batch_size_value: Optional[int] = None
|
|
if batch_size.strip().lower() == "auto":
|
|
batch_size_auto = True
|
|
else:
|
|
try:
|
|
batch_size_value = int(batch_size)
|
|
except ValueError:
|
|
typer.echo(
|
|
f"error: --batch-size must be an integer or 'auto', got {batch_size!r}",
|
|
err=True,
|
|
)
|
|
raise typer.Exit(1)
|
|
|
|
_device = torch.device(device) if device else gconfig.auto_device()
|
|
typer.echo(f"device: {_device}")
|
|
|
|
# --- Load checkpoint ---
|
|
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})")
|
|
|
|
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
|
|
sec_type_topn_map = ctx.sec_type_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(ctx.model_config, training=False)
|
|
try:
|
|
batch_size_value = gconfig.estimate_batch_size(
|
|
est_hidden_dim,
|
|
est_n_blocks,
|
|
_device,
|
|
training=False,
|
|
)
|
|
except ValueError as exc:
|
|
typer.echo(f"error: {exc}", err=True)
|
|
raise typer.Exit(1)
|
|
typer.echo(f"batch_size: {batch_size_value} (auto-estimated from free GPU memory)")
|
|
|
|
assert batch_size_value is not None
|
|
bs = batch_size_value
|
|
|
|
# --- Output path ---
|
|
out, dataset_path, pred_uuid = _resolve_prediction_output(data, out)
|
|
out.parent.mkdir(parents=True, exist_ok=True)
|
|
typer.echo(f"output: {out}")
|
|
|
|
# --- Stream input, generate predictions, write output ---
|
|
files = find_parquet_files(data)
|
|
typer.echo(f"found {len(files)} parquet file(s)")
|
|
|
|
writer: pq.ParquetWriter | None = None
|
|
total = 0
|
|
skipped = 0
|
|
unknown_pdg_counts: Counter[int] = Counter()
|
|
total_rows = sum(pq.ParquetFile(path).metadata.num_rows for path in files)
|
|
|
|
def chunk_iter(path: Path, offset: int):
|
|
if coord == Coord.local:
|
|
return iter_file_chunks(path, offset=offset, k_max=stage2_k_max)
|
|
return iter_cond_chunks(path, offset=offset)
|
|
|
|
# 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}
|
|
|
|
def _process(piece: dict[str, np.ndarray]) -> None:
|
|
nonlocal writer, total
|
|
|
|
if coord == Coord.local:
|
|
feats = build_features(
|
|
piece,
|
|
pdg_map,
|
|
mat_map,
|
|
particle_conditioning=particle_conditioning,
|
|
material_conditioning=material_conditioning,
|
|
pdg_topn_map=cond_pdg_topn,
|
|
mat_topn_map=cond_mat_topn,
|
|
k_max=stage2_k_max,
|
|
)
|
|
cond_cat = feats.cond_cat
|
|
target_raw = feats.target_s1
|
|
cond_cont = cond_norm.transform(feats.cond_cont)
|
|
else:
|
|
cond_cont, cond_cat = build_cond_features(
|
|
piece,
|
|
pdg_map,
|
|
mat_map,
|
|
cond_norm,
|
|
particle_conditioning=particle_conditioning,
|
|
material_conditioning=material_conditioning,
|
|
pdg_topn_map=cond_pdg_topn,
|
|
mat_topn_map=cond_mat_topn,
|
|
)
|
|
|
|
cc = torch.from_numpy(cond_cont).float().to(_device)
|
|
ck = torch.from_numpy(cond_cat).long().to(_device)
|
|
stage1_norm, n_sec_pred = sample_stage1(model, cc, ck, steps=steps, ddpm_steps=stage1_ddpm_steps)
|
|
|
|
if coord == Coord.global_:
|
|
# A fresh v0.3.0 Stage1Model owns no n_sec_head —
|
|
# sample_stage1 returns n_sec_pred=None then, so ask stage 2.
|
|
n_sec_pred = resolve_n_sec(model, sec_decoder, cc, ck, stage1_norm, n_sec_pred)
|
|
sec_cont, sec_type, sec_valid_pred = sample_stage2(sec_decoder, cc, ck, stage1_norm, n_sec_pred, steps)
|
|
# A stop-token decoder resolves n_sec_pred=None above — read the
|
|
# real count back off sec_valid_pred instead (a no-op round trip
|
|
# under every other n_sec.mode, where sec_valid_pred was built
|
|
# FROM n_sec_pred in the first place).
|
|
n_sec_pred_np = sec_valid_pred.sum(dim=-1).cpu().numpy()
|
|
|
|
pred = stage1_norm.cpu().numpy() # normalised
|
|
|
|
# Inverse-normalise → local frame, log-scaled scalars
|
|
raw = tgt_norm.inverse_transform(pred)
|
|
|
|
if coord == Coord.local:
|
|
table = pa.table(
|
|
{
|
|
"event_id": piece["event_id"],
|
|
"pdg": piece["pdg"],
|
|
"pre_x": piece["pre_pos"][:, 0],
|
|
"pre_y": piece["pre_pos"][:, 1],
|
|
"pre_z": piece["pre_pos"][:, 2],
|
|
"pre_E": piece["pre_E"],
|
|
"pre_dx": piece["pre_dir"][:, 0],
|
|
"pre_dy": piece["pre_dir"][:, 1],
|
|
"pre_dz": piece["pre_dir"][:, 2],
|
|
"material": piece["material"],
|
|
"layer_id": piece["layer_id"],
|
|
"n_sec": piece["n_sec"],
|
|
**{f"pred_{name}": raw[:, j] for j, name in enumerate(LOCAL_TARGET_NAMES)},
|
|
**{f"true_{name}": target_raw[:, j] for j, name in enumerate(LOCAL_TARGET_NAMES)},
|
|
}
|
|
)
|
|
else:
|
|
step_length = inv_log_transform(raw[:, 0])
|
|
# Columns 1:3 are ALR coords of the deposit/secondary/post energy
|
|
# simplex; decode them against pre_E so edep + e_sec + post_E == pre_E
|
|
# (hence delta_e == edep + e_sec) holds by construction. e_sec_pred
|
|
# doubles as the stick-breaking energy budget for the Stage-2 decode
|
|
# below, since the model has no other source for it at inference.
|
|
edep, e_sec_pred, _post_E, delta_e = energy_simplex_decode(raw[:, 1:3], piece["pre_E"])
|
|
|
|
# Normalise predicted direction then rotate back to world frame
|
|
post_dir_local = raw[:, 3:6].copy()
|
|
norms = np.linalg.norm(post_dir_local, axis=1, keepdims=True)
|
|
post_dir_local /= np.where(norms < 1e-8, 1.0, norms)
|
|
post_dir_world = inv_local_frame_rotation(piece["pre_dir"], post_dir_local)
|
|
|
|
# Same for the travel direction, then reconstruct post_pos from
|
|
# the single shared step_length so the two stay consistent.
|
|
travel_dir_local = raw[:, 6:9].copy()
|
|
norms = np.linalg.norm(travel_dir_local, axis=1, keepdims=True)
|
|
travel_dir_local /= np.where(norms < 1e-8, 1.0, norms)
|
|
post_pos_world = reconstruct_post_pos(piece["pre_pos"], piece["pre_dir"], step_length, travel_dir_local)
|
|
|
|
# particle_type.target="physical": sec_pdg_code is a reporting-
|
|
# only nearest-known-PDG label (never fed back into the model —
|
|
# "no snapping at inference"). "onehot"/"embedding": PDG
|
|
# resolution IS the secondary's identity — see
|
|
# decode_secondary_identity's docstring.
|
|
sec_E, sec_dir_world, sec_mass, sec_charge, sec_pdg_code, _l1_dist = decode_secondary_identity(
|
|
sec_decoder,
|
|
sec_cont,
|
|
sec_type,
|
|
n_sec_pred_np,
|
|
e_sec_pred,
|
|
piece["pre_dir"],
|
|
sec_phys_norm,
|
|
pdg_map,
|
|
sec_type_topn_map,
|
|
other_policy,
|
|
None,
|
|
)
|
|
sec_pdg_list = [sec_pdg_code[i, :n].tolist() for i, n in enumerate(n_sec_pred_np)]
|
|
sec_E_list = [sec_E[i, :n].tolist() for i, n in enumerate(n_sec_pred_np)]
|
|
sec_dx_list = [sec_dir_world[i, :n, 0].tolist() for i, n in enumerate(n_sec_pred_np)]
|
|
sec_dy_list = [sec_dir_world[i, :n, 1].tolist() for i, n in enumerate(n_sec_pred_np)]
|
|
sec_dz_list = [sec_dir_world[i, :n, 2].tolist() for i, n in enumerate(n_sec_pred_np)]
|
|
|
|
table = pa.table(
|
|
{
|
|
"event_id": piece["event_id"],
|
|
"pdg": piece["pdg"],
|
|
"pre_x": piece["pre_pos"][:, 0],
|
|
"pre_y": piece["pre_pos"][:, 1],
|
|
"pre_z": piece["pre_pos"][:, 2],
|
|
"pre_E": piece["pre_E"],
|
|
"pre_dx": piece["pre_dir"][:, 0],
|
|
"pre_dy": piece["pre_dir"][:, 1],
|
|
"pre_dz": piece["pre_dir"][:, 2],
|
|
"material": piece["material"],
|
|
"layer_id": piece["layer_id"],
|
|
"n_sec": piece["n_sec"],
|
|
"n_sec_pred": n_sec_pred_np,
|
|
"step_length": step_length,
|
|
"delta_e": delta_e,
|
|
"edep": edep,
|
|
"post_dx": post_dir_world[:, 0],
|
|
"post_dy": post_dir_world[:, 1],
|
|
"post_dz": post_dir_world[:, 2],
|
|
"post_x": post_pos_world[:, 0],
|
|
"post_y": post_pos_world[:, 1],
|
|
"post_z": post_pos_world[:, 2],
|
|
"sec_pdg_list": sec_pdg_list,
|
|
"sec_E_list": sec_E_list,
|
|
"sec_dx_list": sec_dx_list,
|
|
"sec_dy_list": sec_dy_list,
|
|
"sec_dz_list": sec_dz_list,
|
|
}
|
|
)
|
|
|
|
table = table.replace_schema_metadata(
|
|
{
|
|
PREDICT_COORD_METADATA_KEY: coord.value,
|
|
PREDICT_SCHEMA_VERSION_KEY: PREDICT_SCHEMA_VERSION,
|
|
}
|
|
)
|
|
|
|
if writer is None:
|
|
writer = pq.ParquetWriter(out, table.schema)
|
|
writer.write_table(table)
|
|
total += len(piece["event_id"])
|
|
|
|
# Buffer rows across row-group boundaries so the inference batch size
|
|
# isn't capped by however the source file happens to be chunked.
|
|
buffer: dict[str, np.ndarray] | None = None
|
|
|
|
bar = tqdm(total=total_rows, desc="predict", unit="row", dynamic_ncols=True)
|
|
for i, path in enumerate(files):
|
|
for chunk in chunk_iter(path, offset=event_id_offset(i)):
|
|
N_in = len(chunk["event_id"])
|
|
|
|
pdg_mask = np.array([int(p) in pdg_map for p in chunk["pdg"]])
|
|
if not pdg_mask.all():
|
|
unknown_pdg_counts.update(int(p) for p in chunk["pdg"][~pdg_mask])
|
|
chunk = {k: v[pdg_mask] for k, v in chunk.items()}
|
|
|
|
skipped += N_in - len(chunk["event_id"])
|
|
bar.update(N_in)
|
|
if len(chunk["event_id"]) == 0:
|
|
continue
|
|
|
|
buffer = chunk if buffer is None else _concat(buffer, chunk)
|
|
while len(buffer["event_id"]) >= bs:
|
|
piece = {k: v[:bs] for k, v in buffer.items()}
|
|
buffer = {k: v[bs:] for k, v in buffer.items()}
|
|
_process(piece)
|
|
|
|
if buffer is not None and len(buffer["event_id"]) > 0:
|
|
_process(buffer)
|
|
|
|
bar.close()
|
|
if writer is not None:
|
|
writer.close()
|
|
|
|
ref_path = _write_prediction_ref(checkpoint, pred_uuid, out, dataset_path, comment)
|
|
typer.echo(f"reference: {ref_path}")
|
|
|
|
if skipped:
|
|
codes = ", ".join(f"{pdg} ({count})" for pdg, count in sorted(unknown_pdg_counts.items()))
|
|
typer.echo(
|
|
f"warning: skipped {skipped:,} row(s) with unknown PDG code(s): {codes}",
|
|
err=True,
|
|
)
|
|
typer.echo(f"wrote {total:,} rows → {out}")
|
|
|
|
|
|
def _seed_from_data(files: list[Path], n_events: int | None) -> dict[str, np.ndarray]:
|
|
"""Pick each event's primary entry state (argmax-pre_E row) as a shower seed.
|
|
|
|
Streams conditioning columns and keeps the highest-pre_E step per event_id —
|
|
the codebase's convention for the primary (a secondary always carries less
|
|
energy than its parent). See giant/analysis/reduce.py:entry_axis.
|
|
"""
|
|
best_E: dict[int, float] = {}
|
|
best: dict[int, tuple] = {}
|
|
for file_idx, path in enumerate(files):
|
|
for chunk in iter_cond_chunks(path, offset=event_id_offset(file_idx)):
|
|
ev = chunk["event_id"]
|
|
pe = chunk["pre_E"]
|
|
for i in range(len(ev)):
|
|
e = int(ev[i])
|
|
if pe[i] > best_E.get(e, -np.inf):
|
|
best_E[e] = float(pe[i])
|
|
best[e] = (
|
|
int(chunk["pdg"][i]),
|
|
chunk["pre_pos"][i].astype(np.float64),
|
|
float(pe[i]),
|
|
chunk["pre_dir"][i].astype(np.float64),
|
|
)
|
|
event_ids = sorted(best)
|
|
if n_events is not None:
|
|
event_ids = event_ids[:n_events]
|
|
if not event_ids:
|
|
raise ValueError("no events found to seed from")
|
|
|
|
return {
|
|
"event_id": np.array(event_ids, dtype=np.int64),
|
|
"pdg": np.array([best[e][0] for e in event_ids], dtype=np.int64),
|
|
"pre_pos": np.stack([best[e][1] for e in event_ids]),
|
|
"pre_E": np.array([best[e][2] for e in event_ids], dtype=np.float64),
|
|
"pre_dir": np.stack([best[e][3] for e in event_ids]),
|
|
}
|
|
|
|
|
|
@app.command()
|
|
def rollout(
|
|
data: Annotated[Path, typer.Argument(help="Parquet file/dir to seed showers from (real events)")],
|
|
checkpoint: Annotated[
|
|
Path,
|
|
typer.Option("--checkpoint", "-c", help="Checkpoint .pt (best.pt/last.pt)"),
|
|
],
|
|
geometry: Annotated[
|
|
Path,
|
|
typer.Option(
|
|
"--geometry",
|
|
"-g",
|
|
help="Geometry oracle .pkl (dwarf build-geometry-oracle)",
|
|
),
|
|
],
|
|
energy_cutoff: Annotated[
|
|
float,
|
|
typer.Option(
|
|
"--energy-cutoff",
|
|
help="Stop a track when its energy drops below this [MeV]",
|
|
),
|
|
] = 0.1,
|
|
max_steps: Annotated[int, typer.Option("--max-steps", help="Max steps per individual track")] = 1000,
|
|
steps: Annotated[
|
|
int,
|
|
typer.Option(
|
|
"--steps",
|
|
"-s",
|
|
help="Flow matching ODE steps per model call (ignored for a wgan checkpoint)",
|
|
),
|
|
] = 10,
|
|
weights: Annotated[
|
|
Weights,
|
|
typer.Option(
|
|
"--weights",
|
|
help="raw: the live training weights. ema: the EMA shadow copy "
|
|
"(see --ema-decay in `giant train`) — usually cleaner samples, "
|
|
"requires a checkpoint trained with EMA enabled.",
|
|
),
|
|
] = Weights.raw,
|
|
batch_size: Annotated[int, typer.Option("--batch-size", "-b", help="Tracks stepped per model forward")] = 4096,
|
|
max_tracks_per_event: Annotated[
|
|
Optional[int],
|
|
typer.Option(
|
|
"--max-tracks-per-event",
|
|
help="Safety cap on tracks per shower (sub-cap secondaries deposit in place)",
|
|
),
|
|
] = None,
|
|
escape_threshold: Annotated[
|
|
Optional[float],
|
|
typer.Option(
|
|
"--escape-threshold",
|
|
help="Override the oracle's NN-distance escape threshold [mm]",
|
|
),
|
|
] = None,
|
|
n_events: Annotated[Optional[int], typer.Option("--n-events", help="Cap number of seed events")] = None,
|
|
device: Annotated[Optional[str], typer.Option("--device", "-d", help="cpu | cuda | mps (auto)")] = None,
|
|
out: Annotated[Optional[Path], typer.Option("--out", "-o", help="Output steps parquet")] = None,
|
|
seed: Annotated[
|
|
Optional[int],
|
|
typer.Option("--seed", help="Torch/numpy seed for reproducibility"),
|
|
] = None,
|
|
) -> None:
|
|
"""Roll the surrogate forward into full showers (autoregressive)."""
|
|
if seed is not None:
|
|
torch.manual_seed(seed)
|
|
np.random.seed(seed)
|
|
|
|
_device = torch.device(device) if device else gconfig.auto_device()
|
|
typer.echo(f"device: {_device}")
|
|
|
|
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})")
|
|
|
|
training_cfg = gconfig.load_checkpoint_config(checkpoint)
|
|
|
|
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
|
|
sec_type_topn_map = ctx.sec_type_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})")
|
|
|
|
files = find_parquet_files(data)
|
|
seeds = _seed_from_data(files, n_events)
|
|
typer.echo(f"seeded {len(seeds['event_id']):,} shower(s)")
|
|
|
|
out, dataset_path, pred_uuid = _resolve_prediction_output(data, out)
|
|
out.parent.mkdir(parents=True, exist_ok=True)
|
|
|
|
# Written incrementally as each batch of steps is produced, rather than
|
|
# buffering the whole run (which scales with n_events * max_steps *
|
|
# avg_tracks_per_event) — mirrors the row-group streaming `giant predict`
|
|
# already does on its input side.
|
|
writer: pq.ParquetWriter | None = None
|
|
|
|
def _write_chunk(row: dict[str, np.ndarray]) -> None:
|
|
nonlocal writer
|
|
table = pa.table(row)
|
|
if writer is None:
|
|
table = table.replace_schema_metadata(
|
|
{
|
|
PREDICT_COORD_METADATA_KEY: ROLLOUT_COORD_VALUE,
|
|
PREDICT_SCHEMA_VERSION_KEY: PREDICT_SCHEMA_VERSION,
|
|
}
|
|
)
|
|
writer = pq.ParquetWriter(out, table.schema)
|
|
writer.write_table(table)
|
|
|
|
# Only meaningful under particle_type.target="embedding" — a
|
|
# no-op collector otherwise, cheaper than branching the call itself.
|
|
l1_dist_collector = L1DistCollector()
|
|
|
|
summary = run_rollout(
|
|
model,
|
|
sec_decoder,
|
|
oracle,
|
|
seeds,
|
|
cond_norm,
|
|
tgt_norm,
|
|
sec_phys_norm,
|
|
pdg_map,
|
|
mat_map,
|
|
energy_cutoff=energy_cutoff,
|
|
max_steps=max_steps,
|
|
steps=steps,
|
|
batch_size=batch_size,
|
|
device=_device,
|
|
max_tracks_per_event=max_tracks_per_event,
|
|
escape_threshold=escape_threshold,
|
|
on_chunk=_write_chunk,
|
|
particle_conditioning=particle_conditioning,
|
|
material_conditioning=material_conditioning,
|
|
pdg_topn_map=pdg_topn_map,
|
|
mat_topn_map=mat_topn_map,
|
|
sec_type_topn_map=sec_type_topn_map,
|
|
other_policy=other_policy,
|
|
seed=seed,
|
|
stage1_ddpm_steps=stage1_ddpm_steps,
|
|
stage2_ddpm_steps=stage2_ddpm_steps,
|
|
l1_dist_collector=l1_dist_collector,
|
|
)
|
|
if writer is not None:
|
|
writer.close()
|
|
|
|
l1_summary = l1_dist_collector.summary()
|
|
|
|
ref_path = _write_prediction_ref(checkpoint, pred_uuid, out, dataset_path)
|
|
ref = yaml.safe_load(ref_path.read_text())
|
|
ref.update(
|
|
{
|
|
"kind": "rollout",
|
|
"geometry_oracle": str(geometry.resolve()),
|
|
"energy_cutoff": energy_cutoff,
|
|
"max_steps": max_steps,
|
|
"steps": steps,
|
|
"max_tracks_per_event": max_tracks_per_event,
|
|
"escape_threshold": escape_threshold,
|
|
"n_events": n_events,
|
|
"n_seed_events": int(len(seeds["event_id"])),
|
|
"weights": weights.value,
|
|
"batch_size": batch_size,
|
|
"device": str(_device),
|
|
"rollout_seed": seed,
|
|
"n_rows": summary["n_rows"],
|
|
"termination_reason_counts": summary["termination_reason_counts"],
|
|
# Diagnostic — only present under
|
|
# stage2_model.particle_type.target="embedding"; omitted (not
|
|
# written as null) otherwise, so giant.analysis can tell "not
|
|
# applicable to this checkpoint" apart from "collector empty".
|
|
**({"type_embedding_l1_dist": l1_summary} if l1_summary is not None else {}),
|
|
# Full architecture spec baked into the checkpoint — includes the
|
|
# entire router sub-dict, not just a hand-picked subset, so any
|
|
# 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": 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", {})),
|
|
"training_meta": dict(training_cfg.get("meta", {})),
|
|
}
|
|
)
|
|
ref_path.write_text(yaml.dump(ref, default_flow_style=False, sort_keys=False))
|
|
|
|
typer.echo(f"wrote {summary['n_rows']:,} step rows → {out}")
|
|
typer.echo(f"terminations: {summary['termination_reason_counts']}")
|
|
typer.echo(f"reference: {ref_path}")
|
|
|
|
|
|
analyze_app = typer.Typer(
|
|
no_args_is_help=True,
|
|
help="Rollout-vs-reference analysis: parallel compute on HTCondor + local render.",
|
|
)
|
|
app.add_typer(analyze_app, name="analyze")
|
|
|
|
|
|
@analyze_app.command("prep")
|
|
def analyze_prep(
|
|
rollout_yaml: Annotated[
|
|
Path,
|
|
typer.Argument(help="giant rollout YAML sidecar (names the rollout + reference files)"),
|
|
],
|
|
run_dir: Annotated[
|
|
Optional[Path],
|
|
typer.Option(
|
|
"--run-dir",
|
|
"-o",
|
|
help="Override the run directory (default: <cwd>/analysis_runs/analysis_<id>)",
|
|
),
|
|
] = None,
|
|
n_energy_bins: Annotated[int, typer.Option("--energy-bins")] = 4,
|
|
n_marginal_bins: Annotated[int, typer.Option("--bins")] = 50,
|
|
top_k_pdg: Annotated[int, typer.Option("--top-pdg")] = 6,
|
|
chunks: Annotated[
|
|
int,
|
|
typer.Option("--chunks", help="Split each plot's data into this many event_id chunks"),
|
|
] = 1,
|
|
) -> None:
|
|
"""Read the rollout YAML → shared.json + run_meta.json in the run directory."""
|
|
from giant.analysis import prep
|
|
|
|
path = prep(
|
|
rollout_yaml,
|
|
run_dir,
|
|
n_chunks=chunks,
|
|
default_base=Path.cwd() / "analysis_runs",
|
|
n_energy_bins=n_energy_bins,
|
|
n_marginal_bins=n_marginal_bins,
|
|
top_k_pdg=top_k_pdg,
|
|
)
|
|
typer.echo(f"run directory: {path}")
|
|
|
|
|
|
@analyze_app.command("compute-one")
|
|
def analyze_compute_one(
|
|
id: Annotated[str, typer.Option("--id", help="Catalog plot id (see `analyze list`)")],
|
|
run_dir: Annotated[Path, typer.Option("--run-dir", help="Run directory from `analyze prep`")],
|
|
chunk: Annotated[int, typer.Option("--chunk", help="Chunk index (see `analyze prep --chunks`)")] = 0,
|
|
) -> None:
|
|
"""Run one (plot, chunk)'s streaming reduction (this is what each condor job runs)."""
|
|
from giant.analysis import compute_one
|
|
|
|
path = compute_one(id, run_dir, chunk_index=chunk)
|
|
typer.echo(f"wrote {path}")
|
|
|
|
|
|
@analyze_app.command("merge-one")
|
|
def analyze_merge_one(
|
|
id: Annotated[str, typer.Option("--id", help="Catalog plot id (see `analyze list`)")],
|
|
run_dir: Annotated[Path, typer.Option("--run-dir", help="Run directory from `analyze prep`")],
|
|
) -> None:
|
|
"""Merge one plot's chunk partials into its final reduced JSON.
|
|
|
|
Runs automatically as part of `analyze render`; useful standalone to
|
|
debug a specific plot without re-rendering everything.
|
|
"""
|
|
from giant.analysis import merge_one
|
|
|
|
path = merge_one(id, run_dir)
|
|
typer.echo(f"wrote {path}")
|
|
|
|
|
|
@analyze_app.command("list")
|
|
def analyze_list() -> None:
|
|
"""Print every catalog plot id."""
|
|
from giant.analysis import catalog_ids
|
|
|
|
for pid in catalog_ids():
|
|
typer.echo(pid)
|
|
|
|
|
|
@analyze_app.command("render")
|
|
def analyze_render(
|
|
run_dir: Annotated[Path, typer.Argument(help="Run directory from `analyze prep` (holds reduced/)")],
|
|
gallery: Annotated[
|
|
bool,
|
|
typer.Option("--gallery/--no-gallery", help="Run `gallery generate` after rendering"),
|
|
] = False,
|
|
) -> None:
|
|
"""Render reduced artifacts to styled PDFs + gallery metadata (local; needs LaTeX)."""
|
|
from giant.analysis.render import render_run
|
|
|
|
pdfs = render_run(run_dir, run_gallery=gallery)
|
|
typer.echo(f"rendered {len(pdfs)} plots → {Path(run_dir) / 'plots'}")
|
|
|
|
|
|
@analyze_app.command("submit")
|
|
def analyze_submit(
|
|
rollout_yaml: Annotated[Path, typer.Argument(help="giant rollout YAML sidecar")],
|
|
accounting_group: Annotated[str, typer.Option("--accounting-group")],
|
|
run_dir: Annotated[
|
|
Optional[Path],
|
|
typer.Option(
|
|
"--run-dir",
|
|
"-o",
|
|
help="Override the run directory (default: <cwd>/analysis_runs/analysis_<id>)",
|
|
),
|
|
] = None,
|
|
docker_image: Annotated[str, typer.Option("--docker-image")] = "cverstege/alma9-gridjob",
|
|
request_memory: Annotated[int, typer.Option("--request-memory", help="MB")] = 8192,
|
|
remote: Annotated[
|
|
bool,
|
|
typer.Option("--remote/--local", help="+RemoteJob vs ProvidesETPResources"),
|
|
] = False,
|
|
chunks: Annotated[
|
|
int,
|
|
typer.Option(
|
|
"--chunks",
|
|
help="Split each plot's data into this many event_id chunks/jobs",
|
|
),
|
|
] = 1,
|
|
n_energy_bins: Annotated[int, typer.Option("--energy-bins")] = 4,
|
|
n_marginal_bins: Annotated[int, typer.Option("--bins")] = 50,
|
|
top_k_pdg: Annotated[int, typer.Option("--top-pdg")] = 6,
|
|
dry_run: Annotated[bool, typer.Option("--dry-run", help="Write files but don't condor_submit")] = False,
|
|
) -> None:
|
|
"""prep + write the HTCondor submit description (one job per plot x chunk), then submit."""
|
|
import subprocess
|
|
|
|
from giant.analysis import SubmitConfig, prep, write_submit
|
|
|
|
path = prep(
|
|
rollout_yaml,
|
|
run_dir,
|
|
n_chunks=chunks,
|
|
default_base=Path.cwd() / "analysis_runs",
|
|
n_energy_bins=n_energy_bins,
|
|
n_marginal_bins=n_marginal_bins,
|
|
top_k_pdg=top_k_pdg,
|
|
)
|
|
cfg = SubmitConfig(
|
|
run_dir=path,
|
|
accounting_group=accounting_group,
|
|
repo_dir=Path.cwd(),
|
|
docker_image=docker_image,
|
|
request_memory_mb=request_memory,
|
|
remote=remote,
|
|
n_chunks=chunks,
|
|
)
|
|
sub = write_submit(cfg)
|
|
typer.echo(f"run directory: {path}")
|
|
typer.echo(f"wrote submit description: {sub}")
|
|
if dry_run:
|
|
typer.echo("dry-run: not submitting")
|
|
return
|
|
subprocess.run(["condor_submit", str(sub)], check=True)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
app()
|