2 Commits

Author SHA1 Message Date
lars ecd6347fca chore: raise pyarrow ceiling to <26, bump to 25.0.1
CI / Sync project version with tag (hand-pushed tags only) (pull_request) Has been skipped
CI / Publish package to Gitea package registry (pull_request) Has been skipped
CI / Type check (ty) (pull_request) Successful in 2m28s
CI / Lint (ruff check) (pull_request) Successful in 2m28s
CI / Format (ruff format) (pull_request) Successful in 2m28s
CI / Tests (pull_request) Successful in 6m2s
CI / Release (bump, changelog, badges, tag) on merge to master (pull_request) Has been skipped
25.0.0's changelog has no breaking changes for our usage (streaming
parquet reads via polars/pyarrow) — the only Python-relevant
deprecation is the `feather` module, which giant doesn't use. Tests,
ruff, and ty all pass against 25.0.1.

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01TMdZFqXXig7i3XkirSUxef
2026-09-04 14:12:55 +02:00
lars c984d0a19d chore: bump uv.lock and fix ruff 0.16 default-rule lint findings
CI / Sync project version with tag (hand-pushed tags only) (pull_request) Has been skipped
CI / Publish package to Gitea package registry (pull_request) Has been skipped
CI / Format (ruff format) (pull_request) Successful in 53s
CI / Type check (ty) (pull_request) Successful in 57s
CI / Lint (ruff check) (pull_request) Successful in 41s
CI / Tests (pull_request) Successful in 8m20s
CI / Release (bump, changelog, badges, tag) on merge to master (pull_request) Has been skipped
uv.lock was stale (ty 0.0.50 -> 0.0.78, ruff 0.15 -> 0.16, polars, numpy,
typer, wandb, pytest, and others), all within existing pyproject.toml
bounds. ruff 0.16 widened its default rule selection, taking this repo
from 0 to 274 lint errors under the same config; --fix handled most of
it (import sorting, Optional[X] -> X | None, ...), and the remainder
(unused unpacked variables, dict()-as-literal, subprocess.run without
explicit check=, a couple of intentional broad excepts/naive datetimes)
were fixed or annotated by hand. Also fixes a real type-narrowing gap
ty 0.0.78 caught in test_config_consumed_keys.py's `or`-combined
isinstance check.

