Compare commits
2 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| ecd6347fca | |||
| c984d0a19d |
+11
-11
@@ -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",
|
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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()))
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -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"])
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
@@ -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()))
|
||||||
|
|||||||
@@ -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
|
||||||
|
|
||||||
|
|||||||
@@ -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.
|
||||||
|
|||||||
@@ -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
@@ -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
@@ -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),
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
@@ -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,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
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
@@ -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):
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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",
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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":
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
@@ -1,4 +1,4 @@
|
|||||||
from typing import Callable
|
from collections.abc import Callable
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
|
|||||||
+3
-3
@@ -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
@@ -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
@@ -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.
|
||||||
|
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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):
|
||||||
|
|||||||
@@ -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
@@ -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[
|
||||||
|
|||||||
@@ -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 "
|
||||||
|
|||||||
@@ -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]},
|
||||||
|
|||||||
@@ -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
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
@@ -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]
|
||||||
|
|||||||
@@ -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
@@ -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",
|
||||||
|
|||||||
@@ -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
@@ -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
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|
||||||
|
|||||||
@@ -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,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
|
||||||
|
|
||||||
|
|||||||
@@ -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):
|
||||||
|
|||||||
@@ -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
@@ -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
@@ -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
@@ -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:
|
||||||
|
|||||||
@@ -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]
|
||||||
|
|||||||
@@ -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):
|
||||||
|
|||||||
@@ -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}
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|
||||||
|
|||||||
@@ -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
@@ -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():
|
||||||
|
|||||||
@@ -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 ----------------------------------------------------------
|
||||||
|
|
||||||
|
|||||||
@@ -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,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user