From 516a8a9ee1ba430a389c8886e3ccd7b94db49a4f Mon Sep 17 00:00:00 2001 From: Lars Bogner Date: Fri, 28 Aug 2026 14:10:19 +0200 Subject: [PATCH] perf: defer heavy imports in giant/dwarf CLIs until commands run torch/pandas/pyarrow/polars/uproot/awkward/particle were all imported at module scope in giant/cli.py and giant/tools/dwarf.py, so even `--help` paid ~1.6-1.9s of import cost. Move those imports into the command bodies that actually need them (following the deferred-import pattern already used for analysis/render/plots/sklearn/wandb), cutting `giant --help` to ~0.3s and `dwarf --help` to ~0.2s with no change to any command's actual behavior. Co-Authored-By: Claude Sonnet 5 --- giant/cli.py | 79 +++++++++++++++++++------------ giant/config.py | 21 ++++++-- giant/tools/dwarf.py | 48 +++++++++++++------ tests/test_cli_train_overrides.py | 12 ++--- 4 files changed, 106 insertions(+), 54 deletions(-) diff --git a/giant/cli.py b/giant/cli.py index e45b78b..9ac1e0e 100644 --- a/giant/cli.py +++ b/giant/cli.py @@ -1,21 +1,19 @@ +from __future__ import annotations + from collections import Counter from datetime import datetime, timezone from enum import Enum import math from pathlib import Path import re -from typing import Optional +from typing import TYPE_CHECKING, Optional import uuid as uuid_mod -import numpy as np -import yaml -import torch import typer from typing_extensions import Annotated -import pyarrow as pa -import pyarrow.parquet as pq -from tqdm import tqdm +if TYPE_CHECKING: + import numpy as np from giant import config as gconfig from giant.constants import ( @@ -25,30 +23,11 @@ from giant.constants import ( PREDICT_SCHEMA_VERSION_KEY, ROLLOUT_COORD_VALUE, ) -from giant.data.loader import ( - event_id_offset, - find_parquet_files, - iter_file_chunks, - iter_cond_chunks, -) -from giant.data.transforms import ( - build_features, - build_cond_features, - energy_simplex_decode, - inv_local_frame_rotation, - inv_log_transform, - reconstruct_post_pos, -) -from giant.checkpoint_io import CheckpointCompatibilityError, load_for_inference -from giant.geometry import GeometryOracle + +# giant.materials only pulls in numpy (no torch/pandas), and MATERIAL_PROPERTIES +# is needed at decoration time below (a Typer option default), so it can't be +# deferred into a command body like the rest of this module's heavy imports. from giant.materials import MATERIAL_PROPERTIES -from giant.pipeline import run_train_job -from giant.rollout import ( - L1DistCollector, - decode_secondary_identity, - rollout as run_rollout, -) -from giant.sample import resolve_n_sec, sample_stage1, sample_stage2 app = typer.Typer(no_args_is_help=True) @@ -210,6 +189,8 @@ def _write_prediction_ref( comment: str | None = None, ) -> Path: """Write a YAML sidecar in the checkpoint directory and return its path.""" + import yaml + ref = { "prediction_id": pred_uuid, "output": str(out), @@ -639,6 +620,10 @@ def train( ] = None, ) -> None: """Train the GIANT surrogate model.""" + import torch + + from giant.pipeline import run_train_job + batch_size_auto = False batch_size_value: Optional[int] = None if batch_size is not None: @@ -1048,6 +1033,25 @@ def predict( ] = None, ) -> None: """Run trained model on a parquet file and save predictions.""" + import numpy as np + import pyarrow as pa + import pyarrow.parquet as pq + import torch + from tqdm import tqdm + + from giant.checkpoint_io import CheckpointCompatibilityError, load_for_inference + from giant.data.loader import event_id_offset, find_parquet_files, iter_cond_chunks, iter_file_chunks + from giant.data.transforms import ( + build_cond_features, + build_features, + energy_simplex_decode, + inv_local_frame_rotation, + inv_log_transform, + reconstruct_post_pos, + ) + from giant.rollout import decode_secondary_identity + from giant.sample import resolve_n_sec, sample_stage1, sample_stage2 + batch_size_auto = False batch_size_value: Optional[int] = None if batch_size.strip().lower() == "auto": @@ -1343,6 +1347,10 @@ def _seed_from_data(files: list[Path], n_events: int | None) -> dict[str, np.nda the codebase's convention for the primary (a secondary always carries less energy than its parent). See giant/analysis/reduce.py:entry_axis. """ + import numpy as np + + from giant.data.loader import event_id_offset, iter_cond_chunks + best_E: dict[int, float] = {} best: dict[int, tuple] = {} for file_idx, path in enumerate(files): @@ -1447,6 +1455,17 @@ def rollout( ] = None, ) -> None: """Roll the surrogate forward into full showers (autoregressive).""" + import numpy as np + import pyarrow as pa + import pyarrow.parquet as pq + import torch + import yaml + + from giant.checkpoint_io import CheckpointCompatibilityError, load_for_inference + from giant.data.loader import find_parquet_files + from giant.geometry import GeometryOracle + from giant.rollout import L1DistCollector, rollout as run_rollout + if seed is not None: torch.manual_seed(seed) np.random.seed(seed) diff --git a/giant/config.py b/giant/config.py index d94ff64..8661e09 100644 --- a/giant/config.py +++ b/giant/config.py @@ -1,3 +1,5 @@ +from __future__ import annotations + import copy import difflib import hashlib @@ -10,12 +12,12 @@ from dataclasses import dataclass, field from datetime import datetime, timezone from enum import Enum from pathlib import Path - -import numpy as np -import torch +from typing import TYPE_CHECKING from giant._migration import V02_FIXED_FACTS, V02_MODEL_KEY_TO_STAGES, reject_legacy_router_expert_sizing -from giant.model.history import HISTORY_REGISTRY + +if TYPE_CHECKING: + import torch class Conditioning(str, Enum): @@ -941,6 +943,8 @@ def git_hash() -> str: def auto_device() -> torch.device: + import torch + if torch.cuda.is_available(): return torch.device("cuda") if torch.backends.mps.is_available(): @@ -986,6 +990,8 @@ def estimate_batch_size( inference (e.g. `predict`), which uses a much lower per-sample memory calibration since there's no backward graph or optimizer state. """ + import torch + if device.type != "cuda": raise ValueError(f"--batch-size auto is only supported on cuda devices, got {device.type!r}") device_index = device.index if device.index is not None else torch.cuda.current_device() @@ -1680,6 +1686,8 @@ def validate_config(cfg: dict, *, resume: bool = False) -> None: "'energy_desc' (the only implemented ordering; see " "AutoregressiveConfig.order's docstring)" ) + from giant.model.history import HISTORY_REGISTRY + history = _get_path(cfg, "stage2_model.autoregressive.history") if history not in HISTORY_REGISTRY: raise ValueError( @@ -1881,6 +1889,9 @@ def resolve_default_out_dir(cfg: dict, base: Path = Path("checkpoints")) -> Path def seed_everything(seed: int) -> None: + import numpy as np + import torch + random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) @@ -1934,6 +1945,8 @@ def build_run_meta( n_val_events: int, n_train_steps: int, ) -> dict: + import torch + return { "config_version": CONFIG_VERSION, "git_hash": git_hash(), diff --git a/giant/tools/dwarf.py b/giant/tools/dwarf.py index 29b0661..8840021 100644 --- a/giant/tools/dwarf.py +++ b/giant/tools/dwarf.py @@ -5,6 +5,8 @@ simulation-fanout tools into one Typer app so there's a single command name (and `--help`) to remember instead of five differently-hyphenated ones. """ +from __future__ import annotations + import os from enum import Enum from pathlib import Path @@ -14,20 +16,15 @@ import typer from typing_extensions import Annotated from giant.config import Conditioning -from giant.tools.bump_dataset_version import ( - run_bump_gen, - run_bump_schema, - run_create_manifest, - run_status, - run_update_manifest, -) -from giant.tools.create_root_files import run_make_root -from giant.tools.geometry_oracle import run_build_geometry_oracle -from giant.tools.hparam_scan import DATA_DEFAULT, SCAN_DIR_DEFAULT, run_hparam_scan -from giant.tools.migrate_geant_steps import run_migration -from giant.tools.steps_to_parquet import convert_steps_to_parquet -from giant.tools.steps_to_parquet_parallel import run_parallel_job -from giant.tools.warm_setup_cache import run_warm_setup_cache + +# DATA_DEFAULT/SCAN_DIR_DEFAULT are Typer option defaults (evaluated at +# decoration time below), so that one name has to stay eager — the module +# itself is stdlib-only, so it costs nothing. Every other giant.tools.* +# import here is deferred into the one command body that uses it, since +# several (steps_to_parquet: uproot/awkward/polars; warm_setup_cache: +# giant.pipeline -> torch; geometry_oracle: pandas) are expensive and +# `dwarf --help`/tab-completion shouldn't pay for all of them upfront. +from giant.tools.hparam_scan import DATA_DEFAULT, SCAN_DIR_DEFAULT app = typer.Typer(no_args_is_help=True) @@ -121,6 +118,9 @@ def convert( ] = None, ) -> None: """Convert ROOT Steps tree(s) to Parquet.""" + from giant.tools.steps_to_parquet import convert_steps_to_parquet + from giant.tools.steps_to_parquet_parallel import run_parallel_job + if jobs < 1: typer.echo("error: --jobs must be >= 1", err=True) raise typer.Exit(1) @@ -183,6 +183,8 @@ def migrate( ] = False, ) -> None: """One-time migration into the versioned raw/processed/pools/derived layout.""" + from giant.tools.migrate_geant_steps import run_migration + run_migration(str(root), execute=execute, copy=copy) @@ -207,6 +209,8 @@ def bump_gen( root: Annotated[Path, typer.Option("--root", help="Dataset root")] = _DATASET_ROOT_DEFAULT, ) -> None: """Cut a new raw generation.""" + from giant.tools.bump_dataset_version import run_bump_gen + run_bump_gen( kind=kind, reason=reason, @@ -240,6 +244,8 @@ def bump_schema( root: Annotated[Path, typer.Option("--root", help="Dataset root")] = _DATASET_ROOT_DEFAULT, ) -> None: """Cut a new schema within a gen.""" + from giant.tools.bump_dataset_version import run_bump_schema + run_bump_schema( kind=kind, gen=gen, @@ -257,6 +263,8 @@ def status( root: Annotated[Path, typer.Option("--root", help="Dataset root")] = _DATASET_ROOT_DEFAULT, ) -> None: """List existing gens/schemas per kind.""" + from giant.tools.bump_dataset_version import run_status + run_status(str(root)) @@ -281,6 +289,8 @@ def update_manifest( ] = False, ) -> None: """Repoint manifest(s) to a new gen and/or schema, verifying all target files exist.""" + from giant.tools.bump_dataset_version import run_update_manifest + run_update_manifest([str(m) for m in manifests], schema=schema, execute=execute, gen=gen) @@ -311,6 +321,8 @@ def create_manifest( ] = False, ) -> None: """Create a new manifest from a list of parquet files.""" + from giant.tools.bump_dataset_version import run_create_manifest + run_create_manifest( [str(f) for f in files], execute=execute, @@ -358,6 +370,8 @@ def make_root( ] = False, ) -> None: """Generate new ROOT shards via a minicalosim executable.""" + from giant.tools.create_root_files import run_make_root + _warn_if_exceeds_shared_quota(jobs, "--jobs") run_make_root( executable=executable, @@ -423,6 +437,8 @@ def build_geometry_oracle( ] = 2000, ) -> None: """Fit a position -> (material, layer_id) oracle for `giant rollout`.""" + from giant.tools.geometry_oracle import run_build_geometry_oracle + run_build_geometry_oracle( data=data, out=out, @@ -515,6 +531,8 @@ def warm_cache( such entry across every run) skips straight to training. See giant/data/setup_cache.py. """ + from giant.tools.warm_setup_cache import run_warm_setup_cache + flag_overrides = { "--val-fraction": val_fraction, "--seed": seed, @@ -557,6 +575,8 @@ def hparam_scan( dry_run: Annotated[bool, typer.Option("--dry-run")] = False, ) -> None: """Grid-scan dropout x n_blocks x hidden_dim via sequential `giant train` runs.""" + from giant.tools.hparam_scan import run_hparam_scan + run_hparam_scan(data=data, scan_dir=scan_dir, seed=seed, dry_run=dry_run) diff --git a/tests/test_cli_train_overrides.py b/tests/test_cli_train_overrides.py index 7e55193..6cec1c2 100644 --- a/tests/test_cli_train_overrides.py +++ b/tests/test_cli_train_overrides.py @@ -20,7 +20,7 @@ def _invoke_and_capture_cfg(monkeypatch, tmp_path: Path, args: list[str]) -> dic def _fake_run_train_job(*, data, cfg, out_dir, **kwargs): captured["cfg"] = cfg - monkeypatch.setattr(cli, "run_train_job", _fake_run_train_job) + monkeypatch.setattr("giant.pipeline.run_train_job", _fake_run_train_job) result = runner.invoke( cli.app, @@ -125,7 +125,7 @@ def test_stage2_init_from_and_freeze_flags_land_in_cfg_and_dont_touch_stage1(mon def test_batch_size_invalid_string_errors(monkeypatch, tmp_path): - monkeypatch.setattr(cli, "run_train_job", lambda *a, **kw: None) + monkeypatch.setattr("giant.pipeline.run_train_job", lambda *a, **kw: None) result = runner.invoke( cli.app, ["train", "dummy.parquet", "--out", str(tmp_path / "run"), "--batch-size", "not-a-number"], @@ -140,7 +140,7 @@ def test_out_dir_resolution_prefers_explicit_out_over_resume(monkeypatch, tmp_pa def _fake_run_train_job(*, data, cfg, out_dir, **kwargs): captured["out_dir"] = out_dir - monkeypatch.setattr(cli, "run_train_job", _fake_run_train_job) + monkeypatch.setattr("giant.pipeline.run_train_job", _fake_run_train_job) resume_dir = tmp_path / "resumed_run" resume_dir.mkdir() @@ -161,7 +161,7 @@ def test_out_dir_resolution_falls_back_to_resume_parent(monkeypatch, tmp_path): def _fake_run_train_job(*, data, cfg, out_dir, **kwargs): captured["out_dir"] = out_dir - monkeypatch.setattr(cli, "run_train_job", _fake_run_train_job) + monkeypatch.setattr("giant.pipeline.run_train_job", _fake_run_train_job) resume_dir = tmp_path / "resumed_run" resume_dir.mkdir() @@ -178,7 +178,7 @@ def test_out_dir_resolution_defaults_when_neither_out_nor_resume_given(monkeypat def _fake_run_train_job(*, data, cfg, out_dir, **kwargs): captured["out_dir"] = out_dir - monkeypatch.setattr(cli, "run_train_job", _fake_run_train_job) + monkeypatch.setattr("giant.pipeline.run_train_job", _fake_run_train_job) monkeypatch.chdir(tmp_path) result = runner.invoke(cli.app, ["train", "dummy.parquet"]) @@ -192,7 +192,7 @@ def test_batch_size_auto_estimates_and_echoes(monkeypatch, tmp_path): def _fake_run_train_job(*, data, cfg, out_dir, num_workers, **kwargs): captured["batch_size"] = cfg["train"]["batch_size"] - monkeypatch.setattr(cli, "run_train_job", _fake_run_train_job) + monkeypatch.setattr("giant.pipeline.run_train_job", _fake_run_train_job) monkeypatch.setattr(cli.gconfig, "estimate_batch_size", lambda hidden_dim, n_blocks, device: 123) result = runner.invoke(