torch stays pinned to 2.3.x (deliberate, see CLAUDE.md); pyarrow's <25
ceiling is left as a separate decision.

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01TMdZFqXXig7i3XkirSUxef
2026-09-04 14:09:29 +02:00
61 changed files with 1700 additions and 1448 deletions
+11 -11
View File
@@ -32,27 +32,27 @@ from giant.analysis.runtime_estimate import RUNTIME_SAFETY_MARGIN, estimate_runt
from giant.analysis.sources import RolloutSpec, Side from giant.analysis.sources import RolloutSpec, Side
__all__ = [ __all__ = [
"build_catalog", "RUNTIME_SAFETY_MARGIN",
"catalog_ids", "Context",
"get_spec",
"LoadedRollout", "LoadedRollout",
"Partial",
"Reduced",
"RolloutSpec",
"RunMeta", "RunMeta",
"Side",
"SubmitConfig", "SubmitConfig",
"build_catalog",
"build_context",
"catalog_ids",
"compute_one", "compute_one",
"compute_reduced", "compute_reduced",
"derive_run_dir", "derive_run_dir",
"estimate_runtime_s",
"get_spec",
"load_rollout_yaml", "load_rollout_yaml",
"load_rollout_yamls", "load_rollout_yamls",
"merge_all", "merge_all",
"merge_one", "merge_one",
"prep", "prep",
"write_submit", "write_submit",
"Context",
"build_context",
"Partial",
"Reduced",
"RolloutSpec",
"Side",
"RUNTIME_SAFETY_MARGIN",
"estimate_runtime_s",
] ]
+2 -2
View File
@@ -38,8 +38,8 @@ variable x grouping, secondaries, ...) into concrete specs.
from __future__ import annotations from __future__ import annotations
from collections.abc import Callable
from dataclasses import asdict, dataclass from dataclasses import asdict, dataclass
from typing import Callable
import numpy as np import numpy as np
import polars as pl import polars as pl
@@ -101,7 +101,7 @@ class Bundle:
reference, reference,
ctx: Context, ctx: Context,
chunk: tuple[int, int] | None = None, chunk: tuple[int, int] | None = None,
) -> "Bundle": ) -> Bundle:
"""Open the reference + every rollout, optionally restricted to one event-disjoint chunk. """Open the reference + every rollout, optionally restricted to one event-disjoint chunk.
``chunk = (chunk_index, n_chunks)`` filters every side to ``chunk = (chunk_index, n_chunks)`` filters every side to
+1 -1
View File
@@ -228,7 +228,7 @@ class RunMeta:
Path(path).write_text(json.dumps(self.__dict__, indent=2)) Path(path).write_text(json.dumps(self.__dict__, indent=2))
@classmethod @classmethod
def load(cls, path: str | Path) -> "RunMeta": def load(cls, path: str | Path) -> RunMeta:
return cls(**json.loads(Path(path).read_text())) return cls(**json.loads(Path(path).read_text()))
+1 -1
View File
@@ -50,7 +50,7 @@ class Context:
Path(path).write_text(json.dumps(asdict(self), indent=2)) Path(path).write_text(json.dumps(asdict(self), indent=2))
@classmethod @classmethod
def load(cls, path: str | Path) -> "Context": def load(cls, path: str | Path) -> Context:
d = json.loads(Path(path).read_text()) d = json.loads(Path(path).read_text())
d["var_ranges"] = {k: tuple(v) for k, v in d["var_ranges"].items()} d["var_ranges"] = {k: tuple(v) for k, v in d["var_ranges"].items()}
d["sec_energy_range"] = tuple(d["sec_energy_range"]) d["sec_energy_range"] = tuple(d["sec_energy_range"])
+1 -1
View File
@@ -46,7 +46,7 @@ def pdg_label(code: int) -> str:
def material_label(name: str) -> str: def material_label(name: str) -> str:
"""Display label for a Geant4 material, dropping the ``G4_`` prefix.""" """Display label for a Geant4 material, dropping the ``G4_`` prefix."""
return name[3:] if name.startswith("G4_") else name return name.removeprefix("G4_")
def energy_bin_edges(incident_E: np.ndarray, n_bins: int = 4) -> np.ndarray: def energy_bin_edges(incident_E: np.ndarray, n_bins: int = 4) -> np.ndarray:
+2 -2
View File
@@ -46,7 +46,7 @@ class Reduced:
Path(path).write_text(json.dumps(asdict(self))) Path(path).write_text(json.dumps(asdict(self)))
@classmethod @classmethod
def load(cls, path: str | Path) -> "Reduced": def load(cls, path: str | Path) -> Reduced:
return cls(**json.loads(Path(path).read_text())) return cls(**json.loads(Path(path).read_text()))
@@ -70,5 +70,5 @@ class Partial:
Path(path).write_text(json.dumps(asdict(self))) Path(path).write_text(json.dumps(asdict(self)))
@classmethod @classmethod
def load(cls, path: str | Path) -> "Partial": def load(cls, path: str | Path) -> Partial:
return cls(**json.loads(Path(path).read_text())) return cls(**json.loads(Path(path).read_text()))
+1 -1
View File
@@ -27,8 +27,8 @@ from pathlib import Path
import numpy as np import numpy as np
import plotstyle as ps import plotstyle as ps
from matplotlib.colors import LogNorm
import yaml import yaml
from matplotlib.colors import LogNorm
from giant.analysis.reduced import Reduced from giant.analysis.reduced import Reduced
+6 -6
View File
@@ -58,10 +58,10 @@ _COLS = (
@dataclass @dataclass
class _RouterHandle: class _RouterHandle:
router: "torch.nn.Module" router: torch.nn.Module
pdg_map: dict[int, int] pdg_map: dict[int, int]
mat_map: dict[str, int] mat_map: dict[str, int]
cond_normalizer: "Normalizer" cond_normalizer: Normalizer
particle_conditioning: str particle_conditioning: str
material_conditioning: str material_conditioning: str
router_type: str router_type: str
@@ -233,7 +233,7 @@ def _gating_entry(checkpoint: str | Path | None, r_phys: pl.LazyFrame, t_phys: p
return {"router_type": handle.router_type, "n_experts": handle.router.n_experts, **sides} return {"router_type": handle.router_type, "n_experts": handle.router.n_experts, **sides}
def compute_router_gating(rollouts: dict[str, "RolloutSide"], t_phys: pl.LazyFrame, seed: int = 0) -> Reduced: def compute_router_gating(rollouts: dict[str, RolloutSide], t_phys: pl.LazyFrame, seed: int = 0) -> Reduced:
"""`Reduced` for the router-gating figure: one panel-pair per rollout with """`Reduced` for the router-gating figure: one panel-pair per rollout with
an enabled MoE router, or an explanatory note if none of them have one.""" an enabled MoE router, or an explanatory note if none of them have one."""
series = {} series = {}
@@ -289,7 +289,7 @@ def _specialization_entry(
} }
def compute_router_specialization(rollouts: dict[str, "RolloutSide"], t_phys: pl.LazyFrame, seed: int = 0) -> Reduced: def compute_router_specialization(rollouts: dict[str, RolloutSide], t_phys: pl.LazyFrame, seed: int = 0) -> Reduced:
"""`Reduced` for the router-specialization figure, one curve per rollout with """`Reduced` for the router-specialization figure, one curve per rollout with
an enabled MoE router (see `_specialization_entry`).""" an enabled MoE router (see `_specialization_entry`)."""
series = {} series = {}
@@ -331,7 +331,7 @@ def _share_by_pdg_entry(
def compute_router_share_by_pdg( def compute_router_share_by_pdg(
rollouts: dict[str, "RolloutSide"], t_phys: pl.LazyFrame, top_pdgs: list[int], seed: int = 0 rollouts: dict[str, RolloutSide], t_phys: pl.LazyFrame, top_pdgs: list[int], seed: int = 0
) -> Reduced: ) -> Reduced:
"""`Reduced` for the router expert-share-by-species figure, one panel-pair """`Reduced` for the router expert-share-by-species figure, one panel-pair
per rollout with an enabled MoE router.""" per rollout with an enabled MoE router."""
@@ -375,7 +375,7 @@ def _share_by_process_entry(checkpoint: str | Path | None, t_phys: pl.LazyFrame,
def compute_router_share_by_process( def compute_router_share_by_process(
rollouts: dict[str, "RolloutSide"], t_phys: pl.LazyFrame, seed: int = 0, top_k: int = _TOP_K_PROCESS rollouts: dict[str, RolloutSide], t_phys: pl.LazyFrame, seed: int = 0, top_k: int = _TOP_K_PROCESS
) -> Reduced: ) -> Reduced:
"""Stacked-bar share of each physics process dispatched to each expert, one """Stacked-bar share of each physics process dispatched to each expert, one
panel per rollout checkpoint with an enabled MoE router. panel per rollout checkpoint with an enabled MoE router.
+1 -1
View File
@@ -35,7 +35,7 @@ _NOTE_NOT_APPLICABLE = (
) )
def compute_type_embedding_l1_distance(rollouts: dict[str, "RolloutSide"]) -> Reduced: def compute_type_embedding_l1_distance(rollouts: dict[str, RolloutSide]) -> Reduced:
"""`Reduced` for the type-embedding-distance figure: one series per rollout """`Reduced` for the type-embedding-distance figure: one series per rollout
whose checkpoint populated the diagnostic, or an explanatory note if none did. whose checkpoint populated the diagnostic, or an explanatory note if none did.
+127 -127
View File
@@ -1,16 +1,15 @@
from __future__ import annotations from __future__ import annotations
from collections import Counter
from datetime import datetime, timezone
from enum import Enum
import math import math
from pathlib import Path
import re import re
from typing import TYPE_CHECKING, Optional, cast
import uuid as uuid_mod import uuid as uuid_mod
from collections import Counter
from datetime import UTC, datetime
from enum import Enum
from pathlib import Path
from typing import TYPE_CHECKING, Annotated, cast
import typer import typer
from typing_extensions import Annotated
if TYPE_CHECKING: if TYPE_CHECKING:
import numpy as np import numpy as np
@@ -120,7 +119,7 @@ def _parse_router_axis_flags(specs: list[str]) -> dict[str, object]:
return out return out
def _parse_set_flags(specs: Optional[list[str]]) -> dict[str, object]: def _parse_set_flags(specs: list[str] | None) -> dict[str, object]:
"""Parse repeated `--set dotted.path=value` flags into a dict, typing """Parse repeated `--set dotted.path=value` flags into a dict, typing
each value with `_coerce_scalar` the same way a TOML file's native types each value with `_coerce_scalar` the same way a TOML file's native types
would arrive. Validation against the inference-safe allowlist happens would arrive. Validation against the inference-safe allowlist happens
@@ -196,7 +195,7 @@ def _write_prediction_ref(
"output": str(out), "output": str(out),
"dataset": str(dataset_path), "dataset": str(dataset_path),
"checkpoint": str(checkpoint.resolve()), "checkpoint": str(checkpoint.resolve()),
"timestamp": datetime.now(timezone.utc).isoformat(), "timestamp": datetime.now(UTC).isoformat(),
} }
if comment is not None: if comment is not None:
ref["comment"] = comment ref["comment"] = comment
@@ -285,52 +284,52 @@ class Weights(str, Enum):
def train( def train(
data: Annotated[Path, typer.Argument(help="Parquet file or directory of parquet files")], data: Annotated[Path, typer.Argument(help="Parquet file or directory of parquet files")],
config: Annotated[ config: Annotated[
Optional[Path], Path | None,
typer.Option("--config", "-c", help="TOML config file (overridden by explicit flags)"), typer.Option("--config", "-c", help="TOML config file (overridden by explicit flags)"),
] = None, ] = None,
mode: Annotated[ mode: Annotated[
Optional[Mode], Mode | None,
typer.Option("--mode", "-m", help="Generative model: flow matching or DDPM"), typer.Option("--mode", "-m", help="Generative model: flow matching or DDPM"),
] = None, ] = None,
epochs: Annotated[Optional[int], typer.Option("--epochs", "-e")] = None, epochs: Annotated[int | None, typer.Option("--epochs", "-e")] = None,
batch_size: Annotated[ batch_size: Annotated[
Optional[str], str | None,
typer.Option( typer.Option(
"--batch-size", "--batch-size",
"-b", "-b",
help="Integer, or 'auto' to estimate from free GPU memory (cuda devices only)", help="Integer, or 'auto' to estimate from free GPU memory (cuda devices only)",
), ),
] = None, ] = None,
lr: Annotated[Optional[float], typer.Option("--lr", "-l")] = None, lr: Annotated[float | None, typer.Option("--lr", "-l")] = None,
weight_decay: Annotated[ weight_decay: Annotated[
Optional[float], float | None,
typer.Option("--weight-decay", "-W", help="AdamW weight decay (default: 0.01)"), typer.Option("--weight-decay", "-W", help="AdamW weight decay (default: 0.01)"),
] = None, ] = None,
ema_decay: Annotated[ ema_decay: Annotated[
Optional[float], float | None,
typer.Option( typer.Option(
"--ema-decay", "--ema-decay",
help="EMA decay for a shadow copy of the model weights, saved " help="EMA decay for a shadow copy of the model weights, saved "
"alongside the raw weights in checkpoints (0 disables; default: 0.9999)", "alongside the raw weights in checkpoints (0 disables; default: 0.9999)",
), ),
] = None, ] = None,
warmup_epochs: Annotated[Optional[int], typer.Option("--warmup-epochs", "-w")] = None, warmup_epochs: Annotated[int | None, typer.Option("--warmup-epochs", "-w")] = None,
hidden_dim: Annotated[Optional[int], typer.Option("--hidden-dim", "-H")] = None, hidden_dim: Annotated[int | None, typer.Option("--hidden-dim", "-H")] = None,
n_blocks: Annotated[Optional[int], typer.Option("--n-blocks", "-n")] = None, n_blocks: Annotated[int | None, typer.Option("--n-blocks", "-n")] = None,
emb_dim: Annotated[Optional[int], typer.Option("--emb-dim", "-E")] = None, emb_dim: Annotated[int | None, typer.Option("--emb-dim", "-E")] = None,
dropout: Annotated[ dropout: Annotated[
Optional[float], float | None,
typer.Option("--dropout", "-d", help="Dropout probability in ResBlocks (default: 0.1)"), typer.Option("--dropout", "-d", help="Dropout probability in ResBlocks (default: 0.1)"),
] = None, ] = None,
stage1_generator: Annotated[ stage1_generator: Annotated[
Optional[Mode], Mode | None,
typer.Option( typer.Option(
"--stage1-generator", "--stage1-generator",
help="Stage 1's generative objective — overrides --mode for stage 1 only", help="Stage 1's generative objective — overrides --mode for stage 1 only",
), ),
] = None, ] = None,
stage1_hidden_dim: Annotated[ stage1_hidden_dim: Annotated[
Optional[int], int | None,
typer.Option( typer.Option(
"--stage1-hidden-dim", "--stage1-hidden-dim",
help="Overrides --hidden-dim for stage 1 only (same effect today; " help="Overrides --hidden-dim for stage 1 only (same effect today; "
@@ -339,15 +338,15 @@ def train(
), ),
] = None, ] = None,
stage1_n_res_blocks: Annotated[ stage1_n_res_blocks: Annotated[
Optional[int], int | None,
typer.Option("--stage1-n-res-blocks", help="Overrides --n-blocks for stage 1 only"), typer.Option("--stage1-n-res-blocks", help="Overrides --n-blocks for stage 1 only"),
] = None, ] = None,
stage1_dropout: Annotated[ stage1_dropout: Annotated[
Optional[float], float | None,
typer.Option("--stage1-dropout", help="Overrides --dropout for stage 1 only"), typer.Option("--stage1-dropout", help="Overrides --dropout for stage 1 only"),
] = None, ] = None,
stage2_generator: Annotated[ stage2_generator: Annotated[
Optional[Mode], Mode | None,
typer.Option( typer.Option(
"--stage2-generator", "--stage2-generator",
help="Stage 2's generative objective — overrides --mode for stage 2 " help="Stage 2's generative objective — overrides --mode for stage 2 "
@@ -356,19 +355,19 @@ def train(
), ),
] = None, ] = None,
stage2_hidden_dim: Annotated[ stage2_hidden_dim: Annotated[
Optional[int], int | None,
typer.Option("--stage2-hidden-dim", help="Stage 2 trunk width"), typer.Option("--stage2-hidden-dim", help="Stage 2 trunk width"),
] = None, ] = None,
stage2_n_res_blocks: Annotated[ stage2_n_res_blocks: Annotated[
Optional[int], int | None,
typer.Option("--stage2-n-res-blocks", help="Stage 2 trunk depth"), typer.Option("--stage2-n-res-blocks", help="Stage 2 trunk depth"),
] = None, ] = None,
stage2_dropout: Annotated[ stage2_dropout: Annotated[
Optional[float], float | None,
typer.Option("--stage2-dropout", help="Dropout inside stage 2's ResBlocks"), typer.Option("--stage2-dropout", help="Dropout inside stage 2's ResBlocks"),
] = None, ] = None,
stage2_decoder: Annotated[ stage2_decoder: Annotated[
Optional[Decoder], Decoder | None,
typer.Option( typer.Option(
"--stage2-decoder", "--stage2-decoder",
help="one_shot: predict all k_max secondary slots at once (v0.2 " help="one_shot: predict all k_max secondary slots at once (v0.2 "
@@ -377,7 +376,7 @@ def train(
), ),
] = None, ] = None,
stage2_k_max: Annotated[ stage2_k_max: Annotated[
Optional[int], int | None,
typer.Option( typer.Option(
"--stage2-k-max", "--stage2-k-max",
help="Maximum secondary slots (fixed width under one_shot, a " help="Maximum secondary slots (fixed width under one_shot, a "
@@ -385,14 +384,14 @@ def train(
), ),
] = None, ] = None,
stage2_context_dim: Annotated[ stage2_context_dim: Annotated[
Optional[int], int | None,
typer.Option( typer.Option(
"--stage2-context-dim", "--stage2-context-dim",
help="Width of the projected stage-1 outcome fed into stage 2's conditioning (default: 64)", help="Width of the projected stage-1 outcome fed into stage 2's conditioning (default: 64)",
), ),
] = None, ] = None,
stage2_stage1_context: Annotated[ stage2_stage1_context: Annotated[
Optional[Stage1Context], Stage1Context | None,
typer.Option( typer.Option(
"--stage2-stage1-context", "--stage2-stage1-context",
help="What stage 2 conditions on during training: 'truth' (the " help="What stage 2 conditions on during training: 'truth' (the "
@@ -402,7 +401,7 @@ def train(
), ),
] = None, ] = None,
conditioning: Annotated[ conditioning: Annotated[
Optional[Conditioning], Conditioning | None,
typer.Option( typer.Option(
"--conditioning", "--conditioning",
help="Input conditioning: continuous physical properties " help="Input conditioning: continuous physical properties "
@@ -411,7 +410,7 @@ def train(
), ),
] = None, ] = None,
router: Annotated[ router: Annotated[
Optional[bool], bool | None,
typer.Option( typer.Option(
"--router/--no-router", "--router/--no-router",
help="Route both stages through a mixture of small experts " help="Route both stages through a mixture of small experts "
@@ -419,12 +418,12 @@ def train(
), ),
] = None, ] = None,
router_type: Annotated[ router_type: Annotated[
Optional[str], str | None,
typer.Option("--router-type", help="Router implementation name (see ROUTER_REGISTRY)"), typer.Option("--router-type", help="Router implementation name (see ROUTER_REGISTRY)"),
] = None, ] = None,
n_experts: Annotated[Optional[int], typer.Option("--n-experts", help="Number of routed experts")] = None, n_experts: Annotated[int | None, typer.Option("--n-experts", help="Number of routed experts")] = None,
router_axis: Annotated[ router_axis: Annotated[
Optional[list[str]], list[str] | None,
typer.Option( typer.Option(
"--router-axis", "--router-axis",
help="Composed-router axis spec 'type:key=val,key=val' (repeatable; " help="Composed-router axis spec 'type:key=val,key=val' (repeatable; "
@@ -434,59 +433,59 @@ def train(
), ),
] = None, ] = None,
n_critic: Annotated[ n_critic: Annotated[
Optional[int], int | None,
typer.Option( typer.Option(
"--n-critic", "--n-critic",
help="WGAN-GP (--mode wgan only): critic updates per generator update (default: 5)", help="WGAN-GP (--mode wgan only): critic updates per generator update (default: 5)",
), ),
] = None, ] = None,
gp_weight: Annotated[ gp_weight: Annotated[
Optional[float], float | None,
typer.Option( typer.Option(
"--gp-weight", "--gp-weight",
help="WGAN-GP (--mode wgan only): gradient-penalty coefficient (default: 10.0)", help="WGAN-GP (--mode wgan only): gradient-penalty coefficient (default: 10.0)",
), ),
] = None, ] = None,
noise_dim: Annotated[ noise_dim: Annotated[
Optional[int], int | None,
typer.Option( typer.Option(
"--noise-dim", "--noise-dim",
help="WGAN (--mode wgan only): generator input noise-vector width (default: 64)", help="WGAN (--mode wgan only): generator input noise-vector width (default: 64)",
), ),
] = None, ] = None,
critic_lr: Annotated[ critic_lr: Annotated[
Optional[float], float | None,
typer.Option( typer.Option(
"--critic-lr", "--critic-lr",
help="WGAN-GP (--mode wgan only): critic learning rate (default: same as --lr)", help="WGAN-GP (--mode wgan only): critic learning rate (default: same as --lr)",
), ),
] = None, ] = None,
stage1_n_critic: Annotated[ stage1_n_critic: Annotated[
Optional[int], int | None,
typer.Option("--stage1-n-critic", help="Overrides --n-critic for stage 1 only"), typer.Option("--stage1-n-critic", help="Overrides --n-critic for stage 1 only"),
] = None, ] = None,
stage1_gp_weight: Annotated[ stage1_gp_weight: Annotated[
Optional[float], float | None,
typer.Option("--stage1-gp-weight", help="Overrides --gp-weight for stage 1 only"), typer.Option("--stage1-gp-weight", help="Overrides --gp-weight for stage 1 only"),
] = None, ] = None,
stage1_noise_dim: Annotated[ stage1_noise_dim: Annotated[
Optional[int], int | None,
typer.Option("--stage1-noise-dim", help="Overrides --noise-dim for stage 1 only"), typer.Option("--stage1-noise-dim", help="Overrides --noise-dim for stage 1 only"),
] = None, ] = None,
stage1_critic_lr: Annotated[ stage1_critic_lr: Annotated[
Optional[float], float | None,
typer.Option("--stage1-critic-lr", help="Overrides --critic-lr for stage 1 only"), typer.Option("--stage1-critic-lr", help="Overrides --critic-lr for stage 1 only"),
] = None, ] = None,
stage2_n_critic: Annotated[ stage2_n_critic: Annotated[
Optional[int], int | None,
typer.Option("--stage2-n-critic", help="Overrides --n-critic for stage 2 only"), typer.Option("--stage2-n-critic", help="Overrides --n-critic for stage 2 only"),
] = None, ] = None,
stage2_gp_weight: Annotated[ stage2_gp_weight: Annotated[
Optional[float], float | None,
typer.Option("--stage2-gp-weight", help="Overrides --gp-weight for stage 2 only"), typer.Option("--stage2-gp-weight", help="Overrides --gp-weight for stage 2 only"),
] = None, ] = None,
stage2_noise_dim: Annotated[ stage2_noise_dim: Annotated[
Optional[int], int | None,
typer.Option( typer.Option(
"--stage2-noise-dim", "--stage2-noise-dim",
help="Overrides --noise-dim for stage 2 only; under " help="Overrides --noise-dim for stage 2 only; under "
@@ -494,39 +493,39 @@ def train(
), ),
] = None, ] = None,
stage2_critic_lr: Annotated[ stage2_critic_lr: Annotated[
Optional[float], float | None,
typer.Option("--stage2-critic-lr", help="Overrides --critic-lr for stage 2 only"), typer.Option("--stage2-critic-lr", help="Overrides --critic-lr for stage 2 only"),
] = None, ] = None,
stage1_critic_hidden_dim: Annotated[ stage1_critic_hidden_dim: Annotated[
Optional[int], int | None,
typer.Option( typer.Option(
"--stage1-critic-hidden-dim", "--stage1-critic-hidden-dim",
help="WGAN-GP (--mode wgan only): critic width for stage 1 (default: same as generator's hidden_dim)", help="WGAN-GP (--mode wgan only): critic width for stage 1 (default: same as generator's hidden_dim)",
), ),
] = None, ] = None,
stage1_critic_n_res_blocks: Annotated[ stage1_critic_n_res_blocks: Annotated[
Optional[int], int | None,
typer.Option( typer.Option(
"--stage1-critic-n-res-blocks", "--stage1-critic-n-res-blocks",
help="WGAN-GP (--mode wgan only): critic depth for stage 1 (default: same as generator's n_res_blocks)", help="WGAN-GP (--mode wgan only): critic depth for stage 1 (default: same as generator's n_res_blocks)",
), ),
] = None, ] = None,
stage2_critic_hidden_dim: Annotated[ stage2_critic_hidden_dim: Annotated[
Optional[int], int | None,
typer.Option( typer.Option(
"--stage2-critic-hidden-dim", "--stage2-critic-hidden-dim",
help="WGAN-GP (--mode wgan only): critic width for stage 2 (default: same as generator's hidden_dim)", help="WGAN-GP (--mode wgan only): critic width for stage 2 (default: same as generator's hidden_dim)",
), ),
] = None, ] = None,
stage2_critic_n_res_blocks: Annotated[ stage2_critic_n_res_blocks: Annotated[
Optional[int], int | None,
typer.Option( typer.Option(
"--stage2-critic-n-res-blocks", "--stage2-critic-n-res-blocks",
help="WGAN-GP (--mode wgan only): critic depth for stage 2 (default: same as generator's n_res_blocks)", help="WGAN-GP (--mode wgan only): critic depth for stage 2 (default: same as generator's n_res_blocks)",
), ),
] = None, ] = None,
stage1_init_from: Annotated[ stage1_init_from: Annotated[
Optional[Path], Path | None,
typer.Option( typer.Option(
"--stage1-init-from", "--stage1-init-from",
help="Checkpoint .pt to load stage 1's weights from before training starts " help="Checkpoint .pt to load stage 1's weights from before training starts "
@@ -535,33 +534,33 @@ def train(
), ),
] = None, ] = None,
stage1_freeze: Annotated[ stage1_freeze: Annotated[
Optional[bool], bool | None,
typer.Option( typer.Option(
"--stage1-freeze/--no-stage1-freeze", "--stage1-freeze/--no-stage1-freeze",
help="Never update stage 1's weights (requires --stage1-init-from, or --resume)", help="Never update stage 1's weights (requires --stage1-init-from, or --resume)",
), ),
] = None, ] = None,
stage2_init_from: Annotated[ stage2_init_from: Annotated[
Optional[Path], Path | None,
typer.Option( typer.Option(
"--stage2-init-from", "--stage2-init-from",
help="Checkpoint .pt to load stage 2's weights from before training starts (gitea #42)", help="Checkpoint .pt to load stage 2's weights from before training starts (gitea #42)",
), ),
] = None, ] = None,
stage2_freeze: Annotated[ stage2_freeze: Annotated[
Optional[bool], bool | None,
typer.Option( typer.Option(
"--stage2-freeze/--no-stage2-freeze", "--stage2-freeze/--no-stage2-freeze",
help="Never update stage 2's weights (requires --stage2-init-from, or --resume)", help="Never update stage 2's weights (requires --stage2-init-from, or --resume)",
), ),
] = None, ] = None,
val_fraction: Annotated[Optional[float], typer.Option("--val-fraction", "-f")] = None, val_fraction: Annotated[float | None, typer.Option("--val-fraction", "-f")] = None,
seed: Annotated[ seed: Annotated[
Optional[int], int | None,
typer.Option("--seed", "-s", help="Random seed for reproducibility"), typer.Option("--seed", "-s", help="Random seed for reproducibility"),
] = None, ] = None,
validate_every: Annotated[ validate_every: Annotated[
Optional[int], int | None,
typer.Option( typer.Option(
"--validate-every", "--validate-every",
"-v", "-v",
@@ -569,7 +568,7 @@ def train(
), ),
] = None, ] = None,
validate_steps: Annotated[ validate_steps: Annotated[
Optional[int], int | None,
typer.Option( typer.Option(
"--validate-steps", "--validate-steps",
"-t", "-t",
@@ -578,7 +577,7 @@ def train(
), ),
] = None, ] = None,
max_val_batches: Annotated[ max_val_batches: Annotated[
Optional[int], int | None,
typer.Option( typer.Option(
"--max-val-batches", "--max-val-batches",
help="Cap the per-epoch val-loss pass to N batches (0 = full val set every epoch; default: 200)", help="Cap the per-epoch val-loss pass to N batches (0 = full val set every epoch; default: 200)",
@@ -609,7 +608,7 @@ def train(
), ),
] = False, ] = False,
out: Annotated[ out: Annotated[
Optional[Path], Path | None,
typer.Option( typer.Option(
"--out", "--out",
"-o", "-o",
@@ -618,31 +617,31 @@ def train(
), ),
] = None, ] = None,
device: Annotated[ device: Annotated[
Optional[str], str | None,
typer.Option("--device", "-D", help="cpu | cuda | mps (default: auto)"), typer.Option("--device", "-D", help="cpu | cuda | mps (default: auto)"),
] = None, ] = None,
num_workers: Annotated[Optional[int], typer.Option("--num-workers", "-j")] = None, num_workers: Annotated[int | None, typer.Option("--num-workers", "-j")] = None,
resume: Annotated[ resume: Annotated[
Optional[Path], Path | None,
typer.Option("--resume", "-r", help="Checkpoint .pt to resume training from"), typer.Option("--resume", "-r", help="Checkpoint .pt to resume training from"),
] = None, ] = None,
wandb: Annotated[ wandb: Annotated[
Optional[bool], bool | None,
typer.Option( typer.Option(
"--wandb/--no-wandb", "--wandb/--no-wandb",
help="Log per-epoch training metrics to Weights & Biases (requires `uv sync --extra wandb`)", help="Log per-epoch training metrics to Weights & Biases (requires `uv sync --extra wandb`)",
), ),
] = None, ] = None,
wandb_project: Annotated[ wandb_project: Annotated[
Optional[str], str | None,
typer.Option("--wandb-project", help="W&B project name (default: giant)"), typer.Option("--wandb-project", help="W&B project name (default: giant)"),
] = None, ] = None,
wandb_run_name: Annotated[ wandb_run_name: Annotated[
Optional[str], str | None,
typer.Option("--wandb-run-name", help="W&B run name (default: out_dir name)"), typer.Option("--wandb-run-name", help="W&B run name (default: out_dir name)"),
] = None, ] = None,
wandb_log_every: Annotated[ wandb_log_every: Annotated[
Optional[int], int | None,
typer.Option( typer.Option(
"--wandb-log-every", "--wandb-log-every",
help="Log batch-level loss/grad_norm/lr to W&B every N optimizer " help="Log batch-level loss/grad_norm/lr to W&B every N optimizer "
@@ -650,7 +649,7 @@ def train(
), ),
] = None, ] = None,
precision: Annotated[ precision: Annotated[
Optional[str], str | None,
typer.Option( typer.Option(
"--precision", "--precision",
help="Training-step autocast precision: 'fp32' (default) or " help="Training-step autocast precision: 'fp32' (default) or "
@@ -664,7 +663,7 @@ def train(
from giant.pipeline import run_train_job from giant.pipeline import run_train_job
batch_size_auto = False batch_size_auto = False
batch_size_value: Optional[int] = None batch_size_value: int | None = None
if batch_size is not None: if batch_size is not None:
if batch_size.strip().lower() == "auto": if batch_size.strip().lower() == "auto":
batch_size_auto = True batch_size_auto = True
@@ -794,52 +793,52 @@ def train(
@app.command("new-run") @app.command("new-run")
def new_run( def new_run(
config: Annotated[ config: Annotated[
Optional[Path], Path | None,
typer.Option( typer.Option(
"--config", "--config",
"-c", "-c",
help="Base TOML to start from (default: built-in defaults)", help="Base TOML to start from (default: built-in defaults)",
), ),
] = None, ] = None,
mode: Annotated[Optional[Mode], typer.Option("--mode", "-m")] = None, mode: Annotated[Mode | None, typer.Option("--mode", "-m")] = None,
epochs: Annotated[Optional[int], typer.Option("--epochs", "-e")] = None, epochs: Annotated[int | None, typer.Option("--epochs", "-e")] = None,
batch_size: Annotated[Optional[int], typer.Option("--batch-size", "-b")] = None, batch_size: Annotated[int | None, typer.Option("--batch-size", "-b")] = None,
lr: Annotated[Optional[float], typer.Option("--lr", "-l")] = None, lr: Annotated[float | None, typer.Option("--lr", "-l")] = None,
hidden_dim: Annotated[Optional[int], typer.Option("--hidden-dim", "-H")] = None, hidden_dim: Annotated[int | None, typer.Option("--hidden-dim", "-H")] = None,
n_blocks: Annotated[Optional[int], typer.Option("--n-blocks", "-n")] = None, n_blocks: Annotated[int | None, typer.Option("--n-blocks", "-n")] = None,
emb_dim: Annotated[Optional[int], typer.Option("--emb-dim", "-E")] = None, emb_dim: Annotated[int | None, typer.Option("--emb-dim", "-E")] = None,
dropout: Annotated[Optional[float], typer.Option("--dropout", "-d")] = None, dropout: Annotated[float | None, typer.Option("--dropout", "-d")] = None,
stage1_generator: Annotated[Optional[Mode], typer.Option("--stage1-generator")] = None, stage1_generator: Annotated[Mode | None, typer.Option("--stage1-generator")] = None,
stage1_hidden_dim: Annotated[Optional[int], typer.Option("--stage1-hidden-dim")] = None, stage1_hidden_dim: Annotated[int | None, typer.Option("--stage1-hidden-dim")] = None,
stage1_n_res_blocks: Annotated[Optional[int], typer.Option("--stage1-n-res-blocks")] = None, stage1_n_res_blocks: Annotated[int | None, typer.Option("--stage1-n-res-blocks")] = None,
stage1_dropout: Annotated[Optional[float], typer.Option("--stage1-dropout")] = None, stage1_dropout: Annotated[float | None, typer.Option("--stage1-dropout")] = None,
stage2_generator: Annotated[Optional[Mode], typer.Option("--stage2-generator")] = None, stage2_generator: Annotated[Mode | None, typer.Option("--stage2-generator")] = None,
stage2_hidden_dim: Annotated[Optional[int], typer.Option("--stage2-hidden-dim")] = None, stage2_hidden_dim: Annotated[int | None, typer.Option("--stage2-hidden-dim")] = None,
stage2_n_res_blocks: Annotated[Optional[int], typer.Option("--stage2-n-res-blocks")] = None, stage2_n_res_blocks: Annotated[int | None, typer.Option("--stage2-n-res-blocks")] = None,
stage2_dropout: Annotated[Optional[float], typer.Option("--stage2-dropout")] = None, stage2_dropout: Annotated[float | None, typer.Option("--stage2-dropout")] = None,
stage2_decoder: Annotated[Optional[Decoder], typer.Option("--stage2-decoder")] = None, stage2_decoder: Annotated[Decoder | None, typer.Option("--stage2-decoder")] = None,
stage2_k_max: Annotated[Optional[int], typer.Option("--stage2-k-max")] = None, stage2_k_max: Annotated[int | None, typer.Option("--stage2-k-max")] = None,
stage2_context_dim: Annotated[Optional[int], typer.Option("--stage2-context-dim")] = None, stage2_context_dim: Annotated[int | None, typer.Option("--stage2-context-dim")] = None,
stage2_stage1_context: Annotated[Optional[Stage1Context], typer.Option("--stage2-stage1-context")] = None, stage2_stage1_context: Annotated[Stage1Context | None, typer.Option("--stage2-stage1-context")] = None,
stage1_init_from: Annotated[Optional[Path], typer.Option("--stage1-init-from")] = None, stage1_init_from: Annotated[Path | None, typer.Option("--stage1-init-from")] = None,
stage1_freeze: Annotated[Optional[bool], typer.Option("--stage1-freeze/--no-stage1-freeze")] = None, stage1_freeze: Annotated[bool | None, typer.Option("--stage1-freeze/--no-stage1-freeze")] = None,
stage2_init_from: Annotated[Optional[Path], typer.Option("--stage2-init-from")] = None, stage2_init_from: Annotated[Path | None, typer.Option("--stage2-init-from")] = None,
stage2_freeze: Annotated[Optional[bool], typer.Option("--stage2-freeze/--no-stage2-freeze")] = None, stage2_freeze: Annotated[bool | None, typer.Option("--stage2-freeze/--no-stage2-freeze")] = None,
conditioning: Annotated[Optional[Conditioning], typer.Option("--conditioning")] = None, conditioning: Annotated[Conditioning | None, typer.Option("--conditioning")] = None,
router: Annotated[Optional[bool], typer.Option("--router/--no-router")] = None, router: Annotated[bool | None, typer.Option("--router/--no-router")] = None,
router_type: Annotated[Optional[str], typer.Option("--router-type")] = None, router_type: Annotated[str | None, typer.Option("--router-type")] = None,
n_experts: Annotated[Optional[int], typer.Option("--n-experts")] = None, n_experts: Annotated[int | None, typer.Option("--n-experts")] = None,
router_axis: Annotated[Optional[list[str]], typer.Option("--router-axis")] = None, router_axis: Annotated[list[str] | None, typer.Option("--router-axis")] = None,
out: Annotated[ out: Annotated[
Optional[Path], Path | None,
typer.Option("--out", "-o", help="Run dir (default: auto from hyperparams)"), typer.Option("--out", "-o", help="Run dir (default: auto from hyperparams)"),
] = None, ] = None,
comment: Annotated[ comment: Annotated[
Optional[str], str | None,
typer.Option("--comment", help="Free-text note recorded in config.toml's meta section"), typer.Option("--comment", help="Free-text note recorded in config.toml's meta section"),
] = None, ] = None,
data: Annotated[ data: Annotated[
Optional[Path], Path | None,
typer.Option( typer.Option(
"--data", "--data",
help="Dataset path to fill in the printed next-step command (not stored in the config)", help="Dataset path to fill in the printed next-step command (not stored in the config)",
@@ -931,7 +930,7 @@ def new_run(
# its stage1_model/stage2_model/conditioning content. # its stage1_model/stage2_model/conditioning content.
"config_version": gconfig.CONFIG_VERSION, "config_version": gconfig.CONFIG_VERSION,
"git_hash": gconfig.git_hash(), "git_hash": gconfig.git_hash(),
"created_at": datetime.now(timezone.utc).isoformat(timespec="seconds"), "created_at": datetime.now(UTC).isoformat(timespec="seconds"),
"created_by": "giant new-run", "created_by": "giant new-run",
} }
if comment: if comment:
@@ -958,7 +957,7 @@ app.add_typer(model_app, name="model")
@model_app.command("summary") @model_app.command("summary")
def model_summary( def model_summary(
config: Annotated[ config: Annotated[
Optional[Path], Path | None,
typer.Option("--config", "-c", help="TOML config file (default: built-in defaults)"), typer.Option("--config", "-c", help="TOML config file (default: built-in defaults)"),
] = None, ] = None,
pdg_vocab: Annotated[ pdg_vocab: Annotated[
@@ -1017,7 +1016,7 @@ def predict(
), ),
] = Coord.global_, ] = Coord.global_,
out: Annotated[ out: Annotated[
Optional[Path], Path | None,
typer.Option( typer.Option(
"--out", "--out",
"-o", "-o",
@@ -1050,11 +1049,11 @@ def predict(
), ),
] = Weights.raw, ] = Weights.raw,
device: Annotated[ device: Annotated[
Optional[str], str | None,
typer.Option("--device", "-d", help="cpu | cuda | mps (default: auto)"), typer.Option("--device", "-d", help="cpu | cuda | mps (default: auto)"),
] = None, ] = None,
comment: Annotated[ comment: Annotated[
Optional[str], str | None,
typer.Option( typer.Option(
"--comment", "--comment",
"-m", "-m",
@@ -1062,7 +1061,7 @@ def predict(
), ),
] = None, ] = None,
set_: Annotated[ set_: Annotated[
Optional[list[str]], list[str] | None,
typer.Option( typer.Option(
"--set", "--set",
help="Override a sampling-only model_config key on this checkpoint, " help="Override a sampling-only model_config key on this checkpoint, "
@@ -1092,7 +1091,7 @@ def predict(
from giant.sample import resolve_n_sec, sample_stage1, sample_stage2 from giant.sample import resolve_n_sec, sample_stage1, sample_stage2
batch_size_auto = False batch_size_auto = False
batch_size_value: Optional[int] = None batch_size_value: int | None = None
if batch_size.strip().lower() == "auto": if batch_size.strip().lower() == "auto":
batch_size_auto = True batch_size_auto = True
else: else:
@@ -1272,7 +1271,7 @@ def predict(
# "no snapping at inference"). "onehot"/"embedding": PDG # "no snapping at inference"). "onehot"/"embedding": PDG
# resolution IS the secondary's identity — see # resolution IS the secondary's identity — see
# decode_secondary_identity's docstring. # decode_secondary_identity's docstring.
sec_E, sec_dir_world, sec_mass, sec_charge, sec_pdg_code, _l1_dist = decode_secondary_identity( sec_E, sec_dir_world, _, _, sec_pdg_code, _l1_dist = decode_secondary_identity(
sec_decoder, sec_decoder,
sec_cont, sec_cont,
sec_type, sec_type,
@@ -1463,28 +1462,28 @@ def rollout(
] = Weights.raw, ] = Weights.raw,
batch_size: Annotated[int, typer.Option("--batch-size", "-b", help="Tracks stepped per model forward")] = 4096, batch_size: Annotated[int, typer.Option("--batch-size", "-b", help="Tracks stepped per model forward")] = 4096,
max_tracks_per_event: Annotated[ max_tracks_per_event: Annotated[
Optional[int], int | None,
typer.Option( typer.Option(
"--max-tracks-per-event", "--max-tracks-per-event",
help="Safety cap on tracks per shower (sub-cap secondaries deposit in place)", help="Safety cap on tracks per shower (sub-cap secondaries deposit in place)",
), ),
] = None, ] = None,
escape_threshold: Annotated[ escape_threshold: Annotated[
Optional[float], float | None,
typer.Option( typer.Option(
"--escape-threshold", "--escape-threshold",
help="Override the oracle's NN-distance escape threshold [mm]", help="Override the oracle's NN-distance escape threshold [mm]",
), ),
] = None, ] = None,
n_events: Annotated[Optional[int], typer.Option("--n-events", help="Cap number of seed events")] = None, n_events: Annotated[int | None, 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, device: Annotated[str | None, typer.Option("--device", "-d", help="cpu | cuda | mps (auto)")] = None,
out: Annotated[Optional[Path], typer.Option("--out", "-o", help="Output steps parquet")] = None, out: Annotated[Path | None, typer.Option("--out", "-o", help="Output steps parquet")] = None,
seed: Annotated[ seed: Annotated[
Optional[int], int | None,
typer.Option("--seed", help="Torch/numpy seed for reproducibility"), typer.Option("--seed", help="Torch/numpy seed for reproducibility"),
] = None, ] = None,
set_: Annotated[ set_: Annotated[
Optional[list[str]], list[str] | None,
typer.Option( typer.Option(
"--set", "--set",
help="Override a sampling-only model_config key on this checkpoint, " help="Override a sampling-only model_config key on this checkpoint, "
@@ -1505,7 +1504,8 @@ def rollout(
from giant.checkpoint_io import CheckpointCompatibilityError, load_for_inference from giant.checkpoint_io import CheckpointCompatibilityError, load_for_inference
from giant.data.loader import find_parquet_files from giant.data.loader import find_parquet_files
from giant.geometry import GeometryOracle from giant.geometry import GeometryOracle
from giant.rollout import L1DistCollector, RolloutSummary, rollout as run_rollout from giant.rollout import L1DistCollector, RolloutSummary
from giant.rollout import rollout as run_rollout
_t_setup_start = time.perf_counter() _t_setup_start = time.perf_counter()
@@ -1640,7 +1640,7 @@ def rollout(
"max_tracks_per_event": max_tracks_per_event, "max_tracks_per_event": max_tracks_per_event,
"escape_threshold": escape_threshold, "escape_threshold": escape_threshold,
"n_events": n_events, "n_events": n_events,
"n_seed_events": int(len(seeds["event_id"])), "n_seed_events": len(seeds["event_id"]),
"weights": weights.value, "weights": weights.value,
"batch_size": batch_size, "batch_size": batch_size,
"device": str(_device), "device": str(_device),
@@ -1700,7 +1700,7 @@ def analyze_prep(
), ),
], ],
label: Annotated[ label: Annotated[
Optional[list[str]], list[str] | None,
typer.Option( typer.Option(
"--label", "--label",
help="Series name for a rollout YAML, positionally matched to it — give none, " help="Series name for a rollout YAML, positionally matched to it — give none, "
@@ -1709,7 +1709,7 @@ def analyze_prep(
), ),
] = None, ] = None,
run_dir: Annotated[ run_dir: Annotated[
Optional[Path], Path | None,
typer.Option( typer.Option(
"--run-dir", "--run-dir",
"-o", "-o",
@@ -1797,7 +1797,7 @@ def analyze_render(
def analyze_metrics( def analyze_metrics(
run_dir: Annotated[Path, typer.Argument(help="Run directory containing metrics.csv (from `giant train`)")], run_dir: Annotated[Path, typer.Argument(help="Run directory containing metrics.csv (from `giant train`)")],
out_dir: Annotated[ out_dir: Annotated[
Optional[Path], Path | None,
typer.Option( typer.Option(
"--out", "--out",
"-o", "-o",
@@ -1823,7 +1823,7 @@ def analyze_submit(
], ],
accounting_group: Annotated[str, typer.Option("--accounting-group")], accounting_group: Annotated[str, typer.Option("--accounting-group")],
label: Annotated[ label: Annotated[
Optional[list[str]], list[str] | None,
typer.Option( typer.Option(
"--label", "--label",
help="Series name for a rollout YAML, positionally matched to it — give none, " help="Series name for a rollout YAML, positionally matched to it — give none, "
@@ -1832,7 +1832,7 @@ def analyze_submit(
), ),
] = None, ] = None,
run_dir: Annotated[ run_dir: Annotated[
Optional[Path], Path | None,
typer.Option( typer.Option(
"--run-dir", "--run-dir",
"-o", "-o",
+23 -23
View File
@@ -9,7 +9,7 @@ import subprocess
import sys import sys
import tomllib import tomllib
from dataclasses import dataclass, field from dataclasses import dataclass, field
from datetime import datetime, timezone from datetime import UTC, datetime
from enum import Enum from enum import Enum
from pathlib import Path from pathlib import Path
from typing import TYPE_CHECKING from typing import TYPE_CHECKING
@@ -80,7 +80,7 @@ class ConditioningAxisConfig:
n_layers: int = 1 n_layers: int = 1
@classmethod @classmethod
def from_dict(cls, d: dict | None) -> "ConditioningAxisConfig": def from_dict(cls, d: dict | None) -> ConditioningAxisConfig:
d = d or {} d = d or {}
return cls( return cls(
type=d.get("type", "physical"), type=d.get("type", "physical"),
@@ -106,7 +106,7 @@ class ConditioningConfig:
material: ConditioningAxisConfig = field(default_factory=ConditioningAxisConfig) material: ConditioningAxisConfig = field(default_factory=ConditioningAxisConfig)
@classmethod @classmethod
def from_dict(cls, d: dict | None) -> "ConditioningConfig": def from_dict(cls, d: dict | None) -> ConditioningConfig:
d = d or {} d = d or {}
return cls( return cls(
out_dim=d.get("out_dim", 128), out_dim=d.get("out_dim", 128),
@@ -130,7 +130,7 @@ class FlowConfig:
time_dim: int = 64 time_dim: int = 64
@classmethod @classmethod
def from_dict(cls, d: dict | None) -> "FlowConfig": def from_dict(cls, d: dict | None) -> FlowConfig:
d = d or {} d = d or {}
return cls(time_dim=d.get("time_dim", 64)) return cls(time_dim=d.get("time_dim", 64))
@@ -144,7 +144,7 @@ class DdpmConfig:
n_steps: int = 1000 n_steps: int = 1000
@classmethod @classmethod
def from_dict(cls, d: dict | None) -> "DdpmConfig": def from_dict(cls, d: dict | None) -> DdpmConfig:
d = d or {} d = d or {}
return cls(time_dim=d.get("time_dim", 64), n_steps=d.get("n_steps", 1000)) return cls(time_dim=d.get("time_dim", 64), n_steps=d.get("n_steps", 1000))
@@ -166,7 +166,7 @@ class Stage1WganConfig:
critic_n_res_blocks: int = 0 critic_n_res_blocks: int = 0
@classmethod @classmethod
def from_dict(cls, d: dict | None) -> "Stage1WganConfig": def from_dict(cls, d: dict | None) -> Stage1WganConfig:
d = d or {} d = d or {}
return cls( return cls(
noise_dim=d.get("noise_dim", 64), noise_dim=d.get("noise_dim", 64),
@@ -198,7 +198,7 @@ class Stage2WganConfig(Stage1WganConfig):
gumbel_tau_end: float = 0.1 gumbel_tau_end: float = 0.1
@classmethod @classmethod
def from_dict(cls, d: dict | None) -> "Stage2WganConfig": def from_dict(cls, d: dict | None) -> Stage2WganConfig:
d = d or {} d = d or {}
return cls( return cls(
noise_dim=d.get("noise_dim", 64), noise_dim=d.get("noise_dim", 64),
@@ -293,7 +293,7 @@ class RouterConfig:
extra: dict = field(default_factory=dict) extra: dict = field(default_factory=dict)
@classmethod @classmethod
def from_dict(cls, d: dict | None) -> "RouterConfig": def from_dict(cls, d: dict | None) -> RouterConfig:
d = d or {} d = d or {}
return cls( return cls(
enabled=d.get("enabled", False), enabled=d.get("enabled", False),
@@ -360,7 +360,7 @@ class TrunkConfig:
block_conditioning: str = "add" block_conditioning: str = "add"
@classmethod @classmethod
def from_dict(cls, d: dict | None) -> "TrunkConfig": def from_dict(cls, d: dict | None) -> TrunkConfig:
d = d or {} d = d or {}
return cls(type=d.get("type", "resmlp"), block_conditioning=d.get("block_conditioning", "add")) return cls(type=d.get("type", "resmlp"), block_conditioning=d.get("block_conditioning", "add"))
@@ -377,7 +377,7 @@ class Stage2RouterConfig(RouterConfig):
tie_to_stage1: bool = False tie_to_stage1: bool = False
@classmethod @classmethod
def from_dict(cls, d: dict | None) -> "Stage2RouterConfig": def from_dict(cls, d: dict | None) -> Stage2RouterConfig:
d = d or {} d = d or {}
known = _ROUTER_KNOWN_KEYS | {"tie_to_stage1"} known = _ROUTER_KNOWN_KEYS | {"tie_to_stage1"}
return cls( return cls(
@@ -447,7 +447,7 @@ class NSecConfig:
sampling: str = "greedy" sampling: str = "greedy"
@classmethod @classmethod
def from_dict(cls, d: dict | None) -> "NSecConfig": def from_dict(cls, d: dict | None) -> NSecConfig:
d = d or {} d = d or {}
return cls( return cls(
mode=d.get("mode", "head"), mode=d.get("mode", "head"),
@@ -498,7 +498,7 @@ class ParticleTypeConfig:
class_weighting: str = "none" class_weighting: str = "none"
@classmethod @classmethod
def from_dict(cls, d: dict | None) -> "ParticleTypeConfig": def from_dict(cls, d: dict | None) -> ParticleTypeConfig:
d = d or {} d = d or {}
return cls( return cls(
target=d.get("target", "onehot"), target=d.get("target", "onehot"),
@@ -537,7 +537,7 @@ class AutoregressiveConfig:
attn_n_layers: int = 2 attn_n_layers: int = 2
@classmethod @classmethod
def from_dict(cls, d: dict | None) -> "AutoregressiveConfig": def from_dict(cls, d: dict | None) -> AutoregressiveConfig:
d = d or {} d = d or {}
return cls( return cls(
order=d.get("order", "energy_desc"), order=d.get("order", "energy_desc"),
@@ -576,7 +576,7 @@ class HeadConfig:
depth: int = 2 # matches build_mlp_head's depth depth: int = 2 # matches build_mlp_head's depth
@classmethod @classmethod
def from_dict(cls, d: dict | None) -> "HeadConfig": def from_dict(cls, d: dict | None) -> HeadConfig:
d = d or {} d = d or {}
return cls(hidden_ratio=d.get("hidden_ratio", 0.5), depth=d.get("depth", 2)) return cls(hidden_ratio=d.get("hidden_ratio", 0.5), depth=d.get("depth", 2))
@@ -593,7 +593,7 @@ class Stage1HeadsConfig:
n_sec: HeadConfig = field(default_factory=HeadConfig) n_sec: HeadConfig = field(default_factory=HeadConfig)
@classmethod @classmethod
def from_dict(cls, d: dict | None) -> "Stage1HeadsConfig": def from_dict(cls, d: dict | None) -> Stage1HeadsConfig:
d = d or {} d = d or {}
return cls(n_sec=HeadConfig.from_dict(d.get("n_sec"))) return cls(n_sec=HeadConfig.from_dict(d.get("n_sec")))
@@ -611,7 +611,7 @@ class Stage2HeadsConfig:
type: HeadConfig = field(default_factory=HeadConfig) type: HeadConfig = field(default_factory=HeadConfig)
@classmethod @classmethod
def from_dict(cls, d: dict | None) -> "Stage2HeadsConfig": def from_dict(cls, d: dict | None) -> Stage2HeadsConfig:
d = d or {} d = d or {}
return cls(n_sec=HeadConfig.from_dict(d.get("n_sec")), type=HeadConfig.from_dict(d.get("type"))) return cls(n_sec=HeadConfig.from_dict(d.get("n_sec")), type=HeadConfig.from_dict(d.get("type")))
@@ -659,7 +659,7 @@ class Stage1ModelConfig:
heads: Stage1HeadsConfig = field(default_factory=Stage1HeadsConfig) heads: Stage1HeadsConfig = field(default_factory=Stage1HeadsConfig)
@classmethod @classmethod
def from_dict(cls, d: dict | None) -> "Stage1ModelConfig": def from_dict(cls, d: dict | None) -> Stage1ModelConfig:
d = d or {} d = d or {}
return cls( return cls(
active=d.get("active", True), active=d.get("active", True),
@@ -745,7 +745,7 @@ class Stage2ModelConfig:
heads: Stage2HeadsConfig = field(default_factory=Stage2HeadsConfig) heads: Stage2HeadsConfig = field(default_factory=Stage2HeadsConfig)
@classmethod @classmethod
def from_dict(cls, d: dict | None) -> "Stage2ModelConfig": def from_dict(cls, d: dict | None) -> Stage2ModelConfig:
d = d or {} d = d or {}
return cls( return cls(
active=d.get("active", True), active=d.get("active", True),
@@ -839,7 +839,7 @@ class TrainConfig:
precision: str = "fp32" precision: str = "fp32"
@classmethod @classmethod
def from_dict(cls, d: dict | None) -> "TrainConfig": def from_dict(cls, d: dict | None) -> TrainConfig:
d = d or {} d = d or {}
return cls( return cls(
epochs=d.get("epochs", 100), epochs=d.get("epochs", 100),
@@ -895,7 +895,7 @@ class GiantConfig:
train: TrainConfig = field(default_factory=TrainConfig) train: TrainConfig = field(default_factory=TrainConfig)
@classmethod @classmethod
def from_dict(cls, d: dict | None) -> "GiantConfig": def from_dict(cls, d: dict | None) -> GiantConfig:
d = d or {} d = d or {}
return cls( return cls(
conditioning=ConditioningConfig.from_dict(d.get("conditioning")), conditioning=ConditioningConfig.from_dict(d.get("conditioning")),
@@ -938,7 +938,7 @@ def leaf_paths(node: dict, prefix: str = "") -> list[str]:
def git_hash() -> str: def git_hash() -> str:
try: try:
return subprocess.check_output(["git", "rev-parse", "HEAD"], stderr=subprocess.DEVNULL).decode().strip() return subprocess.check_output(["git", "rev-parse", "HEAD"], stderr=subprocess.DEVNULL).decode().strip()
except Exception: except Exception: # noqa: BLE001 - any failure (no git, no repo, ...) degrades to "unknown"
return "unknown" return "unknown"
@@ -1852,7 +1852,7 @@ def default_out_dir_name(cfg: dict, now: datetime | None = None) -> str:
name unboundedly. This name doubles as the run's W&B id (see name unboundedly. This name doubles as the run's W&B id (see
giant.training), which is the reason a timestamp is always included. giant.training), which is the reason a timestamp is always included.
""" """
now = now or datetime.now() now = now or datetime.now() # noqa: DTZ005 - human-readable local wall-clock time for run/W&B naming, not stored
tokens = [] tokens = []
overflow = [] overflow = []
for label, candidate in _OUT_DIR_NAME_CANDIDATES: for label, candidate in _OUT_DIR_NAME_CANDIDATES:
@@ -1951,7 +1951,7 @@ def build_run_meta(
"config_version": CONFIG_VERSION, "config_version": CONFIG_VERSION,
"git_hash": git_hash(), "git_hash": git_hash(),
"seed": seed, "seed": seed,
"timestamp_utc": datetime.now(timezone.utc).isoformat(timespec="seconds"), "timestamp_utc": datetime.now(UTC).isoformat(timespec="seconds"),
"python_version": sys.version.split()[0], "python_version": sys.version.split()[0],
"torch_version": torch.__version__, "torch_version": torch.__version__,
"command": " ".join(sys.argv), "command": " ".join(sys.argv),
+2 -1
View File
@@ -1,6 +1,7 @@
from collections.abc import Iterator, Mapping
from dataclasses import dataclass, field from dataclasses import dataclass, field
from pathlib import Path from pathlib import Path
from typing import TYPE_CHECKING, Any, Iterator, Mapping from typing import TYPE_CHECKING, Any
import numpy as np import numpy as np
import polars as pl import polars as pl
+4 -4
View File
@@ -172,7 +172,7 @@ class NormalizerEntry:
} }
@classmethod @classmethod
def from_json(cls, d: dict) -> "NormalizerEntry": def from_json(cls, d: dict) -> NormalizerEntry:
return cls( return cls(
cond_norm=Normalizer.from_dict(d["cond_norm"]), cond_norm=Normalizer.from_dict(d["cond_norm"]),
tgt_norm=Normalizer.from_dict(d["tgt_norm"]), tgt_norm=Normalizer.from_dict(d["tgt_norm"]),
@@ -194,7 +194,7 @@ class SetupCache:
"""Keyed by `topn_key(axis, n_classes)`.""" """Keyed by `topn_key(axis, n_classes)`."""
@classmethod @classmethod
def empty(cls, files: list[Path]) -> "SetupCache": def empty(cls, files: list[Path]) -> SetupCache:
return cls(fingerprint=fingerprint_files(files)) return cls(fingerprint=fingerprint_files(files))
def to_json(self) -> dict: def to_json(self) -> dict:
@@ -222,7 +222,7 @@ class SetupCache:
return d return d
@classmethod @classmethod
def from_json(cls, d: dict) -> "SetupCache": def from_json(cls, d: dict) -> SetupCache:
vocab = None vocab = None
if "vocab" in d: if "vocab" in d:
pdg_map = {int(k): v for k, v in d["vocab"]["pdg_map"].items()} pdg_map = {int(k): v for k, v in d["vocab"]["pdg_map"].items()}
@@ -247,7 +247,7 @@ class SetupCache:
topn_maps=topn_maps, topn_maps=topn_maps,
) )
def merge(self, other: "SetupCache") -> "SetupCache": def merge(self, other: SetupCache) -> SetupCache:
"""Union of both caches; `other`'s populated fields win on a shared key. """Union of both caches; `other`'s populated fields win on a shared key.
Used by `save` to combine freshly-computed sections with whatever a Used by `save` to combine freshly-computed sections with whatever a
+6 -5
View File
@@ -23,9 +23,10 @@ install stays lean.
from __future__ import annotations from __future__ import annotations
from collections.abc import Iterable
from dataclasses import dataclass, field from dataclasses import dataclass, field
from pathlib import Path from pathlib import Path
from typing import Any, Iterable from typing import Any
import numpy as np import numpy as np
import polars as pl import polars as pl
@@ -155,7 +156,7 @@ class GeometryOracle:
) )
@classmethod @classmethod
def load(cls, path: str | Path) -> "GeometryOracle": def load(cls, path: str | Path) -> GeometryOracle:
_require_sklearn() _require_sklearn()
import joblib import joblib
@@ -353,7 +354,7 @@ def _fit_slab_lookup(
radius_max=radius_max, radius_max=radius_max,
) )
info = { info = {
"n_segments": int(len(materials)), "n_segments": len(materials),
"z_range": (z_min, z_max), "z_range": (z_min, z_max),
"median_z_spacing": median_spacing, "median_z_spacing": median_spacing,
"radius_max": radius_max, "radius_max": radius_max,
@@ -412,7 +413,7 @@ def build_geometry_oracle(
"method": "slab", "method": "slab",
"depth_axis": depth_axis, "depth_axis": depth_axis,
"n_bins": n_bins, "n_bins": n_bins,
"n_reference_points": int(len(pos)), "n_reference_points": len(pos),
"escape_factor": escape_factor, "escape_factor": escape_factor,
"n_files": len(files), "n_files": len(files),
**info, **info,
@@ -458,7 +459,7 @@ def build_geometry_oracle(
metadata={ metadata={
"method": method, "method": method,
"k": k, "k": k,
"n_reference_points": int(len(X)), "n_reference_points": len(X),
"median_nn_dist": median_nn, "median_nn_dist": median_nn,
"escape_factor": escape_factor, "escape_factor": escape_factor,
"n_files": len(files), "n_files": len(files),
+1 -1
View File
@@ -1,7 +1,7 @@
"""Factories: `build_models`/`build_critics` assemble the top-level stage """Factories: `build_models`/`build_critics` assemble the top-level stage
models from a config dict (issues.md Issue 8).""" models from a config dict (issues.md Issue 8)."""
import torch.nn as nn from torch import nn
from giant.config import ConditioningConfig, Stage1ModelConfig, Stage2ModelConfig from giant.config import ConditioningConfig, Stage1ModelConfig, Stage2ModelConfig
from giant.constants import X_DIM from giant.constants import X_DIM
+1 -1
View File
@@ -2,8 +2,8 @@
identity (issues.md Issue 8).""" identity (issues.md Issue 8)."""
import torch import torch
import torch.nn as nn
import torch.nn.functional as F import torch.nn.functional as F
from torch import nn
from giant.cond_layout import CondLayout from giant.cond_layout import CondLayout
from giant.config import ConditioningAxisConfig from giant.config import ConditioningAxisConfig
+2 -2
View File
@@ -7,7 +7,7 @@ import inspect
from typing import cast from typing import cast
import torch import torch
import torch.nn as nn from torch import nn
class HistoryEncoder(nn.Module): class HistoryEncoder(nn.Module):
@@ -197,7 +197,7 @@ class AttentionHistory(HistoryEncoder):
return self.in_proj(x) return self.in_proj(x)
def forward(self, feat: torch.Tensor, has_prev: torch.Tensor) -> torch.Tensor: def forward(self, feat: torch.Tensor, has_prev: torch.Tensor) -> torch.Tensor:
B, K, _ = feat.shape _, K, _ = feat.shape
x = self._embed(feat, has_prev) x = self._embed(feat, has_prev)
mask = nn.Transformer.generate_square_subsequent_mask(K, device=feat.device) mask = nn.Transformer.generate_square_subsequent_mask(K, device=feat.device)
for block in self.blocks: for block in self.blocks:
+1 -1
View File
@@ -4,7 +4,7 @@ no dependency on any other `giant.model` submodule (issues.md Issue 8)."""
import math import math
import torch import torch
import torch.nn as nn from torch import nn
class SinusoidalEmbedding(nn.Module): class SinusoidalEmbedding(nn.Module):
+1 -1
View File
@@ -2,7 +2,7 @@
`CriticModel` composed from encoders/trunks/history (issues.md Issue 8).""" `CriticModel` composed from encoders/trunks/history (issues.md Issue 8)."""
import torch import torch
import torch.nn as nn from torch import nn
from giant.config import ConditioningAxisConfig, HeadConfig, ParticleTypeConfig from giant.config import ConditioningAxisConfig, HeadConfig, ParticleTypeConfig
from giant.constants import CONT_SLOT_DIM, K_MAX, PARTICLE_PHYS_DIM, SEC_DIM, SEC_SLOT_DIM, X_DIM from giant.constants import CONT_SLOT_DIM, K_MAX, PARTICLE_PHYS_DIM, SEC_DIM, SEC_SLOT_DIM, X_DIM
+5 -5
View File
@@ -79,9 +79,13 @@ from giant.model.trunks import (
) )
__all__ = [ __all__ = [
"BLOCK_REGISTRY",
"HISTORY_REGISTRY",
"OBJECTIVE_REGISTRY",
"ROUTER_REGISTRY",
"TRUNK_REGISTRY",
"AdaLNResBlock", "AdaLNResBlock",
"AttentionHistory", "AttentionHistory",
"BLOCK_REGISTRY",
"ComposedRouter", "ComposedRouter",
"ConditionEncoder", "ConditionEncoder",
"ContextAdapter", "ContextAdapter",
@@ -91,17 +95,14 @@ __all__ = [
"ExpertTrunk", "ExpertTrunk",
"FilmResBlock", "FilmResBlock",
"FlowObjective", "FlowObjective",
"HISTORY_REGISTRY",
"HistoryEncoder", "HistoryEncoder",
"LinearTrunk", "LinearTrunk",
"MarkovHistory", "MarkovHistory",
"NoHistory", "NoHistory",
"NoneRouter", "NoneRouter",
"OBJECTIVE_REGISTRY",
"Objective", "Objective",
"PdgRouter", "PdgRouter",
"ProcessRouter", "ProcessRouter",
"ROUTER_REGISTRY",
"ResBlock", "ResBlock",
"RoutedTrunk", "RoutedTrunk",
"Router", "Router",
@@ -110,7 +111,6 @@ __all__ = [
"Stage2Autoregressive", "Stage2Autoregressive",
"Stage2OneShot", "Stage2OneShot",
"StageModel", "StageModel",
"TRUNK_REGISTRY",
"Trunk", "Trunk",
"WganObjective", "WganObjective",
"_CausalAttnBlock", "_CausalAttnBlock",
+2 -2
View File
@@ -8,8 +8,8 @@ import re
from collections.abc import Sequence from collections.abc import Sequence
import torch import torch
import torch.nn as nn
import torch.nn.functional as F import torch.nn.functional as F
from torch import nn
from giant.cond_layout import CondLayout from giant.cond_layout import CondLayout
from giant.constants import COND_DIM from giant.constants import COND_DIM
@@ -398,7 +398,7 @@ def _build_router_from_cfg(
"""Resolve one stage's `router` config into a `Router`, single-axis or """Resolve one stage's `router` config into a `Router`, single-axis or
composed. `gumbel` is set as a post-construction attribute (shared by composed. `gumbel` is set as a post-construction attribute (shared by
every router type, not a per-type constructor kwarg).""" every router type, not a per-type constructor kwarg)."""
shared_vocab = dict(pdg_vocab=pdg_vocab, mat_vocab=mat_vocab) shared_vocab = {"pdg_vocab": pdg_vocab, "mat_vocab": mat_vocab}
if router_cfg["type"] == "composed": if router_cfg["type"] == "composed":
axes = _parse_composed_axes(router_cfg) axes = _parse_composed_axes(router_cfg)
_check_router_conditioning_compat([a["type"] for a in axes], particle_conditioning) _check_router_conditioning_compat([a["type"] for a in axes], particle_conditioning)
+1 -1
View File
@@ -14,7 +14,7 @@ class CosineSchedule:
betas = np.clip(1.0 - alpha_bars[1:] / alpha_bars[:-1], 0.0, 0.999).astype(np.float32) betas = np.clip(1.0 - alpha_bars[1:] / alpha_bars[:-1], 0.0, 0.999).astype(np.float32)
self.betas = torch.from_numpy(betas) self.betas = torch.from_numpy(betas)
self.alphas = torch.from_numpy((1.0 - betas)) self.alphas = torch.from_numpy(1.0 - betas)
self.alpha_bars = torch.from_numpy(alpha_bars[1:]) self.alpha_bars = torch.from_numpy(alpha_bars[1:])
def to(self, device: torch.device) -> "CosineSchedule": def to(self, device: torch.device) -> "CosineSchedule":
+2 -2
View File
@@ -35,7 +35,7 @@ finding those two tests independently converge on.
import copy import copy
from dataclasses import dataclass, field from dataclasses import dataclass, field
import torch.nn as nn from torch import nn
from giant.config import INFERENCE_OVERRIDES, _get_path, _set_path, leaf_paths from giant.config import INFERENCE_OVERRIDES, _get_path, _set_path, leaf_paths
from giant.model.builders import build_critics, build_models from giant.model.builders import build_critics, build_models
@@ -224,7 +224,7 @@ def summarize_model(cfg: dict, pdg_vocab: int, mat_vocab: int) -> ModelSummary:
_set_path(probe_cfg, path, candidate) _set_path(probe_cfg, path, candidate)
try: try:
changed = _fingerprint(_built_modules(probe_cfg, pdg_vocab, mat_vocab)) != baseline_fp changed = _fingerprint(_built_modules(probe_cfg, pdg_vocab, mat_vocab)) != baseline_fp
except Exception: except Exception: # noqa: BLE001 - a perturbation that fails to even build counts as "consumed"
changed = True changed = True
if changed: if changed:
break break
+1 -1
View File
@@ -10,7 +10,7 @@ free — no separate "routed transformer trunk" class needed.
""" """
import torch import torch
import torch.nn as nn from torch import nn
from giant.model.layers import build_block from giant.model.layers import build_block
from giant.model.routers import Router from giant.model.routers import Router
+1 -1
View File
@@ -1,4 +1,4 @@
from typing import Callable from collections.abc import Callable
import torch import torch
+3 -3
View File
@@ -19,7 +19,7 @@ free-running history representation stays unsnapped — see
from __future__ import annotations from __future__ import annotations
from functools import lru_cache from functools import cache
from typing import TYPE_CHECKING from typing import TYPE_CHECKING
import numpy as np import numpy as np
@@ -42,7 +42,7 @@ _CHARGE_WEIGHT = 50.0
_LOG_EPS = 1e-8 _LOG_EPS = 1e-8
@lru_cache(maxsize=None) @cache
def particle_mass_charge(pdg: int) -> tuple[float, float]: def particle_mass_charge(pdg: int) -> tuple[float, float]:
"""Return (mass_MeV, charge_e) for a raw PDG code. """Return (mass_MeV, charge_e) for a raw PDG code.
@@ -135,7 +135,7 @@ def invert_dense_map(m: dict[int, int]) -> dict[int, int]:
def decode_topn_class( def decode_topn_class(
class_idx: np.ndarray, class_idx: np.ndarray,
topn_map: "TopNMap", topn_map: TopNMap,
n_classes: int, n_classes: int,
other_policy: str = "sample", other_policy: str = "sample",
rng: np.random.Generator | None = None, rng: np.random.Generator | None = None,
+4 -4
View File
@@ -13,6 +13,7 @@ from giant.constants import (
X_DIM, X_DIM,
) )
from giant.data import setup_cache from giant.data import setup_cache
from giant.data.dataset import StreamingStepsDataset, make_event_split
from giant.data.loader import ( from giant.data.loader import (
TopNMap, TopNMap,
_topn_plus_other_map, _topn_plus_other_map,
@@ -23,13 +24,12 @@ from giant.data.loader import (
from giant.data.scan import MetadataScan, ScanRequest, scan_metadata from giant.data.scan import MetadataScan, ScanRequest, scan_metadata
from giant.data.transforms import ( from giant.data.transforms import (
Normalizer, Normalizer,
build_features,
_WelfordAccumulator,
_ReservoirSampler, _ReservoirSampler,
_WelfordAccumulator,
build_features,
sorted_membership, sorted_membership,
) )
from giant.data.dataset import make_event_split, StreamingStepsDataset from giant.model.network import build_critics, build_models, resolve_type_n_classes
from giant.model.network import build_models, build_critics, resolve_type_n_classes
from giant.training import train as run_training from giant.training import train as run_training
+35 -34
View File
@@ -17,7 +17,8 @@ treated as detector leakage and not deposited.
from __future__ import annotations from __future__ import annotations
from collections import Counter from collections import Counter
from typing import TYPE_CHECKING, Callable, TypedDict from collections.abc import Callable
from typing import TYPE_CHECKING, TypedDict
import numpy as np import numpy as np
import torch import torch
@@ -117,7 +118,7 @@ def decode_secondary_identity(
pre_dir: np.ndarray, pre_dir: np.ndarray,
sec_phys_norm: Normalizer, sec_phys_norm: Normalizer,
pdg_map: dict[int, int], pdg_map: dict[int, int],
sec_type_topn_map: "TopNMap | None", sec_type_topn_map: TopNMap | None,
other_policy: str, other_policy: str,
rng: np.random.Generator | None, rng: np.random.Generator | None,
) -> tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray, np.ndarray, np.ndarray | None]: ) -> tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray, np.ndarray, np.ndarray | None]:
@@ -390,34 +391,34 @@ def _terminal_rows(tr: dict[str, np.ndarray], sel: np.ndarray, reason: str, edep
pos = tr["pre_pos"][sel] pos = tr["pre_pos"][sel]
dir_ = tr["pre_dir"][sel] dir_ = tr["pre_dir"][sel]
n = int(sel.sum()) n = int(sel.sum())
return dict( return {
event_id=tr["event_id"][sel], "event_id": tr["event_id"][sel],
track_id=tr["track_id"][sel], "track_id": tr["track_id"][sel],
parent_id=tr["parent_id"][sel], "parent_id": tr["parent_id"][sel],
generation=tr["generation"][sel], "generation": tr["generation"][sel],
step_no=tr["step_in_track"][sel], "step_no": tr["step_in_track"][sel],
pdg=tr["pdg"][sel], "pdg": tr["pdg"][sel],
pre_x=pos[:, 0], "pre_x": pos[:, 0],
pre_y=pos[:, 1], "pre_y": pos[:, 1],
pre_z=pos[:, 2], "pre_z": pos[:, 2],
pre_E=tr["pre_E"][sel], "pre_E": tr["pre_E"][sel],
pre_dx=dir_[:, 0], "pre_dx": dir_[:, 0],
pre_dy=dir_[:, 1], "pre_dy": dir_[:, 1],
pre_dz=dir_[:, 2], "pre_dz": dir_[:, 2],
post_x=pos[:, 0], "post_x": pos[:, 0],
post_y=pos[:, 1], "post_y": pos[:, 1],
post_z=pos[:, 2], "post_z": pos[:, 2],
post_E=np.zeros(n), "post_E": np.zeros(n),
post_dx=dir_[:, 0], "post_dx": dir_[:, 0],
post_dy=dir_[:, 1], "post_dy": dir_[:, 1],
post_dz=dir_[:, 2], "post_dz": dir_[:, 2],
edep=np.asarray(edep, dtype=np.float64).reshape(n), "edep": np.asarray(edep, dtype=np.float64).reshape(n),
step_length=np.zeros(n), "step_length": np.zeros(n),
material=tr.get("_material", np.full(len(sel), "", dtype=object))[sel], "material": tr.get("_material", np.full(len(sel), "", dtype=object))[sel],
layer_id=tr.get("_layer_id", np.zeros(len(sel), dtype=np.int64))[sel], "layer_id": tr.get("_layer_id", np.zeros(len(sel), dtype=np.int64))[sel],
n_sec_pred=np.zeros(n, dtype=np.int64), "n_sec_pred": np.zeros(n, dtype=np.int64),
termination_reason=np.full(n, reason, dtype=object), "termination_reason": np.full(n, reason, dtype=object),
) }
@torch.no_grad() @torch.no_grad()
@@ -442,14 +443,14 @@ def rollout(
on_chunk: Callable[[dict[str, np.ndarray]], None] | None = None, on_chunk: Callable[[dict[str, np.ndarray]], None] | None = None,
particle_conditioning: str = "embedding", particle_conditioning: str = "embedding",
material_conditioning: str = "embedding", material_conditioning: str = "embedding",
pdg_topn_map: "TopNMap | None" = None, pdg_topn_map: TopNMap | None = None,
mat_topn_map: "TopNMap | None" = None, mat_topn_map: TopNMap | None = None,
sec_type_topn_map: "TopNMap | None" = None, sec_type_topn_map: TopNMap | None = None,
other_policy: str = "sample", other_policy: str = "sample",
seed: int | None = None, seed: int | None = None,
stage1_ddpm_steps: int = 1000, stage1_ddpm_steps: int = 1000,
stage2_ddpm_steps: int = 1000, stage2_ddpm_steps: int = 1000,
l1_dist_collector: "L1DistCollector | None" = None, l1_dist_collector: L1DistCollector | None = None,
) -> dict[str, np.ndarray] | RolloutSummary: ) -> dict[str, np.ndarray] | RolloutSummary:
"""Run showers to completion. """Run showers to completion.
+2 -2
View File
@@ -82,7 +82,7 @@ def _max_index(parent: Path, pattern: re.Pattern) -> int:
def _git_user_name() -> str | None: def _git_user_name() -> str | None:
try: try:
out = subprocess.run(["git", "config", "user.name"], capture_output=True, text=True, timeout=2) out = subprocess.run(["git", "config", "user.name"], capture_output=True, text=True, timeout=2, check=False)
except (OSError, subprocess.SubprocessError): except (OSError, subprocess.SubprocessError):
# OSError (e.g. git not on PATH) and subprocess.SubprocessError # OSError (e.g. git not on PATH) and subprocess.SubprocessError
# (e.g. TimeoutExpired) are unrelated hierarchies — TimeoutExpired # (e.g. TimeoutExpired) are unrelated hierarchies — TimeoutExpired
@@ -604,7 +604,7 @@ def _run_bump(
if not root_path.is_dir(): if not root_path.is_dir():
raise SystemExit(f"error: {root_path} is not a directory") raise SystemExit(f"error: {root_path} is not a directory")
date = date or dt.date.today().isoformat() date = date or dt.date.today().isoformat() # noqa: DTZ011 - local calendar date for the dataset-version log, not stored
by = by if by is not None else _git_user_name() by = by if by is not None else _git_user_name()
if gen is None: if gen is None:
new_dirs, log_line = plan_bump_gen(root_path, kind, reason, by, date, to) new_dirs, log_line = plan_bump_gen(root_path, kind, reason, by, date, to)
+3 -3
View File
@@ -31,9 +31,9 @@ import torch
sys.path.insert(0, str(Path(__file__).resolve().parent.parent)) sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
from giant.constants import COND_DIM, SEC_SLOT_DIM, X_DIM # noqa: E402 from giant.constants import COND_DIM, SEC_SLOT_DIM, X_DIM
from giant.model import network as net # noqa: E402 from giant.model import network as net
from tests.legacy import network_v02_snapshot as legacy # noqa: E402 from tests.legacy import network_v02_snapshot as legacy
def _random_batch(model_config: dict, batch: int, seed: int): def _random_batch(model_config: dict, batch: int, seed: int):
+2 -2
View File
@@ -28,8 +28,8 @@ import subprocess
import sys import sys
import uuid import uuid
import zlib import zlib
from dataclasses import dataclass
from concurrent.futures import ThreadPoolExecutor, as_completed from concurrent.futures import ThreadPoolExecutor, as_completed
from dataclasses import dataclass
from pathlib import Path from pathlib import Path
# Must match giant/tools/bump_dataset_version.py's GEN_RE. # Must match giant/tools/bump_dataset_version.py's GEN_RE.
@@ -147,7 +147,7 @@ def run_job(
cmd = build_cmd(executable, job, events_per_file, energy_gev) cmd = build_cmd(executable, job, events_per_file, energy_gev)
env = dict(os.environ, MINICALOSIM_SEED=str(job_seed(kind, gen, job, energy_gev))) env = dict(os.environ, MINICALOSIM_SEED=str(job_seed(kind, gen, job, energy_gev)))
result = subprocess.run(cmd, cwd=workdir, capture_output=True, text=True, env=env) result = subprocess.run(cmd, cwd=workdir, capture_output=True, text=True, env=env, check=False)
if result.returncode != 0: if result.returncode != 0:
return JobResult( return JobResult(
+22 -23
View File
@@ -10,10 +10,9 @@ from __future__ import annotations
import os import os
from enum import Enum from enum import Enum
from pathlib import Path from pathlib import Path
from typing import Optional from typing import Annotated
import typer import typer
from typing_extensions import Annotated
from giant.config import Conditioning from giant.config import Conditioning
@@ -73,7 +72,7 @@ class PoolType(str, Enum):
def convert( def convert(
root_files: Annotated[list[Path], typer.Argument(help="Input ROOT file(s)")], root_files: Annotated[list[Path], typer.Argument(help="Input ROOT file(s)")],
output: Annotated[ output: Annotated[
Optional[Path], Path | None,
typer.Option( typer.Option(
"--output", "--output",
"-o", "-o",
@@ -108,7 +107,7 @@ def convert(
), ),
] = _DATASET_ROOT_DEFAULT, ] = _DATASET_ROOT_DEFAULT,
schema: Annotated[ schema: Annotated[
Optional[str], str | None,
typer.Option( typer.Option(
"--schema", "--schema",
help="Schema tag to write parquets under, e.g. schema2 (only used " help="Schema tag to write parquets under, e.g. schema2 (only used "
@@ -192,13 +191,13 @@ def migrate(
def bump_gen( def bump_gen(
reason: Annotated[str, typer.Option("--reason", help="Why this gen exists")], reason: Annotated[str, typer.Option("--reason", help="Why this gen exists")],
kind: Annotated[str, typer.Option("--kind", help="steps | hits | ... (default: steps)")] = "steps", kind: Annotated[str, typer.Option("--kind", help="steps | hits | ... (default: steps)")] = "steps",
by: Annotated[Optional[str], typer.Option("--by", help="Attribution (default: git user.name)")] = None, by: Annotated[str | None, typer.Option("--by", help="Attribution (default: git user.name)")] = None,
date: Annotated[ date: Annotated[
Optional[str], str | None,
typer.Option("--date", help="Override date (default: today, ISO)"), typer.Option("--date", help="Override date (default: today, ISO)"),
] = None, ] = None,
to: Annotated[ to: Annotated[
Optional[str], str | None,
typer.Option( typer.Option(
"--to", "--to",
metavar="genN", metavar="genN",
@@ -227,13 +226,13 @@ def bump_schema(
gen: Annotated[str, typer.Option("--gen", help="Existing gen tag, e.g. gen1")], gen: Annotated[str, typer.Option("--gen", help="Existing gen tag, e.g. gen1")],
reason: Annotated[str, typer.Option("--reason", help="Why this schema exists")], reason: Annotated[str, typer.Option("--reason", help="Why this schema exists")],
kind: Annotated[str, typer.Option("--kind", help="steps | hits | ... (default: steps)")] = "steps", kind: Annotated[str, typer.Option("--kind", help="steps | hits | ... (default: steps)")] = "steps",
by: Annotated[Optional[str], typer.Option("--by", help="Attribution (default: git user.name)")] = None, by: Annotated[str | None, typer.Option("--by", help="Attribution (default: git user.name)")] = None,
date: Annotated[ date: Annotated[
Optional[str], str | None,
typer.Option("--date", help="Override date (default: today, ISO)"), typer.Option("--date", help="Override date (default: today, ISO)"),
] = None, ] = None,
to: Annotated[ to: Annotated[
Optional[str], str | None,
typer.Option( typer.Option(
"--to", "--to",
metavar="schemaN", metavar="schemaN",
@@ -272,7 +271,7 @@ def status(
def update_manifest( def update_manifest(
manifests: Annotated[list[Path], typer.Argument(help="One or more .manifest files to update")], manifests: Annotated[list[Path], typer.Argument(help="One or more .manifest files to update")],
schema: Annotated[ schema: Annotated[
Optional[str], str | None,
typer.Option( typer.Option(
"--schema", "--schema",
metavar="schemaN", metavar="schemaN",
@@ -280,7 +279,7 @@ def update_manifest(
), ),
] = None, ] = None,
gen: Annotated[ gen: Annotated[
Optional[str], str | None,
typer.Option("--gen", metavar="genN", help="Target gen tag (default: keep existing gen)"), typer.Option("--gen", metavar="genN", help="Target gen tag (default: keep existing gen)"),
] = None, ] = None,
execute: Annotated[ execute: Annotated[
@@ -298,11 +297,11 @@ def update_manifest(
def create_manifest( def create_manifest(
files: Annotated[list[Path], typer.Argument(help="Parquet files to include")], files: Annotated[list[Path], typer.Argument(help="Parquet files to include")],
output: Annotated[ output: Annotated[
Optional[Path], Path | None,
typer.Option("--output", "-o", help="Explicit path for the new .manifest file"), typer.Option("--output", "-o", help="Explicit path for the new .manifest file"),
] = None, ] = None,
pool: Annotated[ pool: Annotated[
Optional[str], str | None,
typer.Option( typer.Option(
"--pool", "--pool",
metavar="DETECTOR", metavar="DETECTOR",
@@ -310,7 +309,7 @@ def create_manifest(
), ),
] = None, ] = None,
type_: Annotated[ type_: Annotated[
Optional[PoolType], PoolType | None,
typer.Option("--type", help="Pool type — full, holdout, or dev (required with --pool)"), typer.Option("--type", help="Pool type — full, holdout, or dev (required with --pool)"),
] = None, ] = None,
root: Annotated[Path, typer.Option("--root", help="Dataset root (used with --pool)")] = _DATASET_ROOT_DEFAULT, root: Annotated[Path, typer.Option("--root", help="Dataset root (used with --pool)")] = _DATASET_ROOT_DEFAULT,
@@ -459,7 +458,7 @@ def warm_cache(
typer.Argument(help="Parquet file, directory, or .manifest — same as `giant train`'s"), typer.Argument(help="Parquet file, directory, or .manifest — same as `giant train`'s"),
], ],
config: Annotated[ config: Annotated[
Optional[Path], Path | None,
typer.Option( typer.Option(
"--config", "--config",
"-c", "-c",
@@ -469,7 +468,7 @@ def warm_cache(
), ),
] = None, ] = None,
val_fraction: Annotated[ val_fraction: Annotated[
Optional[float], float | None,
typer.Option( typer.Option(
"--val-fraction", "--val-fraction",
"-f", "-f",
@@ -477,7 +476,7 @@ def warm_cache(
), ),
] = None, ] = None,
seed: Annotated[ seed: Annotated[
Optional[int], int | None,
typer.Option( typer.Option(
"--seed", "--seed",
"-s", "-s",
@@ -485,7 +484,7 @@ def warm_cache(
), ),
] = None, ] = None,
particle_conditioning: Annotated[ particle_conditioning: Annotated[
Optional[Conditioning], Conditioning | None,
typer.Option( typer.Option(
"--particle-conditioning", "--particle-conditioning",
help="Must match the `giant train` run(s)' conditioning.particle.type to warm for. " help="Must match the `giant train` run(s)' conditioning.particle.type to warm for. "
@@ -493,7 +492,7 @@ def warm_cache(
), ),
] = None, ] = None,
material_conditioning: Annotated[ material_conditioning: Annotated[
Optional[Conditioning], Conditioning | None,
typer.Option( typer.Option(
"--material-conditioning", "--material-conditioning",
help="Must match the `giant train` run(s)' conditioning.material.type " help="Must match the `giant train` run(s)' conditioning.material.type "
@@ -502,7 +501,7 @@ def warm_cache(
), ),
] = None, ] = None,
router: Annotated[ router: Annotated[
Optional[bool], bool | None,
typer.Option( typer.Option(
"--router/--no-router", "--router/--no-router",
help="Warm the process vocabulary too (only takes effect with --router-type process). " help="Warm the process vocabulary too (only takes effect with --router-type process). "
@@ -510,11 +509,11 @@ def warm_cache(
), ),
] = None, ] = None,
router_type: Annotated[ router_type: Annotated[
Optional[str], str | None,
typer.Option("--router-type", help="Router implementation name. Not allowed together with --config"), typer.Option("--router-type", help="Router implementation name. Not allowed together with --config"),
] = None, ] = None,
n_experts: Annotated[ n_experts: Annotated[
Optional[int], int | None,
typer.Option("--n-experts", help="Number of routed experts. Not allowed together with --config"), typer.Option("--n-experts", help="Number of routed experts. Not allowed together with --config"),
] = None, ] = None,
rebuild: Annotated[ rebuild: Annotated[
+1 -1
View File
@@ -144,7 +144,7 @@ def run_hparam_scan(
start = time.monotonic() start = time.monotonic()
try: try:
with open(out_dir / "train.log", "a") as log: with open(out_dir / "train.log", "a") as log:
subprocess.run(cmd, env=env, stdout=log, stderr=subprocess.STDOUT) subprocess.run(cmd, env=env, stdout=log, stderr=subprocess.STDOUT, check=False)
except KeyboardInterrupt: except KeyboardInterrupt:
print( print(
f"\ninterrupted during {name} — re-run this script to resume " f"\ninterrupted during {name} — re-run this script to resume "
+1 -1
View File
@@ -32,7 +32,7 @@ MANIFEST_SUFFIX = ".manifest"
# since today's pool assignment is encoded only by *which folder a file's # since today's pool assignment is encoded only by *which folder a file's
# parquet was copied into* — not by anything in the filename itself. # parquet was copied into* — not by anything in the filename itself.
POOL_ASSIGNMENT: dict[str, dict[str, range | list[int]]] = { POOL_ASSIGNMENT: dict[str, dict[str, range | list[int]]] = {
"pbwo4": {"full": range(0, 6), "holdout": range(6, 10)}, "pbwo4": {"full": range(6), "holdout": range(6, 10)},
"sampling_fe_scint": {"dev": [0], "full": [1, 2], "holdout": [3]}, "sampling_fe_scint": {"dev": [0], "full": [1, 2], "holdout": [3]},
"sampling_pb_lar": {"dev": [0], "full": [1, 2], "holdout": [3]}, "sampling_pb_lar": {"dev": [0], "full": [1, 2], "holdout": [3]},
"sampling_pb_scint": {"dev": [0], "full": [1, 2], "holdout": [3]}, "sampling_pb_scint": {"dev": [0], "full": [1, 2], "holdout": [3]},
+1 -1
View File
@@ -106,7 +106,7 @@ def _convert_one(
if output_path is not None: if output_path is not None:
output_path.parent.mkdir(parents=True, exist_ok=True) output_path.parent.mkdir(parents=True, exist_ok=True)
cmd += ["--output", str(output_path)] cmd += ["--output", str(output_path)]
result = subprocess.run(cmd, capture_output=True, text=True) result = subprocess.run(cmd, capture_output=True, text=True, check=False)
return root_file, result.returncode, result.stdout, result.stderr return root_file, result.returncode, result.stdout, result.stderr
+1 -1
View File
@@ -6,8 +6,8 @@ that tests and tooling construct directly.
""" """
from giant.training.checkpoint import build_checkpoint, init_stages_from_checkpoints, 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.loop import train
from giant.training.metrics import MetricsCollector, MetricSpec
from giant.training.trainers import ( from giant.training.trainers import (
FlowDDPMStageTrainer, FlowDDPMStageTrainer,
StageSpec, StageSpec,
+3 -2
View File
@@ -9,9 +9,10 @@
import os import os
import signal import signal
import time import time
from collections.abc import Callable
from pathlib import Path from pathlib import Path
from types import FrameType from types import FrameType
from typing import Callable from typing import Self
import numpy as np import numpy as np
import torch import torch
@@ -47,7 +48,7 @@ class _GracefulShutdown:
Callable[[int, FrameType | None], object] | signal.Handlers | int | None, Callable[[int, FrameType | None], object] | signal.Handlers | int | None,
] = {} ] = {}
def __enter__(self) -> "_GracefulShutdown": def __enter__(self) -> Self:
for sig in _CATCHABLE_SIGNALS: for sig in _CATCHABLE_SIGNALS:
self._previous[sig] = signal.getsignal(sig) self._previous[sig] = signal.getsignal(sig)
signal.signal(sig, self._handle) signal.signal(sig, self._handle)
+1 -1
View File
@@ -188,7 +188,7 @@ class MetricsCollector:
self.fieldnames = self._build_fieldnames() self.fieldnames = self._build_fieldnames()
metrics_path = out_dir / "metrics.csv" metrics_path = out_dir / "metrics.csv"
append = resume and metrics_path.exists() append = resume and metrics_path.exists()
self._file = open(metrics_path, "a" if append else "w", newline="") self._file = open(metrics_path, "a" if append else "w", newline="") # noqa: SIM115 - kept open for the object's lifetime, closed in .close()
self._writer = csv.DictWriter(self._file, fieldnames=self.fieldnames) self._writer = csv.DictWriter(self._file, fieldnames=self.fieldnames)
if not append: if not append:
self._writer.writeheader() self._writer.writeheader()
+1 -1
View File
@@ -47,7 +47,7 @@ class MetricsTable:
columns: dict[str, list[float]] columns: dict[str, list[float]]
@classmethod @classmethod
def load(cls, path: str | Path) -> "MetricsTable": def load(cls, path: str | Path) -> MetricsTable:
with open(path, newline="") as f: with open(path, newline="") as f:
rows = list(csv.DictReader(f)) rows = list(csv.DictReader(f))
epochs = [int(float(r["epoch"])) for r in rows] epochs = [int(float(r["epoch"])) for r in rows]
+2 -2
View File
@@ -20,7 +20,7 @@ from typing import NamedTuple
import torch import torch
import torch.nn.functional as F import torch.nn.functional as F
import torch.optim as optim from torch import optim
from giant.config import ParticleTypeConfig, Stage1ModelConfig, Stage2ModelConfig, TrainConfig from giant.config import ParticleTypeConfig, Stage1ModelConfig, Stage2ModelConfig, TrainConfig
from giant.constants import CONT_SLOT_DIM from giant.constants import CONT_SLOT_DIM
@@ -333,7 +333,7 @@ class StageTrainer:
#: "sampled"` — the stage-1 `StageTrainer` this (stage-2) trainer #: "sampled"` — the stage-1 `StageTrainer` this (stage-2) trainer
#: draws its context sample from. `None` for stage 1 itself, and for #: draws its context sample from. `None` for stage 1 itself, and for
#: stage 2 under "truth". #: stage 2 under "truth".
self.stage1_source: "StageTrainer | None" = None self.stage1_source: StageTrainer | None = None
def attach_stage1(self, stage1_trainer: "StageTrainer") -> None: def attach_stage1(self, stage1_trainer: "StageTrainer") -> None:
"""Wires this (stage-2) trainer to the stage-1 trainer it should """Wires this (stage-2) trainer to the stage-1 trainer it should
+1 -1
View File
@@ -7,7 +7,7 @@ requires-python = ">=3.12"
dependencies = [ dependencies = [
"numpy>=1.26,<3", "numpy>=1.26,<3",
"polars>=1.0,<2", "polars>=1.0,<2",
"pyarrow>=16,<25", "pyarrow>=16,<26",
"tqdm>=4.60,<5", "tqdm>=4.60,<5",
"typer>=0.12,<1", "typer>=0.12,<1",
"pyyaml>=6,<7", "pyyaml>=6,<7",
+11 -11
View File
@@ -17,8 +17,8 @@ import re
from collections.abc import Sequence from collections.abc import Sequence
import torch import torch
import torch.nn as nn
import torch.nn.functional as F import torch.nn.functional as F
from torch import nn
from giant.constants import ( from giant.constants import (
COND_DIM, COND_DIM,
@@ -949,7 +949,7 @@ def _check_router_conditioning_compat(router_types: list[str], conditioning: str
def _build_router_from_cfg(router_cfg: dict, pdg_vocab: int, mat_vocab: int, conditioning: str = "embedding") -> Router: def _build_router_from_cfg(router_cfg: dict, pdg_vocab: int, mat_vocab: int, conditioning: str = "embedding") -> Router:
shared_vocab = dict(pdg_vocab=pdg_vocab, mat_vocab=mat_vocab) shared_vocab = {"pdg_vocab": pdg_vocab, "mat_vocab": mat_vocab}
if router_cfg["type"] == "composed": if router_cfg["type"] == "composed":
axes = _parse_composed_axes(router_cfg) axes = _parse_composed_axes(router_cfg)
_check_router_conditioning_compat([a["type"] for a in axes], conditioning) _check_router_conditioning_compat([a["type"] for a in axes], conditioning)
@@ -977,15 +977,15 @@ def build_models(model_config: dict) -> tuple[nn.Module, nn.Module]:
if router_cfg and router_cfg.get("enabled"): if router_cfg and router_cfg.get("enabled"):
pdg_vocab = model_config["pdg_vocab"] pdg_vocab = model_config["pdg_vocab"]
mat_vocab = model_config["mat_vocab"] mat_vocab = model_config["mat_vocab"]
shared = dict( shared = {
pdg_vocab=pdg_vocab, "pdg_vocab": pdg_vocab,
mat_vocab=mat_vocab, "mat_vocab": mat_vocab,
expert_hidden_dim=model_config.get("expert_hidden_dim") or model_config.get("hidden_dim", 128), "expert_hidden_dim": model_config.get("expert_hidden_dim") or model_config.get("hidden_dim", 128),
expert_n_blocks=model_config.get("expert_n_blocks") or model_config.get("n_blocks", 3), "expert_n_blocks": model_config.get("expert_n_blocks") or model_config.get("n_blocks", 3),
emb_dim=model_config.get("emb_dim", EMB_DIM), "emb_dim": model_config.get("emb_dim", EMB_DIM),
dropout=model_config.get("dropout", 0.1), "dropout": model_config.get("dropout", 0.1),
conditioning=model_config.get("conditioning", "embedding"), "conditioning": model_config.get("conditioning", "embedding"),
) }
conditioning = shared["conditioning"] conditioning = shared["conditioning"]
stage1 = RoutedDenoisingMLP( stage1 = RoutedDenoisingMLP(
router=_build_router_from_cfg(router_cfg, pdg_vocab, mat_vocab, conditioning), router=_build_router_from_cfg(router_cfg, pdg_vocab, mat_vocab, conditioning),
+1 -1
View File
@@ -5,12 +5,12 @@ from pathlib import Path
import pytest import pytest
import torch import torch
from test_train import _base_cfg, _run_train
from giant.model.routers import EnergyRouter from giant.model.routers import EnergyRouter
from giant.model.wgan import gradient_penalty from giant.model.wgan import gradient_penalty
from giant.training.amp import resolve_autocast from giant.training.amp import resolve_autocast
from giant.training.stage2_inputs import _remaining_energy_fraction from giant.training.stage2_inputs import _remaining_energy_fraction
from test_train import _base_cfg, _run_train
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# resolve_autocast # resolve_autocast
+1 -1
View File
@@ -184,7 +184,7 @@ def test_weighted_profile_matches_manual_bincount():
ea = R.entry_axis(lf) ea = R.entry_axis(lf)
lf2 = R.attach_entry_axis(lf, ea) lf2 = R.attach_entry_axis(lf, ea)
edges = np.linspace(0.0, 3.0, 4) # depth bins along +z edges = np.linspace(0.0, 3.0, 4) # depth bins along +z
mean, std = R.weighted_profile(lf2, R.depth_expr(), edges, pl.col("edep")) mean, _ = R.weighted_profile(lf2, R.depth_expr(), edges, pl.col("edep"))
assert mean.shape == (3,) assert mean.shape == (3,)
# totals conserved: sum over bins == mean total edep per event # totals conserved: sum over bins == mean total edep per event
assert np.isclose(mean.sum() * 1, (90.0 + 30.0) / 2) # 2 events assert np.isclose(mean.sum() * 1, (90.0 + 30.0) / 2) # 2 events
+2 -2
View File
@@ -197,7 +197,7 @@ def test_update_manifest_reports_missing_targets(tmp_path):
# schema2 dir exists but the parquet file does not # schema2 dir exists but the parquet file does not
(tmp_path / "processed" / "steps" / "gen1" / "schema2").mkdir(parents=True) (tmp_path / "processed" / "steps" / "gen1" / "schema2").mkdir(parents=True)
lines, missing = plan_update_manifest(manifest, "schema2") _, missing = plan_update_manifest(manifest, "schema2")
assert len(missing) == 1 assert len(missing) == 1
assert "schema2" in str(missing[0]) assert "schema2" in str(missing[0])
@@ -313,7 +313,7 @@ def test_create_manifest_writes_relative_paths(tmp_path):
def test_create_manifest_reports_missing_files(tmp_path): def test_create_manifest_reports_missing_files(tmp_path):
ghost = tmp_path / "processed" / "gen1" / "schema2" / "shard-000.parquet" ghost = tmp_path / "processed" / "gen1" / "schema2" / "shard-000.parquet"
output = tmp_path / "pools" / "full.manifest" output = tmp_path / "pools" / "full.manifest"
lines, missing, _ = plan_create_manifest(output, [ghost]) _, missing, _ = plan_create_manifest(output, [ghost])
assert len(missing) == 1 assert len(missing) == 1
assert missing[0] == ghost.resolve() assert missing[0] == ghost.resolve()
+1 -1
View File
@@ -9,7 +9,7 @@ from pathlib import Path
from typer.testing import CliRunner from typer.testing import CliRunner
import giant.cli as cli from giant import cli
runner = CliRunner() runner = CliRunner()
+1
View File
@@ -1,4 +1,5 @@
import pytest import pytest
from giant.cond_layout import AXIS_TYPES, CondLayout from giant.cond_layout import AXIS_TYPES, CondLayout
from giant.constants import COND_DIM, COND_DIM_BASE, MATERIAL_PHYS_DIM, PARTICLE_PHYS_DIM from giant.constants import COND_DIM, COND_DIM_BASE, MATERIAL_PHYS_DIM, PARTICLE_PHYS_DIM
+1 -1
View File
@@ -576,7 +576,7 @@ def test_save_config_round_trips_three_level_nesting(tmp_path):
# default_out_dir_name # default_out_dir_name
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
_NOW = datetime(2026, 7, 29, 14, 30) _NOW = datetime(2026, 7, 29, 14, 30) # noqa: DTZ001 - naive, matching default_out_dir_name's naive datetime.now()
def _cfg_with(**dotted_overrides): def _cfg_with(**dotted_overrides):
+1 -1
View File
@@ -90,7 +90,7 @@ def _collect_names(source: str, filename: str) -> set[str]:
names.add(node.attr) names.add(node.attr)
elif isinstance(node, ast.Constant) and isinstance(node.value, str) and id(node) not in docstring_ids: elif isinstance(node, ast.Constant) and isinstance(node.value, str) and id(node) not in docstring_ids:
names.add(node.value) names.add(node.value)
elif isinstance(node, ast.arg): elif isinstance(node, ast.arg): # noqa: SIM114 - kept separate so ty narrows node.arg to str, not str | None
names.add(node.arg) names.add(node.arg)
elif isinstance(node, ast.keyword) and node.arg is not None: elif isinstance(node, ast.keyword) and node.arg is not None:
names.add(node.arg) names.add(node.arg)
+1 -1
View File
@@ -1,3 +1,4 @@
from test_pipeline import _make_synthetic_steps
from typer.testing import CliRunner from typer.testing import CliRunner
from giant import cli as giant_cli from giant import cli as giant_cli
@@ -5,7 +6,6 @@ from giant.config import Conditioning
from giant.data import setup_cache from giant.data import setup_cache
from giant.tools import dwarf from giant.tools import dwarf
from giant.tools.dwarf import app from giant.tools.dwarf import app
from test_pipeline import _make_synthetic_steps
runner = CliRunner() runner = CliRunner()
+2 -1
View File
@@ -1,9 +1,10 @@
import torch import torch
from giant.config import ConditioningAxisConfig from giant.config import ConditioningAxisConfig
from giant.constants import COND_DIM from giant.constants import COND_DIM
from giant.model.network import Stage1Model from giant.model.network import Stage1Model
from giant.model.schedule import CosineSchedule, flow_matching_loss from giant.model.schedule import CosineSchedule, flow_matching_loss
from giant.sample import sample_flow, sample_ddim from giant.sample import sample_ddim, sample_flow
PARTICLE_CFG = ConditioningAxisConfig(type="physical", emb_dim=8, n_layers=1) PARTICLE_CFG = ConditioningAxisConfig(type="physical", emb_dim=8, n_layers=1)
MATERIAL_CFG = ConditioningAxisConfig(type="physical", emb_dim=8, n_layers=1) MATERIAL_CFG = ConditioningAxisConfig(type="physical", emb_dim=8, n_layers=1)
+14 -13
View File
@@ -2,8 +2,9 @@ import copy
import pytest import pytest
import torch import torch
from giant import config as gconfig from giant import config as gconfig
from giant.constants import CONT_SLOT_DIM, COND_DIM, PARTICLE_PHYS_DIM, SEC_SLOT_DIM from giant.constants import COND_DIM, CONT_SLOT_DIM, PARTICLE_PHYS_DIM, SEC_SLOT_DIM
from giant.model.network import ( from giant.model.network import (
HISTORY_REGISTRY, HISTORY_REGISTRY,
AttentionHistory, AttentionHistory,
@@ -1217,18 +1218,18 @@ def test_stage_classes_are_stagemodel_subclasses(cls):
@pytest.mark.parametrize("cls", [Stage1Model, Stage2OneShot, Stage2Autoregressive]) @pytest.mark.parametrize("cls", [Stage1Model, Stage2OneShot, Stage2Autoregressive])
@pytest.mark.parametrize("generator", ["flow", "ddpm", "wgan"]) @pytest.mark.parametrize("generator", ["flow", "ddpm", "wgan"])
def test_stagemodel_time_emb_matches_objective_needs_time(cls, generator): def test_stagemodel_time_emb_matches_objective_needs_time(cls, generator):
kwargs = dict( kwargs = {
pdg_vocab=5, "pdg_vocab": 5,
mat_vocab=3, "mat_vocab": 3,
particle_cfg=PARTICLE_CFG, "particle_cfg": PARTICLE_CFG,
material_cfg=MATERIAL_CFG, "material_cfg": MATERIAL_CFG,
hidden_dim=_STAGE_HIDDEN_DIM, "hidden_dim": _STAGE_HIDDEN_DIM,
n_res_blocks=_STAGE_N_BLOCKS, "n_res_blocks": _STAGE_N_BLOCKS,
cond_out_dim=_STAGE_COND_OUT_DIM, "cond_out_dim": _STAGE_COND_OUT_DIM,
generator=generator, "generator": generator,
time_dim=8, "time_dim": 8,
noise_dim=8, "noise_dim": 8,
) }
if cls is Stage1Model: if cls is Stage1Model:
kwargs["n_sec_head_k_max"] = 15 kwargs["n_sec_head_k_max"] = 15
else: else:
+3 -4
View File
@@ -20,7 +20,6 @@ from giant.model.schedule import (
) )
from giant.sample import sample_secondaries from giant.sample import sample_secondaries
# ── helpers ────────────────────────────────────────────────────────────────── # ── helpers ──────────────────────────────────────────────────────────────────
@@ -341,7 +340,7 @@ def test_encode_secondaries_energy_conservation():
def test_encode_secondaries_stick_logits_match_naive_reference(): def test_encode_secondaries_stick_logits_match_naive_reference():
"""Cumsum-based remaining-budget computation must match a naive """Cumsum-based remaining-budget computation must match a naive
per-row, per-slot Python reference (no cumsum) within float tolerance.""" per-row, per-slot Python reference (no cumsum) within float tolerance."""
from giant.data.transforms import encode_secondaries, _EPS, _STICK_LOGIT_CLIP from giant.data.transforms import _EPS, _STICK_LOGIT_CLIP, encode_secondaries
rng = np.random.default_rng(11) rng = np.random.default_rng(11)
N = 25 N = 25
@@ -569,7 +568,7 @@ def test_decode_secondaries_degenerate_row_falls_back_to_even_split():
e_sec = np.array([0.0, 4.0, 9.0, 30.0], dtype=np.float32) e_sec = np.array([0.0, 4.0, 9.0, 30.0], dtype=np.float32)
pre_dir = np.tile([0.0, 0.0, 1.0], (N, 1)).astype(np.float32) pre_dir = np.tile([0.0, 0.0, 1.0], (N, 1)).astype(np.float32)
sec_E, _sec_dir, _mass, _charge, sec_valid = decode_secondaries(sec_cont, n_sec, e_sec, pre_dir) sec_E, _sec_dir, _mass, _charge, _ = decode_secondaries(sec_cont, n_sec, e_sec, pre_dir)
for i, k in enumerate(n_sec): for i, k in enumerate(n_sec):
if k == 0: if k == 0:
@@ -593,7 +592,7 @@ def test_decode_secondaries_rescale_preserves_relative_shares():
n_sec = np.array([4]) n_sec = np.array([4])
pre_dir = np.array([[0.0, 0.0, 1.0]], dtype=np.float32) pre_dir = np.array([[0.0, 0.0, 1.0]], dtype=np.float32)
sec_E_small, _, _, _, sec_valid = decode_secondaries(sec_cont, n_sec, np.array([5.0], dtype=np.float32), pre_dir) sec_E_small, _, _, _, _ = decode_secondaries(sec_cont, n_sec, np.array([5.0], dtype=np.float32), pre_dir)
sec_E_large, _, _, _, _ = decode_secondaries(sec_cont, n_sec, np.array([50.0], dtype=np.float32), pre_dir) sec_E_large, _, _, _, _ = decode_secondaries(sec_cont, n_sec, np.array([50.0], dtype=np.float32), pre_dir)
ratio_small = sec_E_small[0, :4] / sec_E_small[0, 0] ratio_small = sec_E_small[0, :4] / sec_E_small[0, 0]
+2 -2
View File
@@ -9,8 +9,8 @@ import pytest
pytest.importorskip("plotstyle") pytest.importorskip("plotstyle")
from giant.analysis import render as render_mod # noqa: E402 from giant.analysis import render as render_mod
from giant.analysis.reduced import Reduced # noqa: E402 from giant.analysis.reduced import Reduced
def _try_render(reduced: list[Reduced], out: Path): def _try_render(reduced: list[Reduced], out: Path):
+2 -2
View File
@@ -9,7 +9,7 @@ import pytest
import torch import torch
from giant.config import ConditioningAxisConfig, ParticleTypeConfig from giant.config import ConditioningAxisConfig, ParticleTypeConfig
from giant.constants import TERM_ESCAPED, TERM_MAX_STEPS, TERM_UNKNOWN_PDG, K_MAX from giant.constants import K_MAX, TERM_ESCAPED, TERM_MAX_STEPS, TERM_UNKNOWN_PDG
from giant.data.loader import TopNMap from giant.data.loader import TopNMap
from giant.data.transforms import Normalizer from giant.data.transforms import Normalizer
from giant.model.network import ( from giant.model.network import (
@@ -21,7 +21,7 @@ from giant.model.network import (
from giant.rollout import L1DistCollector, make_seed_frontier, rollout from giant.rollout import L1DistCollector, make_seed_frontier, rollout
pytest.importorskip("sklearn") pytest.importorskip("sklearn")
from giant import geometry as g # noqa: E402 from giant import geometry as g
PDG_MAP = {22: 0, 11: 1, -11: 2} PDG_MAP = {22: 0, 11: 1, -11: 2}
MAT_MAP = {"G4_AIR": 0, "G4_PbWO4": 1} MAT_MAP = {"G4_AIR": 0, "G4_PbWO4": 1}
+6 -8
View File
@@ -1,5 +1,7 @@
"""Tests for the mixture-of-experts routing prototype (giant/model/network.py).""" """Tests for the mixture-of-experts routing prototype (giant/model/network.py)."""
import itertools
import pytest import pytest
import torch import torch
@@ -7,6 +9,7 @@ from giant.config import ConditioningAxisConfig
from giant.constants import COND_DIM, K_MAX, SEC_DIM, X_DIM from giant.constants import COND_DIM, K_MAX, SEC_DIM, X_DIM
from giant.model.network import ( from giant.model.network import (
BLOCK_REGISTRY, BLOCK_REGISTRY,
ROUTER_REGISTRY,
TRUNK_REGISTRY, TRUNK_REGISTRY,
AdaLNResBlock, AdaLNResBlock,
ComposedRouter, ComposedRouter,
@@ -17,7 +20,6 @@ from giant.model.network import (
NoneRouter, NoneRouter,
PdgRouter, PdgRouter,
ProcessRouter, ProcessRouter,
ROUTER_REGISTRY,
ResBlock, ResBlock,
RoutedTrunk, RoutedTrunk,
Stage1Model, Stage1Model,
@@ -397,7 +399,7 @@ def test_energy_router_own_width_controls_own_coverage_independent_of_others():
router.raw_width[0] = raw router.raw_width[0] = raw
shares.append(router.gate(cond_cont, cond_cat)[0, 0].item()) shares.append(router.gate(cond_cont, cond_cat)[0, 0].item())
assert all(a <= b + 1e-6 for a, b in zip(shares, shares[1:])) assert all(a <= b + 1e-6 for a, b in itertools.pairwise(shares))
def test_build_router_threads_learn_width_kwargs_through(): def test_build_router_threads_learn_width_kwargs_through():
@@ -1010,9 +1012,7 @@ def test_build_models_routed_pair_composed_router_is_drop_in_for_sample_flow():
n_sec_pred = stage2.predict_n_sec(cond_cont, cond_cat, stage1_norm).argmax(dim=-1) n_sec_pred = stage2.predict_n_sec(cond_cont, cond_cat, stage1_norm).argmax(dim=-1)
assert n_sec_pred.shape == (B,) assert n_sec_pred.shape == (B,)
sec_cont, sec_type_emb, sec_valid = sample_secondaries( sec_cont, _, sec_valid = sample_secondaries(stage2, cond_cont, cond_cat, stage1_norm, n_sec_pred, steps=2)
stage2, cond_cont, cond_cat, stage1_norm, n_sec_pred, steps=2
)
assert sec_cont.shape == (B, K_MAX, 4) assert sec_cont.shape == (B, K_MAX, 4)
assert sec_valid.shape == (B, K_MAX) assert sec_valid.shape == (B, K_MAX)
@@ -1261,9 +1261,7 @@ def test_build_models_routed_pair_is_drop_in_for_sample_flow():
n_sec_pred = stage2.predict_n_sec(cond_cont, cond_cat, stage1_norm).argmax(dim=-1) n_sec_pred = stage2.predict_n_sec(cond_cont, cond_cat, stage1_norm).argmax(dim=-1)
assert n_sec_pred.shape == (B,) assert n_sec_pred.shape == (B,)
sec_cont, sec_type_emb, sec_valid = sample_secondaries( sec_cont, _, sec_valid = sample_secondaries(stage2, cond_cont, cond_cat, stage1_norm, n_sec_pred, steps=2)
stage2, cond_cont, cond_cat, stage1_norm, n_sec_pred, steps=2
)
assert sec_cont.shape == (B, K_MAX, 4) assert sec_cont.shape == (B, K_MAX, 4)
assert sec_valid.shape == (B, K_MAX) assert sec_valid.shape == (B, K_MAX)
+4 -4
View File
@@ -262,7 +262,7 @@ def test_sample_secondaries_ar_first_slot_has_no_history():
cond_cont, cond_cat = _cond(B) cond_cont, cond_cat = _cond(B)
stage1_out = torch.randn(B, X_DIM) stage1_out = torch.randn(B, X_DIM)
n_sec_pred = torch.tensor([0, 1, 1]) n_sec_pred = torch.tensor([0, 1, 1])
sec_cont, sec_type, sec_valid = sample_secondaries_ar(decoder, cond_cont, cond_cat, stage1_out, n_sec_pred, steps=2) sec_cont, _, sec_valid = sample_secondaries_ar(decoder, cond_cont, cond_cat, stage1_out, n_sec_pred, steps=2)
assert sec_cont.shape == (B, 1, CONT_SLOT_DIM) assert sec_cont.shape == (B, 1, CONT_SLOT_DIM)
assert sec_valid.tolist() == [[False], [True], [True]] assert sec_valid.tolist() == [[False], [True], [True]]
@@ -281,7 +281,7 @@ def test_sample_secondaries_ar_stop_token_forced_stop_gives_zero_secondaries(n_s
_force_stop_head_logit(decoder, 50.0) _force_stop_head_logit(decoder, 50.0)
cond_cont, cond_cat = _cond(B) cond_cont, cond_cat = _cond(B)
stage1_out = torch.randn(B, X_DIM) stage1_out = torch.randn(B, X_DIM)
sec_cont, sec_type, sec_valid = sample_secondaries_ar(decoder, cond_cont, cond_cat, stage1_out, None, steps=2) _, _, sec_valid = sample_secondaries_ar(decoder, cond_cont, cond_cat, stage1_out, None, steps=2)
assert sec_valid.shape == (B, k_max) assert sec_valid.shape == (B, k_max)
assert not sec_valid.any() assert not sec_valid.any()
@@ -296,7 +296,7 @@ def test_sample_secondaries_ar_stop_token_forced_never_stop_runs_to_k_max(n_sec_
_force_stop_head_logit(decoder, -50.0) _force_stop_head_logit(decoder, -50.0)
cond_cont, cond_cat = _cond(B) cond_cont, cond_cat = _cond(B)
stage1_out = torch.randn(B, X_DIM) stage1_out = torch.randn(B, X_DIM)
sec_cont, sec_type, sec_valid = sample_secondaries_ar(decoder, cond_cont, cond_cat, stage1_out, None, steps=2) _, _, sec_valid = sample_secondaries_ar(decoder, cond_cont, cond_cat, stage1_out, None, steps=2)
assert sec_valid.all() assert sec_valid.all()
@@ -412,7 +412,7 @@ def test_sample_secondaries_ar_full_length_ignores_n_sec_pred_zero_rows():
cond_cont, cond_cat = _cond(B) cond_cont, cond_cat = _cond(B)
stage1_out = torch.randn(B, X_DIM) stage1_out = torch.randn(B, X_DIM)
n_sec_pred = torch.tensor([0, 0, 0]) n_sec_pred = torch.tensor([0, 0, 0])
sec_cont, sec_type, sec_valid = sample_secondaries_ar( sec_cont, _, sec_valid = sample_secondaries_ar(
decoder, cond_cont, cond_cat, stage1_out, n_sec_pred, steps=2, full_length=True decoder, cond_cont, cond_cat, stage1_out, n_sec_pred, steps=2, full_length=True
) )
assert not sec_valid.any() assert not sec_valid.any()
+14 -15
View File
@@ -12,6 +12,7 @@ import pytest
import torch import torch
import torch.nn.functional as F import torch.nn.functional as F
from giant.checkpoint_io import load_for_inference
from giant.config import ParticleTypeConfig from giant.config import ParticleTypeConfig
from giant.constants import ( from giant.constants import (
COND_DIM, COND_DIM,
@@ -21,7 +22,6 @@ from giant.constants import (
SEC_SLOT_DIM, SEC_SLOT_DIM,
X_DIM, X_DIM,
) )
from giant.checkpoint_io import load_for_inference
from giant.data.dataset import StepBatch from giant.data.dataset import StepBatch
from giant.data.transforms import Normalizer from giant.data.transforms import Normalizer
from giant.model.network import Stage2Autoregressive, build_critics, build_models from giant.model.network import Stage2Autoregressive, build_critics, build_models
@@ -36,7 +36,6 @@ from giant.training import (
train, train,
) )
from giant.training.metrics import _wandb_run_config from giant.training.metrics import _wandb_run_config
from giant.training.trainers import _type_class_weight_vector
from giant.training.stage2_inputs import ( from giant.training.stage2_inputs import (
_ar_has_prev, _ar_has_prev,
_assemble_stage2_ar_inputs, _assemble_stage2_ar_inputs,
@@ -51,6 +50,7 @@ from giant.training.stage2_inputs import (
_stop_target_and_mask, _stop_target_and_mask,
_type_repr, _type_repr,
) )
from giant.training.trainers import _type_class_weight_vector
PDG_VOCAB = 6 PDG_VOCAB = 6
MAT_VOCAB = 3 MAT_VOCAB = 3
@@ -543,19 +543,18 @@ def test_train_raises_when_no_active_stage():
model_config = _model_config(cfg) model_config = _model_config(cfg)
models = build_models(model_config) models = build_models(model_config)
critics = build_critics(model_config) critics = build_critics(model_config)
with tempfile.TemporaryDirectory() as tmp: with tempfile.TemporaryDirectory() as tmp, pytest.raises(ValueError, match="no active stage"):
with pytest.raises(ValueError, match="no active stage"): train(
train( cfg=cfg,
cfg=cfg, models=models,
models=models, critics=critics,
critics=critics, train_loader=_fake_batches(1, 8),
train_loader=_fake_batches(1, 8), val_loader=_fake_batches(1, 8),
val_loader=_fake_batches(1, 8), device=torch.device("cpu"),
device=torch.device("cpu"), out_dir=Path(tmp) / "run",
out_dir=Path(tmp) / "run", model_config=model_config,
model_config=model_config, total_train_batches=1,
total_train_batches=1, )
)
def test_metrics_csv_columns_are_stage_prefixed(): def test_metrics_csv_columns_are_stage_prefixed():
+2 -2
View File
@@ -11,8 +11,8 @@ import pytest
pytest.importorskip("plotstyle") pytest.importorskip("plotstyle")
from giant.training import plots as plots_mod # noqa: E402 from giant.training import plots as plots_mod
from giant.training.plots import MetricsTable, derive_metrics_dir, render_metrics # noqa: E402 from giant.training.plots import MetricsTable, derive_metrics_dir, render_metrics
# --- fixtures ---------------------------------------------------------- # --- fixtures ----------------------------------------------------------
+4 -3
View File
@@ -2,9 +2,13 @@ import warnings
import numpy as np import numpy as np
import pytest import pytest
from giant.cond_layout import AXIS_TYPES, CondLayout from giant.cond_layout import AXIS_TYPES, CondLayout
from giant.constants import COND_DIM, COND_DIM_BASE, K_MAX from giant.constants import COND_DIM, COND_DIM_BASE, K_MAX
from giant.data.transforms import ( from giant.data.transforms import (
Normalizer,
_vectorized_map_lookup,
_WelfordAccumulator,
build_cond_features, build_cond_features,
build_features, build_features,
encode_secondaries, encode_secondaries,
@@ -14,12 +18,9 @@ from giant.data.transforms import (
inv_log_transform, inv_log_transform,
local_frame_rotation, local_frame_rotation,
log_transform, log_transform,
Normalizer,
reconstruct_post_pos, reconstruct_post_pos,
sorted_membership, sorted_membership,
travel_direction, travel_direction,
_vectorized_map_lookup,
_WelfordAccumulator,
) )
Generated
+1338 -1089
View File
File diff suppressed because it is too large Load Diff