Compare commits
16 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 50d8368415 | |||
| 95d5fc6d89 | |||
| b8bd1ec982 | |||
| 70d018982b | |||
| 8d1c29efdd | |||
| 516a8a9ee1 | |||
| c12acfdade | |||
| e06d9e9581 | |||
| b0998a7d86 | |||
| cc11efb3ae | |||
| 7de3e92871 | |||
| f80fc90758 | |||
| 1cf16526c9 | |||
| 5c93457081 | |||
| 1ec333ff6d | |||
| bd255419e1 |
+2
-2
@@ -1,5 +1,5 @@
|
|||||||
[tool.bumpversion]
|
[tool.bumpversion]
|
||||||
current_version = "0.3.12"
|
current_version = "0.3.15"
|
||||||
parse = "(?P<major>\\d+)\\.(?P<minor>\\d+)\\.(?P<patch>\\d+)"
|
parse = "(?P<major>\\d+)\\.(?P<minor>\\d+)\\.(?P<patch>\\d+)"
|
||||||
serialize = ["{major}.{minor}.{patch}"]
|
serialize = ["{major}.{minor}.{patch}"]
|
||||||
search = "{current_version}"
|
search = "{current_version}"
|
||||||
@@ -8,7 +8,7 @@ regex = false
|
|||||||
allow_dirty = false
|
allow_dirty = false
|
||||||
commit = true
|
commit = true
|
||||||
tag = false
|
tag = false
|
||||||
message = "chore: bump version {current_version} -> {new_version} [skip ci]"
|
message = "chore: bump version {current_version} -> {new_version}"
|
||||||
pre_commit_hooks = ["uv lock", "git add uv.lock"]
|
pre_commit_hooks = ["uv lock", "git add uv.lock"]
|
||||||
|
|
||||||
[[tool.bumpversion.files]]
|
[[tool.bumpversion.files]]
|
||||||
|
|||||||
+34
-5
@@ -2,10 +2,9 @@ name: CI
|
|||||||
|
|
||||||
"on":
|
"on":
|
||||||
push:
|
push:
|
||||||
branches: ["**"]
|
branches: ["master"]
|
||||||
tags: ["**"]
|
tags: ["**"]
|
||||||
pull_request:
|
pull_request: {}
|
||||||
branches: [master]
|
|
||||||
|
|
||||||
env:
|
env:
|
||||||
UV_CACHE_DIR: /uv-cache
|
UV_CACHE_DIR: /uv-cache
|
||||||
@@ -156,7 +155,7 @@ jobs:
|
|||||||
uv run git-cliff --tag "$TAG" --unreleased --prepend CHANGELOG.md
|
uv run git-cliff --tag "$TAG" --unreleased --prepend CHANGELOG.md
|
||||||
git add CHANGELOG.md
|
git add CHANGELOG.md
|
||||||
if ! git diff --cached --quiet -- CHANGELOG.md; then
|
if ! git diff --cached --quiet -- CHANGELOG.md; then
|
||||||
git commit -m "chore: update changelog for $TAG [skip ci]"
|
git commit -m "chore: update changelog for $TAG"
|
||||||
else
|
else
|
||||||
git restore --staged CHANGELOG.md
|
git restore --staged CHANGELOG.md
|
||||||
fi
|
fi
|
||||||
@@ -193,7 +192,7 @@ jobs:
|
|||||||
git config user.name "gitea-actions"
|
git config user.name "gitea-actions"
|
||||||
git config user.email "actions@git.larsbogner.de"
|
git config user.email "actions@git.larsbogner.de"
|
||||||
git add pyproject.toml uv.lock
|
git add pyproject.toml uv.lock
|
||||||
git commit -m "chore: sync project version to tag ${GITHUB_REF_NAME} [skip ci]"
|
git commit -m "chore: sync project version to tag ${GITHUB_REF_NAME}"
|
||||||
git push origin HEAD:master
|
git push origin HEAD:master
|
||||||
git push origin ":refs/tags/${GITHUB_REF_NAME}"
|
git push origin ":refs/tags/${GITHUB_REF_NAME}"
|
||||||
git tag -f "${GITHUB_REF_NAME}" HEAD
|
git tag -f "${GITHUB_REF_NAME}" HEAD
|
||||||
@@ -201,3 +200,33 @@ jobs:
|
|||||||
else
|
else
|
||||||
echo "Tag version matches project version ($CURRENT_VERSION)"
|
echo "Tag version matches project version ($CURRENT_VERSION)"
|
||||||
fi
|
fi
|
||||||
|
|
||||||
|
publish-package:
|
||||||
|
name: Publish package to Gitea package registry
|
||||||
|
needs: [ruff-check, ruff-format, type-check, test, sync-version-on-tag]
|
||||||
|
if: startsWith(github.ref, 'refs/tags/')
|
||||||
|
runs-on: ubuntu-latest
|
||||||
|
container:
|
||||||
|
image: docker.gitea.com/runner-images:ubuntu-latest
|
||||||
|
volumes:
|
||||||
|
- /srv/act-runner-cache/uv:/uv-cache
|
||||||
|
steps:
|
||||||
|
# Check out by tag name (not the triggering SHA) since sync-version-on-tag
|
||||||
|
# may have force-moved the tag to a version-corrected commit.
|
||||||
|
- uses: actions/checkout@v4
|
||||||
|
with:
|
||||||
|
ref: ${{ github.ref_name }}
|
||||||
|
- uses: astral-sh/setup-uv@v5
|
||||||
|
with:
|
||||||
|
enable-cache: false
|
||||||
|
- run: |
|
||||||
|
echo "UV_CACHE_DIR=/uv-cache" >> "$GITHUB_ENV"
|
||||||
|
echo "UV_LINK_MODE=copy" >> "$GITHUB_ENV"
|
||||||
|
- run: uv build
|
||||||
|
# CI_TOKEN needs write:package scope (in addition to write:repository,
|
||||||
|
# used elsewhere) for this upload to authenticate.
|
||||||
|
- run: |
|
||||||
|
uv publish \
|
||||||
|
--publish-url "https://git.larsbogner.de/api/packages/lars/pypi" \
|
||||||
|
--username gitea-actions \
|
||||||
|
--password "${{ secrets.CI_TOKEN }}"
|
||||||
|
|||||||
@@ -1,5 +1,25 @@
|
|||||||
# Changelog
|
# Changelog
|
||||||
|
|
||||||
|
## [0.3.15] - 2026-08-28
|
||||||
|
|
||||||
|
### Changed
|
||||||
|
|
||||||
|
- Perf: defer heavy imports in giant/dwarf CLIs until commands run
|
||||||
|
|
||||||
|
## [0.3.14] - 2026-08-28
|
||||||
|
|
||||||
|
### Changed
|
||||||
|
|
||||||
|
- Ci: give automated commits visible checks, scope CI triggers, publish releases
|
||||||
|
|
||||||
|
- Ci: fix pull_request trigger not registering
|
||||||
|
|
||||||
|
## [0.3.13] - 2026-08-28
|
||||||
|
|
||||||
|
### Added
|
||||||
|
|
||||||
|
- Add inference-time model_config overrides with a sampling-key allowlist [gitea #87](https://git.larsbogner.de/lars/giant/issues/87)
|
||||||
|
|
||||||
## [0.3.12] - 2026-08-28
|
## [0.3.12] - 2026-08-28
|
||||||
|
|
||||||
### Added
|
### Added
|
||||||
|
|||||||
@@ -113,6 +113,8 @@ Secondary energies are a **stick-breaking partition of the `e_sec` budget** from
|
|||||||
|
|
||||||
**v0.3.0 — Stage-2 autoregressive redesign (implemented, released; on `master` since 2026-08-13):** motivated by the 2026-08-03 WGAN rollout benchmark, which failed specifically at the secondary-species level (zero photon secondaries, ~4M hallucinated `-14` muon antineutrinos). Stage 2 became autoregressive in descending-energy order with teacher forcing, and the particle-type representation went back to **categorical** (`particle_type.target = "onehot"`), reversing the 2026-07-17 continuous `(log-mass, charge)` target. The config break (`[conditioning]`/`[stage1_model]`/`[stage2_model]`/`[train]` replacing the flat `train.mode` + `[model]`) makes per-stage generators, stage-2-only training, and one-shot-vs-autoregressive comparison all expressible, and the `network.py` refactor into composable parts (encoder × trunk × objective) also makes routed WGAN work for the first time.
|
**v0.3.0 — Stage-2 autoregressive redesign (implemented, released; on `master` since 2026-08-13):** motivated by the 2026-08-03 WGAN rollout benchmark, which failed specifically at the secondary-species level (zero photon secondaries, ~4M hallucinated `-14` muon antineutrinos). Stage 2 became autoregressive in descending-energy order with teacher forcing, and the particle-type representation went back to **categorical** (`particle_type.target = "onehot"`), reversing the 2026-07-17 continuous `(log-mass, charge)` target. The config break (`[conditioning]`/`[stage1_model]`/`[stage2_model]`/`[train]` replacing the flat `train.mode` + `[model]`) makes per-stage generators, stage-2-only training, and one-shot-vs-autoregressive comparison all expressible, and the `network.py` refactor into composable parts (encoder × trunk × objective) also makes routed WGAN work for the first time.
|
||||||
|
|
||||||
|
**Baseline benchmark (done, 2026-08-26):** `configs/baseline.toml`'s first full rollout-vs-Geant4 validation (`analysis_341dfb14`, checkpoint `20260814_1743_s2-flow_h512_s2h512_bs36864_ep50/best.pt`, epoch 50/50). Confirms the v0.3.0 pivot fixed the species collapse — zero photon secondaries / hallucinated `-14` muon antineutrinos are both gone (γ at 95% of truth, no `-14` in the top species) — and rules out `conditioning.*.type = "physical"` as the cause, since this checkpoint pairs it with `flow`/no-router and still doesn't collapse. Bulk shower observables are close to Geant4 (total deposited energy +1.9%, containment depth-90%/95% both 0.986×), but steps/event now *over*-shoots by 1.32× (the opposite sign from every pre-v0.3.0 checkpoint), no hadronic/nuclear secondaries are produced at all, and event-to-event energy variance is ~16× too narrow. Writeup: `/home/lars/knowledge-base/experiments/giant-baseline-flow-ar-rollout-validation.md`.
|
||||||
|
|
||||||
v0.2 configs and checkpoints are auto-migrated (`config.migrate_config`, `model._legacy._migrate_legacy_model_config`, both drawing on shared facts in `giant/_migration.py`). **v0.2 checkpoint-loading support has no expiry decided yet**: `/ceph` still holds pre-v0.3.0 checkpoints and analysis runs referencing them, so don't delete or substantially alter either migration function or `tests/legacy/network_v02_snapshot.py` (the frozen v0.2 snapshot they're tested against) without an explicit decision to do so first.
|
v0.2 configs and checkpoints are auto-migrated (`config.migrate_config`, `model._legacy._migrate_legacy_model_config`, both drawing on shared facts in `giant/_migration.py`). **v0.2 checkpoint-loading support has no expiry decided yet**: `/ceph` still holds pre-v0.3.0 checkpoints and analysis runs referencing them, so don't delete or substantially alter either migration function or `tests/legacy/network_v02_snapshot.py` (the frozen v0.2 snapshot they're tested against) without an explicit decision to do so first.
|
||||||
|
|
||||||
**Faster-eval architectures — both implemented, neither validated.** Target is a ~10× native-Geant4 eval budget; no eval-latency number exists for any configuration yet, so that budget is unverified across the board.
|
**Faster-eval architectures — both implemented, neither validated.** Target is a ~10× native-Geant4 eval budget; no eval-latency number exists for any configuration yet, so that budget is unverified across the board.
|
||||||
|
|||||||
+1
-1
@@ -37,7 +37,7 @@ commit_preprocessors = [
|
|||||||
protect_breaking_commits = false
|
protect_breaking_commits = false
|
||||||
commit_parsers = [
|
commit_parsers = [
|
||||||
{ message = "^Merge ", skip = true },
|
{ message = "^Merge ", skip = true },
|
||||||
{ message = "\\[skip ci\\]", skip = true },
|
{ message = "^chore: (bump version|update changelog|sync project version)", skip = true },
|
||||||
{ message = "^Add", group = "<!-- 0 -->Added" },
|
{ message = "^Add", group = "<!-- 0 -->Added" },
|
||||||
{ message = "^(Fix|Clamp|Clip)", group = "<!-- 1 -->Fixed" },
|
{ message = "^(Fix|Clamp|Clip)", group = "<!-- 1 -->Fixed" },
|
||||||
{ message = "^(Remove|Drop|Deprecate)", group = "<!-- 2 -->Removed" },
|
{ message = "^(Remove|Drop|Deprecate)", group = "<!-- 2 -->Removed" },
|
||||||
|
|||||||
+15
-5
@@ -29,11 +29,21 @@
|
|||||||
# capacity overfitting is not the binding constraint, and every recent
|
# capacity overfitting is not the binding constraint, and every recent
|
||||||
# run used 0.0.
|
# run used 0.0.
|
||||||
#
|
#
|
||||||
# Known weak spots this baseline is expected to *exhibit* (they are the
|
# Known weak spots, now measured against this exact config rather than
|
||||||
# reason for the comparisons, not a reason to retune this file): every model
|
# extrapolated from the pre-v0.3 field (analysis_341dfb14, best.pt @ epoch
|
||||||
# on record under-produces steps per event by ~2x (rollout ~7e4 vs Geant4
|
# 50/50, full writeup: knowledge-base/experiments/
|
||||||
# ~1.4e5) and secondaries per event by 2-3.5x (~2-3e4 vs 7.2e4), and n_sec
|
# giant-baseline-flow-ar-rollout-validation.md). Unlike every pre-v0.3
|
||||||
# head accuracy sits at 0.863-0.867 regardless of size or objective.
|
# checkpoint (which under-produced steps/event by 1.6-5x), this baseline
|
||||||
|
# OVER-produces steps/event by 1.32x (1.86e5 vs Geant4 1.41e5) and
|
||||||
|
# under-produces secondaries/event by 0.84x (5.97e4 vs 7.14e4) — the sign on
|
||||||
|
# steps flipped with the v0.3 autoregressive pivot, so don't assume it still
|
||||||
|
# undershoots. Secondary-species hallucination (zero photons, hallucinated
|
||||||
|
# `-14` muon antineutrinos) that broke every prior checkpoint is gone; the
|
||||||
|
# remaining species gap is a total absence of hadronic/nuclear secondaries
|
||||||
|
# (protons, neutrons, ion recoils), not miscalibration of the ones produced.
|
||||||
|
# Total deposited energy/event is +1.9% high but its event-to-event spread is
|
||||||
|
# ~16x too narrow (31 MeV vs Geant4's 491 MeV). Per-step deposited energy is
|
||||||
|
# the worst per-step marginal (KS 0.179 vs 0.004-0.071 for the others).
|
||||||
|
|
||||||
[meta]
|
[meta]
|
||||||
# REQUIRED. Without it config.migrate_config reads this file as v0.2 and
|
# REQUIRED. Without it config.migrate_config reads this file as v0.2 and
|
||||||
|
|||||||
+52
-3
@@ -14,11 +14,18 @@ directly and imported from non-CLI code (`giant.analysis.router_gating`,
|
|||||||
lazily — see that module's docstring for why). Failures raise
|
lazily — see that module's docstring for why). Failures raise
|
||||||
`CheckpointCompatibilityError` with the same wording the CLI has always
|
`CheckpointCompatibilityError` with the same wording the CLI has always
|
||||||
shown; the CLI layer catches it and does the `typer.echo`/`Exit(1)`.
|
shown; the CLI layer catches it and does the `typer.echo`/`Exit(1)`.
|
||||||
|
|
||||||
|
`load_for_inference`'s `config_overrides` (gitea #87) lets a caller change a
|
||||||
|
checkpoint's `model_config` at load time, restricted to
|
||||||
|
`giant.config.INFERENCE_OVERRIDES` — the allowlist of keys that only affect
|
||||||
|
sampling, never module construction/shapes or the preprocessing normalizers/
|
||||||
|
vocab maps were fit under.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from dataclasses import dataclass
|
import copy
|
||||||
|
from dataclasses import dataclass, field
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
@@ -29,13 +36,45 @@ from giant.constants import K_MAX
|
|||||||
from giant.data.loader import TopNMap
|
from giant.data.loader import TopNMap
|
||||||
from giant.data.setup_cache import topnmap_from_json
|
from giant.data.setup_cache import topnmap_from_json
|
||||||
from giant.data.transforms import Normalizer
|
from giant.data.transforms import Normalizer
|
||||||
from giant.model.network import build_models
|
from giant.model.network import _migrate_legacy_model_config, build_models
|
||||||
|
|
||||||
|
|
||||||
class CheckpointCompatibilityError(Exception):
|
class CheckpointCompatibilityError(Exception):
|
||||||
"""Checkpoint is missing something `load_for_inference` needs."""
|
"""Checkpoint is missing something `load_for_inference` needs."""
|
||||||
|
|
||||||
|
|
||||||
|
def apply_config_overrides(model_cfg: dict, overrides: dict[str, object] | None) -> dict:
|
||||||
|
"""Deep-merge dotted-path *overrides* into a checkpoint's `model_config`,
|
||||||
|
validated against `giant.config.INFERENCE_OVERRIDES` — the allowlist of
|
||||||
|
keys that only affect sampling, not module construction/shapes or the
|
||||||
|
preprocessing normalizers/vocab maps were fit under (gitea #87).
|
||||||
|
|
||||||
|
Migrates a v0.2 flat `model_config` to the nested v0.3 shape first: a
|
||||||
|
dotted path like "stage1_model.ddpm.n_steps" would otherwise silently
|
||||||
|
write into a dict that `build_models` still reads as flat (it decides
|
||||||
|
v0.2-vs-v0.3 by `"stage1_model" in model_config`), suppressing migration.
|
||||||
|
|
||||||
|
Raises `CheckpointCompatibilityError` — never a bare `ValueError` or a
|
||||||
|
downstream `load_state_dict` size mismatch — for an unknown/disallowed
|
||||||
|
path or a value that fails its allowlisted check.
|
||||||
|
"""
|
||||||
|
if not overrides:
|
||||||
|
return model_cfg
|
||||||
|
cfg = model_cfg if "stage1_model" in model_cfg else _migrate_legacy_model_config(model_cfg)
|
||||||
|
cfg = copy.deepcopy(cfg)
|
||||||
|
for path, value in overrides.items():
|
||||||
|
spec = gconfig.INFERENCE_OVERRIDES.get(path)
|
||||||
|
if spec is None:
|
||||||
|
allowed = ", ".join(sorted(gconfig.INFERENCE_OVERRIDES))
|
||||||
|
raise CheckpointCompatibilityError(f"{path!r} is not an inference-safe override — allowed paths: {allowed}")
|
||||||
|
try:
|
||||||
|
spec.check(path, value)
|
||||||
|
except ValueError as exc:
|
||||||
|
raise CheckpointCompatibilityError(str(exc)) from exc
|
||||||
|
gconfig._set_path(cfg, path, value)
|
||||||
|
return cfg
|
||||||
|
|
||||||
|
|
||||||
def conditioning_axes(model_cfg: dict, default: str = "embedding") -> tuple[str, str]:
|
def conditioning_axes(model_cfg: dict, default: str = "embedding") -> tuple[str, str]:
|
||||||
"""(particle_conditioning, material_conditioning) for
|
"""(particle_conditioning, material_conditioning) for
|
||||||
`giant.data.transforms.build_cond_features`/`build_features` — from
|
`giant.data.transforms.build_cond_features`/`build_features` — from
|
||||||
@@ -128,6 +167,7 @@ class InferenceContext:
|
|||||||
model_config: dict
|
model_config: dict
|
||||||
epoch: int | None
|
epoch: int | None
|
||||||
best_val_loss: float | None
|
best_val_loss: float | None
|
||||||
|
config_overrides: dict[str, object] = field(default_factory=dict)
|
||||||
|
|
||||||
|
|
||||||
def load_for_inference(
|
def load_for_inference(
|
||||||
@@ -136,6 +176,7 @@ def load_for_inference(
|
|||||||
command_name: str,
|
command_name: str,
|
||||||
weights: str = "raw",
|
weights: str = "raw",
|
||||||
require_stage2: bool = True,
|
require_stage2: bool = True,
|
||||||
|
config_overrides: dict[str, object] | None = None,
|
||||||
) -> InferenceContext:
|
) -> InferenceContext:
|
||||||
"""Load *checkpoint* and reconstruct everything `predict`/`rollout` need
|
"""Load *checkpoint* and reconstruct everything `predict`/`rollout` need
|
||||||
to run it forward, on *device*, in `eval()` mode.
|
to run it forward, on *device*, in `eval()` mode.
|
||||||
@@ -148,6 +189,13 @@ def load_for_inference(
|
|||||||
both stages) or an acceptable `stage2 = None` result — kept as a real
|
both stages) or an acceptable `stage2 = None` result — kept as a real
|
||||||
parameter since `stage{1,2}_model.active` is a real, if currently
|
parameter since `stage{1,2}_model.active` is a real, if currently
|
||||||
stage1+stage2-only-in-practice, config option.
|
stage1+stage2-only-in-practice, config option.
|
||||||
|
|
||||||
|
*config_overrides* deep-merges dotted `model_config` paths (e.g.
|
||||||
|
`{"stage2_model.n_sec.sampling": "sample"}`) before anything is
|
||||||
|
derived from `model_config` or built — see `apply_config_overrides` for
|
||||||
|
the allowlist and validation. Every derived `InferenceContext` field
|
||||||
|
(`other_policy`, `stage{1,2}_ddpm_steps`, the built modules, ...)
|
||||||
|
reflects the overridden config.
|
||||||
"""
|
"""
|
||||||
ckpt = torch.load(checkpoint, map_location="cpu", weights_only=False)
|
ckpt = torch.load(checkpoint, map_location="cpu", weights_only=False)
|
||||||
for key in ("model_config", "sec_decoder"):
|
for key in ("model_config", "sec_decoder"):
|
||||||
@@ -159,7 +207,7 @@ def load_for_inference(
|
|||||||
|
|
||||||
gconfig.warn_if_checkpoint_config_mismatch(checkpoint)
|
gconfig.warn_if_checkpoint_config_mismatch(checkpoint)
|
||||||
|
|
||||||
model_cfg = ckpt["model_config"]
|
model_cfg = apply_config_overrides(ckpt["model_config"], config_overrides)
|
||||||
particle_conditioning, material_conditioning = conditioning_axes(model_cfg)
|
particle_conditioning, material_conditioning = conditioning_axes(model_cfg)
|
||||||
pdg_topn_map = load_pdg_topn_map(ckpt)
|
pdg_topn_map = load_pdg_topn_map(ckpt)
|
||||||
mat_topn_map = load_mat_topn_map(ckpt)
|
mat_topn_map = load_mat_topn_map(ckpt)
|
||||||
@@ -231,4 +279,5 @@ def load_for_inference(
|
|||||||
model_config=model_cfg,
|
model_config=model_cfg,
|
||||||
epoch=ckpt.get("epoch"),
|
epoch=ckpt.get("epoch"),
|
||||||
best_val_loss=ckpt.get("best_val_loss"),
|
best_val_loss=ckpt.get("best_val_loss"),
|
||||||
|
config_overrides=dict(config_overrides) if config_overrides else {},
|
||||||
)
|
)
|
||||||
|
|||||||
+93
-32
@@ -1,21 +1,19 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
from collections import Counter
|
from collections import Counter
|
||||||
from datetime import datetime, timezone
|
from datetime import datetime, timezone
|
||||||
from enum import Enum
|
from enum import Enum
|
||||||
import math
|
import math
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
import re
|
import re
|
||||||
from typing import Optional
|
from typing import TYPE_CHECKING, Optional
|
||||||
import uuid as uuid_mod
|
import uuid as uuid_mod
|
||||||
|
|
||||||
import numpy as np
|
|
||||||
import yaml
|
|
||||||
import torch
|
|
||||||
import typer
|
import typer
|
||||||
from typing_extensions import Annotated
|
from typing_extensions import Annotated
|
||||||
|
|
||||||
import pyarrow as pa
|
if TYPE_CHECKING:
|
||||||
import pyarrow.parquet as pq
|
import numpy as np
|
||||||
from tqdm import tqdm
|
|
||||||
|
|
||||||
from giant import config as gconfig
|
from giant import config as gconfig
|
||||||
from giant.constants import (
|
from giant.constants import (
|
||||||
@@ -25,30 +23,11 @@ from giant.constants import (
|
|||||||
PREDICT_SCHEMA_VERSION_KEY,
|
PREDICT_SCHEMA_VERSION_KEY,
|
||||||
ROLLOUT_COORD_VALUE,
|
ROLLOUT_COORD_VALUE,
|
||||||
)
|
)
|
||||||
from giant.data.loader import (
|
|
||||||
event_id_offset,
|
# giant.materials only pulls in numpy (no torch/pandas), and MATERIAL_PROPERTIES
|
||||||
find_parquet_files,
|
# is needed at decoration time below (a Typer option default), so it can't be
|
||||||
iter_file_chunks,
|
# deferred into a command body like the rest of this module's heavy imports.
|
||||||
iter_cond_chunks,
|
|
||||||
)
|
|
||||||
from giant.data.transforms import (
|
|
||||||
build_features,
|
|
||||||
build_cond_features,
|
|
||||||
energy_simplex_decode,
|
|
||||||
inv_local_frame_rotation,
|
|
||||||
inv_log_transform,
|
|
||||||
reconstruct_post_pos,
|
|
||||||
)
|
|
||||||
from giant.checkpoint_io import CheckpointCompatibilityError, load_for_inference
|
|
||||||
from giant.geometry import GeometryOracle
|
|
||||||
from giant.materials import MATERIAL_PROPERTIES
|
from giant.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)
|
app = typer.Typer(no_args_is_help=True)
|
||||||
|
|
||||||
@@ -141,6 +120,23 @@ 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]:
|
||||||
|
"""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
|
||||||
|
would arrive. Validation against the inference-safe allowlist happens
|
||||||
|
downstream in `giant.checkpoint_io.apply_config_overrides` — this only
|
||||||
|
parses syntax.
|
||||||
|
"""
|
||||||
|
out: dict[str, object] = {}
|
||||||
|
for spec in specs or []:
|
||||||
|
path, sep, val = spec.partition("=")
|
||||||
|
if not sep:
|
||||||
|
typer.echo(f"error: --set {spec!r} must be 'dotted.path=value'", err=True)
|
||||||
|
raise typer.Exit(1)
|
||||||
|
out[path] = _coerce_scalar(val)
|
||||||
|
return out
|
||||||
|
|
||||||
|
|
||||||
def _router_cli_overrides(
|
def _router_cli_overrides(
|
||||||
router: bool | None,
|
router: bool | None,
|
||||||
router_type: str | None,
|
router_type: str | None,
|
||||||
@@ -193,6 +189,8 @@ def _write_prediction_ref(
|
|||||||
comment: str | None = None,
|
comment: str | None = None,
|
||||||
) -> Path:
|
) -> Path:
|
||||||
"""Write a YAML sidecar in the checkpoint directory and return its path."""
|
"""Write a YAML sidecar in the checkpoint directory and return its path."""
|
||||||
|
import yaml
|
||||||
|
|
||||||
ref = {
|
ref = {
|
||||||
"prediction_id": pred_uuid,
|
"prediction_id": pred_uuid,
|
||||||
"output": str(out),
|
"output": str(out),
|
||||||
@@ -622,6 +620,10 @@ def train(
|
|||||||
] = None,
|
] = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Train the GIANT surrogate model."""
|
"""Train the GIANT surrogate model."""
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from giant.pipeline import run_train_job
|
||||||
|
|
||||||
batch_size_auto = False
|
batch_size_auto = False
|
||||||
batch_size_value: Optional[int] = None
|
batch_size_value: Optional[int] = None
|
||||||
if batch_size is not None:
|
if batch_size is not None:
|
||||||
@@ -1020,8 +1022,36 @@ def predict(
|
|||||||
help="Free-text note recorded in the prediction's YAML sidecar",
|
help="Free-text note recorded in the prediction's YAML sidecar",
|
||||||
),
|
),
|
||||||
] = None,
|
] = None,
|
||||||
|
set_: Annotated[
|
||||||
|
Optional[list[str]],
|
||||||
|
typer.Option(
|
||||||
|
"--set",
|
||||||
|
help="Override a sampling-only model_config key on this checkpoint, "
|
||||||
|
"'dotted.path=value' (repeatable) — see giant.config.INFERENCE_OVERRIDES "
|
||||||
|
"for the allowlist, e.g. --set stage2_model.n_sec.sampling=sample",
|
||||||
|
),
|
||||||
|
] = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Run trained model on a parquet file and save predictions."""
|
"""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_auto = False
|
||||||
batch_size_value: Optional[int] = None
|
batch_size_value: Optional[int] = None
|
||||||
if batch_size.strip().lower() == "auto":
|
if batch_size.strip().lower() == "auto":
|
||||||
@@ -1040,8 +1070,11 @@ def predict(
|
|||||||
typer.echo(f"device: {_device}")
|
typer.echo(f"device: {_device}")
|
||||||
|
|
||||||
# --- Load checkpoint ---
|
# --- Load checkpoint ---
|
||||||
|
config_overrides = _parse_set_flags(set_)
|
||||||
try:
|
try:
|
||||||
ctx = load_for_inference(checkpoint, _device, "predict", weights=weights.value)
|
ctx = load_for_inference(
|
||||||
|
checkpoint, _device, "predict", weights=weights.value, config_overrides=config_overrides
|
||||||
|
)
|
||||||
except CheckpointCompatibilityError as exc:
|
except CheckpointCompatibilityError as exc:
|
||||||
typer.echo(f"error: {exc}", err=True)
|
typer.echo(f"error: {exc}", err=True)
|
||||||
raise typer.Exit(1)
|
raise typer.Exit(1)
|
||||||
@@ -1314,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
|
the codebase's convention for the primary (a secondary always carries less
|
||||||
energy than its parent). See giant/analysis/reduce.py:entry_axis.
|
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_E: dict[int, float] = {}
|
||||||
best: dict[int, tuple] = {}
|
best: dict[int, tuple] = {}
|
||||||
for file_idx, path in enumerate(files):
|
for file_idx, path in enumerate(files):
|
||||||
@@ -1407,8 +1444,28 @@ def rollout(
|
|||||||
Optional[int],
|
Optional[int],
|
||||||
typer.Option("--seed", help="Torch/numpy seed for reproducibility"),
|
typer.Option("--seed", help="Torch/numpy seed for reproducibility"),
|
||||||
] = None,
|
] = None,
|
||||||
|
set_: Annotated[
|
||||||
|
Optional[list[str]],
|
||||||
|
typer.Option(
|
||||||
|
"--set",
|
||||||
|
help="Override a sampling-only model_config key on this checkpoint, "
|
||||||
|
"'dotted.path=value' (repeatable) — see giant.config.INFERENCE_OVERRIDES "
|
||||||
|
"for the allowlist, e.g. --set stage2_model.n_sec.sampling=sample",
|
||||||
|
),
|
||||||
|
] = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Roll the surrogate forward into full showers (autoregressive)."""
|
"""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:
|
if seed is not None:
|
||||||
torch.manual_seed(seed)
|
torch.manual_seed(seed)
|
||||||
np.random.seed(seed)
|
np.random.seed(seed)
|
||||||
@@ -1416,8 +1473,11 @@ def rollout(
|
|||||||
_device = torch.device(device) if device else gconfig.auto_device()
|
_device = torch.device(device) if device else gconfig.auto_device()
|
||||||
typer.echo(f"device: {_device}")
|
typer.echo(f"device: {_device}")
|
||||||
|
|
||||||
|
config_overrides = _parse_set_flags(set_)
|
||||||
try:
|
try:
|
||||||
ctx = load_for_inference(checkpoint, _device, "rollout", weights=weights.value)
|
ctx = load_for_inference(
|
||||||
|
checkpoint, _device, "rollout", weights=weights.value, config_overrides=config_overrides
|
||||||
|
)
|
||||||
except CheckpointCompatibilityError as exc:
|
except CheckpointCompatibilityError as exc:
|
||||||
typer.echo(f"error: {exc}", err=True)
|
typer.echo(f"error: {exc}", err=True)
|
||||||
raise typer.Exit(1)
|
raise typer.Exit(1)
|
||||||
@@ -1532,6 +1592,7 @@ def rollout(
|
|||||||
# model knob (router type/n_experts, noise_dim, vocab sizes, ...)
|
# model knob (router type/n_experts, noise_dim, vocab sizes, ...)
|
||||||
# is available downstream without touching this command again.
|
# is available downstream without touching this command again.
|
||||||
"model_config": dict(model_cfg),
|
"model_config": dict(model_cfg),
|
||||||
|
"config_overrides": dict(ctx.config_overrides),
|
||||||
"training_epoch": ctx.epoch,
|
"training_epoch": ctx.epoch,
|
||||||
"best_val_loss": ctx.best_val_loss,
|
"best_val_loss": ctx.best_val_loss,
|
||||||
# [train]/[meta] from the sibling config.toml (giant.config.save_config)
|
# [train]/[meta] from the sibling config.toml (giant.config.save_config)
|
||||||
|
|||||||
+94
-5
@@ -1,3 +1,5 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
import copy
|
import copy
|
||||||
import difflib
|
import difflib
|
||||||
import hashlib
|
import hashlib
|
||||||
@@ -10,12 +12,12 @@ from dataclasses import dataclass, field
|
|||||||
from datetime import datetime, timezone
|
from datetime import datetime, timezone
|
||||||
from enum import Enum
|
from enum import Enum
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
import numpy as np
|
|
||||||
import torch
|
|
||||||
|
|
||||||
from giant._migration import V02_FIXED_FACTS, V02_MODEL_KEY_TO_STAGES, reject_legacy_router_expert_sizing
|
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):
|
class Conditioning(str, Enum):
|
||||||
@@ -404,6 +406,16 @@ class Stage2RouterConfig(RouterConfig):
|
|||||||
return {"tie_to_stage1": self.tie_to_stage1, **super().to_dict()}
|
return {"tie_to_stage1": self.tie_to_stage1, **super().to_dict()}
|
||||||
|
|
||||||
|
|
||||||
|
# stage2_model.n_sec.sampling choices — single source of truth for both
|
||||||
|
# validate_config's train-time check and INFERENCE_OVERRIDES below.
|
||||||
|
STOP_SAMPLING_CHOICES = ("greedy", "sample")
|
||||||
|
|
||||||
|
# stage2_model.particle_type.other_policy choices — see ParticleTypeConfig's
|
||||||
|
# docstring for what each means; only documented there until now, since
|
||||||
|
# nothing validated it at train time.
|
||||||
|
OTHER_POLICY_CHOICES = ("sample", "modal", "drop")
|
||||||
|
|
||||||
|
|
||||||
@dataclass(frozen=True)
|
@dataclass(frozen=True)
|
||||||
class NSecConfig:
|
class NSecConfig:
|
||||||
# "head": a classifier over {0..k_max} on the condition encoding alone
|
# "head": a classifier over {0..k_max} on the condition encoding alone
|
||||||
@@ -931,6 +943,8 @@ def git_hash() -> str:
|
|||||||
|
|
||||||
|
|
||||||
def auto_device() -> torch.device:
|
def auto_device() -> torch.device:
|
||||||
|
import torch
|
||||||
|
|
||||||
if torch.cuda.is_available():
|
if torch.cuda.is_available():
|
||||||
return torch.device("cuda")
|
return torch.device("cuda")
|
||||||
if torch.backends.mps.is_available():
|
if torch.backends.mps.is_available():
|
||||||
@@ -976,6 +990,8 @@ def estimate_batch_size(
|
|||||||
inference (e.g. `predict`), which uses a much lower per-sample memory
|
inference (e.g. `predict`), which uses a much lower per-sample memory
|
||||||
calibration since there's no backward graph or optimizer state.
|
calibration since there's no backward graph or optimizer state.
|
||||||
"""
|
"""
|
||||||
|
import torch
|
||||||
|
|
||||||
if device.type != "cuda":
|
if device.type != "cuda":
|
||||||
raise ValueError(f"--batch-size auto is only supported on cuda devices, got {device.type!r}")
|
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()
|
device_index = device.index if device.index is not None else torch.cuda.current_device()
|
||||||
@@ -1118,6 +1134,72 @@ def _deep_merge(base: dict, override: dict) -> dict:
|
|||||||
return result
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class InferenceOverride:
|
||||||
|
"""One dotted `model_config` path that `giant.checkpoint_io.load_for_inference`
|
||||||
|
is allowed to change on an already-trained checkpoint, without retraining.
|
||||||
|
|
||||||
|
A path only belongs here if it affects neither module construction/tensor
|
||||||
|
shapes nor the data preprocessing the normalizers/vocab maps were fit
|
||||||
|
under — see the module docstring on `giant.model.summary` for the class
|
||||||
|
of key this targets (`_fingerprint`'s "plain scalar attribute" leaves),
|
||||||
|
and `giant.checkpoint_io.apply_config_overrides` for where this is used.
|
||||||
|
"""
|
||||||
|
|
||||||
|
why: str
|
||||||
|
choices: tuple[str, ...] | None = None
|
||||||
|
minimum: float | None = None
|
||||||
|
numeric: bool = False # int/float leaf (vs. str, the default)
|
||||||
|
|
||||||
|
def check(self, path: str, value: object) -> None:
|
||||||
|
if self.choices is not None:
|
||||||
|
if value not in self.choices:
|
||||||
|
raise ValueError(f"{path} = {value!r} — must be one of {self.choices}")
|
||||||
|
return
|
||||||
|
if self.numeric:
|
||||||
|
if isinstance(value, bool) or not isinstance(value, (int, float)):
|
||||||
|
raise ValueError(f"{path} = {value!r} — must be a number")
|
||||||
|
if self.minimum is not None and value < self.minimum:
|
||||||
|
raise ValueError(f"{path} = {value!r} — must be >= {self.minimum}")
|
||||||
|
|
||||||
|
|
||||||
|
# Inference-safe dotted `model_config` paths — the allowlist gitea #87 asked
|
||||||
|
# for, so a typo or a shape-bearing key (e.g. "stage1_model.hidden_dim")
|
||||||
|
# raises a clear CheckpointCompatibilityError instead of surfacing as an
|
||||||
|
# opaque load_state_dict size mismatch later. Extend this table, not a
|
||||||
|
# per-call bypass, when a new inference-only key needs the capability.
|
||||||
|
INFERENCE_OVERRIDES: dict[str, InferenceOverride] = {
|
||||||
|
"stage2_model.n_sec.sampling": InferenceOverride(
|
||||||
|
why="giant.sample's n_sec head/stop-token sampling reads this at sample time only (gitea #86)",
|
||||||
|
choices=STOP_SAMPLING_CHOICES,
|
||||||
|
),
|
||||||
|
"stage1_model.ddpm.n_steps": InferenceOverride(
|
||||||
|
why="giant.model.schedule.CosineSchedule's step count, resolved at sample time",
|
||||||
|
numeric=True,
|
||||||
|
minimum=1,
|
||||||
|
),
|
||||||
|
"stage2_model.ddpm.n_steps": InferenceOverride(
|
||||||
|
why="giant.model.schedule.CosineSchedule's step count, resolved at sample time",
|
||||||
|
numeric=True,
|
||||||
|
minimum=1,
|
||||||
|
),
|
||||||
|
"stage2_model.particle_type.other_policy": InferenceOverride(
|
||||||
|
why="giant.rollout resolves an 'other'-bucket secondary's PDG code with this at rollout time",
|
||||||
|
choices=OTHER_POLICY_CHOICES,
|
||||||
|
),
|
||||||
|
"stage1_model.router.temperature": InferenceOverride(
|
||||||
|
why="giant.model.routers.EnergyRouter.temperature, a plain constructor attribute",
|
||||||
|
numeric=True,
|
||||||
|
minimum=1e-6,
|
||||||
|
),
|
||||||
|
"stage2_model.router.temperature": InferenceOverride(
|
||||||
|
why="giant.model.routers.EnergyRouter.temperature, a plain constructor attribute",
|
||||||
|
numeric=True,
|
||||||
|
minimum=1e-6,
|
||||||
|
),
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
@dataclass(frozen=True)
|
@dataclass(frozen=True)
|
||||||
class FlagSpec:
|
class FlagSpec:
|
||||||
"""One CLI flag's mapping into the config-overrides tree.
|
"""One CLI flag's mapping into the config-overrides tree.
|
||||||
@@ -1548,7 +1630,7 @@ def validate_config(cfg: dict, *, resume: bool = False) -> None:
|
|||||||
)
|
)
|
||||||
|
|
||||||
n_sec_sampling = _get_path(cfg, "stage2_model.n_sec.sampling")
|
n_sec_sampling = _get_path(cfg, "stage2_model.n_sec.sampling")
|
||||||
if n_sec_sampling not in ("greedy", "sample"):
|
if n_sec_sampling not in STOP_SAMPLING_CHOICES:
|
||||||
raise ValueError(f"stage2_model.n_sec.sampling = {n_sec_sampling!r} — must be 'greedy' or 'sample'")
|
raise ValueError(f"stage2_model.n_sec.sampling = {n_sec_sampling!r} — must be 'greedy' or 'sample'")
|
||||||
|
|
||||||
precision = _get_path(cfg, "train.precision")
|
precision = _get_path(cfg, "train.precision")
|
||||||
@@ -1604,6 +1686,8 @@ def validate_config(cfg: dict, *, resume: bool = False) -> None:
|
|||||||
"'energy_desc' (the only implemented ordering; see "
|
"'energy_desc' (the only implemented ordering; see "
|
||||||
"AutoregressiveConfig.order's docstring)"
|
"AutoregressiveConfig.order's docstring)"
|
||||||
)
|
)
|
||||||
|
from giant.model.history import HISTORY_REGISTRY
|
||||||
|
|
||||||
history = _get_path(cfg, "stage2_model.autoregressive.history")
|
history = _get_path(cfg, "stage2_model.autoregressive.history")
|
||||||
if history not in HISTORY_REGISTRY:
|
if history not in HISTORY_REGISTRY:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
@@ -1805,6 +1889,9 @@ def resolve_default_out_dir(cfg: dict, base: Path = Path("checkpoints")) -> Path
|
|||||||
|
|
||||||
|
|
||||||
def seed_everything(seed: int) -> None:
|
def seed_everything(seed: int) -> None:
|
||||||
|
import numpy as np
|
||||||
|
import torch
|
||||||
|
|
||||||
random.seed(seed)
|
random.seed(seed)
|
||||||
np.random.seed(seed)
|
np.random.seed(seed)
|
||||||
torch.manual_seed(seed)
|
torch.manual_seed(seed)
|
||||||
@@ -1858,6 +1945,8 @@ def build_run_meta(
|
|||||||
n_val_events: int,
|
n_val_events: int,
|
||||||
n_train_steps: int,
|
n_train_steps: int,
|
||||||
) -> dict:
|
) -> dict:
|
||||||
|
import torch
|
||||||
|
|
||||||
return {
|
return {
|
||||||
"config_version": CONFIG_VERSION,
|
"config_version": CONFIG_VERSION,
|
||||||
"git_hash": git_hash(),
|
"git_hash": git_hash(),
|
||||||
|
|||||||
+11
-1
@@ -37,7 +37,7 @@ from dataclasses import dataclass, field
|
|||||||
|
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
|
|
||||||
from giant.config import _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
|
||||||
from giant.model.trunks import RoutedTrunk
|
from giant.model.trunks import RoutedTrunk
|
||||||
|
|
||||||
@@ -111,6 +111,7 @@ class ModelSummary:
|
|||||||
pdg_vocab: int
|
pdg_vocab: int
|
||||||
mat_vocab: int
|
mat_vocab: int
|
||||||
vocab_caveats: list[str] = field(default_factory=list)
|
vocab_caveats: list[str] = field(default_factory=list)
|
||||||
|
overridable: list[str] = field(default_factory=list)
|
||||||
|
|
||||||
|
|
||||||
def _build_model_config(cfg: dict, pdg_vocab: int, mat_vocab: int) -> dict:
|
def _build_model_config(cfg: dict, pdg_vocab: int, mat_vocab: int) -> dict:
|
||||||
@@ -234,6 +235,8 @@ def summarize_model(cfg: dict, pdg_vocab: int, mat_vocab: int) -> ModelSummary:
|
|||||||
else:
|
else:
|
||||||
inert.append(path)
|
inert.append(path)
|
||||||
|
|
||||||
|
overridable = sorted(p for p in in_scope if p in INFERENCE_OVERRIDES)
|
||||||
|
|
||||||
return ModelSummary(
|
return ModelSummary(
|
||||||
modules=modules,
|
modules=modules,
|
||||||
consumed=sorted(consumed),
|
consumed=sorted(consumed),
|
||||||
@@ -242,6 +245,7 @@ def summarize_model(cfg: dict, pdg_vocab: int, mat_vocab: int) -> ModelSummary:
|
|||||||
pdg_vocab=pdg_vocab,
|
pdg_vocab=pdg_vocab,
|
||||||
mat_vocab=mat_vocab,
|
mat_vocab=mat_vocab,
|
||||||
vocab_caveats=_vocab_caveats(cfg),
|
vocab_caveats=_vocab_caveats(cfg),
|
||||||
|
overridable=overridable,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -309,6 +313,12 @@ def render_summary(summary: ModelSummary) -> str:
|
|||||||
else:
|
else:
|
||||||
lines.append(" (none)")
|
lines.append(" (none)")
|
||||||
|
|
||||||
|
if summary.overridable:
|
||||||
|
lines.append("")
|
||||||
|
lines.append("inference-overridable without retraining (giant predict/rollout --set):")
|
||||||
|
for path in summary.overridable:
|
||||||
|
lines.append(f" {path} ({INFERENCE_OVERRIDES[path].why})")
|
||||||
|
|
||||||
if summary.vocab_caveats:
|
if summary.vocab_caveats:
|
||||||
lines.append("")
|
lines.append("")
|
||||||
lines.append("vocab placeholder caveats:")
|
lines.append("vocab placeholder caveats:")
|
||||||
|
|||||||
+34
-14
@@ -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.
|
(and `--help`) to remember instead of five differently-hyphenated ones.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
import os
|
import os
|
||||||
from enum import Enum
|
from enum import Enum
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
@@ -14,20 +16,15 @@ import typer
|
|||||||
from typing_extensions import Annotated
|
from typing_extensions import Annotated
|
||||||
|
|
||||||
from giant.config import Conditioning
|
from giant.config import Conditioning
|
||||||
from giant.tools.bump_dataset_version import (
|
|
||||||
run_bump_gen,
|
# DATA_DEFAULT/SCAN_DIR_DEFAULT are Typer option defaults (evaluated at
|
||||||
run_bump_schema,
|
# decoration time below), so that one name has to stay eager — the module
|
||||||
run_create_manifest,
|
# itself is stdlib-only, so it costs nothing. Every other giant.tools.*
|
||||||
run_status,
|
# import here is deferred into the one command body that uses it, since
|
||||||
run_update_manifest,
|
# several (steps_to_parquet: uproot/awkward/polars; warm_setup_cache:
|
||||||
)
|
# giant.pipeline -> torch; geometry_oracle: pandas) are expensive and
|
||||||
from giant.tools.create_root_files import run_make_root
|
# `dwarf --help`/tab-completion shouldn't pay for all of them upfront.
|
||||||
from giant.tools.geometry_oracle import run_build_geometry_oracle
|
from giant.tools.hparam_scan import DATA_DEFAULT, SCAN_DIR_DEFAULT
|
||||||
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
|
|
||||||
|
|
||||||
app = typer.Typer(no_args_is_help=True)
|
app = typer.Typer(no_args_is_help=True)
|
||||||
|
|
||||||
@@ -121,6 +118,9 @@ def convert(
|
|||||||
] = None,
|
] = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Convert ROOT Steps tree(s) to Parquet."""
|
"""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:
|
if jobs < 1:
|
||||||
typer.echo("error: --jobs must be >= 1", err=True)
|
typer.echo("error: --jobs must be >= 1", err=True)
|
||||||
raise typer.Exit(1)
|
raise typer.Exit(1)
|
||||||
@@ -183,6 +183,8 @@ def migrate(
|
|||||||
] = False,
|
] = False,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""One-time migration into the versioned raw/processed/pools/derived layout."""
|
"""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)
|
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,
|
root: Annotated[Path, typer.Option("--root", help="Dataset root")] = _DATASET_ROOT_DEFAULT,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Cut a new raw generation."""
|
"""Cut a new raw generation."""
|
||||||
|
from giant.tools.bump_dataset_version import run_bump_gen
|
||||||
|
|
||||||
run_bump_gen(
|
run_bump_gen(
|
||||||
kind=kind,
|
kind=kind,
|
||||||
reason=reason,
|
reason=reason,
|
||||||
@@ -240,6 +244,8 @@ def bump_schema(
|
|||||||
root: Annotated[Path, typer.Option("--root", help="Dataset root")] = _DATASET_ROOT_DEFAULT,
|
root: Annotated[Path, typer.Option("--root", help="Dataset root")] = _DATASET_ROOT_DEFAULT,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Cut a new schema within a gen."""
|
"""Cut a new schema within a gen."""
|
||||||
|
from giant.tools.bump_dataset_version import run_bump_schema
|
||||||
|
|
||||||
run_bump_schema(
|
run_bump_schema(
|
||||||
kind=kind,
|
kind=kind,
|
||||||
gen=gen,
|
gen=gen,
|
||||||
@@ -257,6 +263,8 @@ def status(
|
|||||||
root: Annotated[Path, typer.Option("--root", help="Dataset root")] = _DATASET_ROOT_DEFAULT,
|
root: Annotated[Path, typer.Option("--root", help="Dataset root")] = _DATASET_ROOT_DEFAULT,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""List existing gens/schemas per kind."""
|
"""List existing gens/schemas per kind."""
|
||||||
|
from giant.tools.bump_dataset_version import run_status
|
||||||
|
|
||||||
run_status(str(root))
|
run_status(str(root))
|
||||||
|
|
||||||
|
|
||||||
@@ -281,6 +289,8 @@ def update_manifest(
|
|||||||
] = False,
|
] = False,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Repoint manifest(s) to a new gen and/or schema, verifying all target files exist."""
|
"""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)
|
run_update_manifest([str(m) for m in manifests], schema=schema, execute=execute, gen=gen)
|
||||||
|
|
||||||
|
|
||||||
@@ -311,6 +321,8 @@ def create_manifest(
|
|||||||
] = False,
|
] = False,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Create a new manifest from a list of parquet files."""
|
"""Create a new manifest from a list of parquet files."""
|
||||||
|
from giant.tools.bump_dataset_version import run_create_manifest
|
||||||
|
|
||||||
run_create_manifest(
|
run_create_manifest(
|
||||||
[str(f) for f in files],
|
[str(f) for f in files],
|
||||||
execute=execute,
|
execute=execute,
|
||||||
@@ -358,6 +370,8 @@ def make_root(
|
|||||||
] = False,
|
] = False,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Generate new ROOT shards via a minicalosim executable."""
|
"""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")
|
_warn_if_exceeds_shared_quota(jobs, "--jobs")
|
||||||
run_make_root(
|
run_make_root(
|
||||||
executable=executable,
|
executable=executable,
|
||||||
@@ -423,6 +437,8 @@ def build_geometry_oracle(
|
|||||||
] = 2000,
|
] = 2000,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Fit a position -> (material, layer_id) oracle for `giant rollout`."""
|
"""Fit a position -> (material, layer_id) oracle for `giant rollout`."""
|
||||||
|
from giant.tools.geometry_oracle import run_build_geometry_oracle
|
||||||
|
|
||||||
run_build_geometry_oracle(
|
run_build_geometry_oracle(
|
||||||
data=data,
|
data=data,
|
||||||
out=out,
|
out=out,
|
||||||
@@ -515,6 +531,8 @@ def warm_cache(
|
|||||||
such entry across every run) skips straight to training. See
|
such entry across every run) skips straight to training. See
|
||||||
giant/data/setup_cache.py.
|
giant/data/setup_cache.py.
|
||||||
"""
|
"""
|
||||||
|
from giant.tools.warm_setup_cache import run_warm_setup_cache
|
||||||
|
|
||||||
flag_overrides = {
|
flag_overrides = {
|
||||||
"--val-fraction": val_fraction,
|
"--val-fraction": val_fraction,
|
||||||
"--seed": seed,
|
"--seed": seed,
|
||||||
@@ -557,6 +575,8 @@ def hparam_scan(
|
|||||||
dry_run: Annotated[bool, typer.Option("--dry-run")] = False,
|
dry_run: Annotated[bool, typer.Option("--dry-run")] = False,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Grid-scan dropout x n_blocks x hidden_dim via sequential `giant train` runs."""
|
"""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)
|
run_hparam_scan(data=data, scan_dir=scan_dir, seed=seed, dry_run=dry_run)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
+1
-1
@@ -1,6 +1,6 @@
|
|||||||
[project]
|
[project]
|
||||||
name = "giant"
|
name = "giant"
|
||||||
version = "0.3.12"
|
version = "0.3.15"
|
||||||
description = "Geant4 step-function surrogate via conditional flow matching"
|
description = "Geant4 step-function surrogate via conditional flow matching"
|
||||||
readme = "README.md"
|
readme = "README.md"
|
||||||
requires-python = ">=3.12"
|
requires-python = ">=3.12"
|
||||||
|
|||||||
@@ -14,6 +14,7 @@ from giant import config as gconfig
|
|||||||
from giant.checkpoint_io import (
|
from giant.checkpoint_io import (
|
||||||
CheckpointCompatibilityError,
|
CheckpointCompatibilityError,
|
||||||
InferenceContext,
|
InferenceContext,
|
||||||
|
apply_config_overrides,
|
||||||
conditioning_axes,
|
conditioning_axes,
|
||||||
load_for_inference,
|
load_for_inference,
|
||||||
stage_cfg,
|
stage_cfg,
|
||||||
@@ -278,3 +279,122 @@ def test_stage_cfg_new_shape_returns_subdict():
|
|||||||
def test_stage_cfg_v02_flat_shape_returns_empty_dict():
|
def test_stage_cfg_v02_flat_shape_returns_empty_dict():
|
||||||
model_cfg = {"hidden_dim": 32, "n_blocks": 4}
|
model_cfg = {"hidden_dim": 32, "n_blocks": 4}
|
||||||
assert stage_cfg(model_cfg, "stage2") == {}
|
assert stage_cfg(model_cfg, "stage2") == {}
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# config_overrides (gitea #87)
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def _router_model_cfg() -> dict:
|
||||||
|
cfg = _model_cfg()
|
||||||
|
cfg["stage1_model"]["router"] = {"enabled": True, "type": "energy", "n_experts": 2}
|
||||||
|
return cfg
|
||||||
|
|
||||||
|
|
||||||
|
def test_config_override_n_sec_sampling_changes_stage2_attribute(tmp_path):
|
||||||
|
checkpoint = _write_checkpoint(tmp_path)
|
||||||
|
ctx = load_for_inference(
|
||||||
|
checkpoint,
|
||||||
|
torch.device("cpu"),
|
||||||
|
"predict",
|
||||||
|
config_overrides={"stage2_model.n_sec.sampling": "sample"},
|
||||||
|
)
|
||||||
|
assert ctx.stage2 is not None
|
||||||
|
assert ctx.stage2.n_sec_sampling == "sample"
|
||||||
|
assert ctx.config_overrides == {"stage2_model.n_sec.sampling": "sample"}
|
||||||
|
|
||||||
|
|
||||||
|
def test_config_override_ddpm_n_steps_changes_context_fields(tmp_path):
|
||||||
|
checkpoint = _write_checkpoint(tmp_path)
|
||||||
|
ctx = load_for_inference(
|
||||||
|
checkpoint,
|
||||||
|
torch.device("cpu"),
|
||||||
|
"predict",
|
||||||
|
config_overrides={"stage1_model.ddpm.n_steps": 42, "stage2_model.ddpm.n_steps": 7},
|
||||||
|
)
|
||||||
|
assert ctx.stage1_ddpm_steps == 42
|
||||||
|
assert ctx.stage2_ddpm_steps == 7
|
||||||
|
|
||||||
|
|
||||||
|
def test_config_override_other_policy_changes_context_field(tmp_path):
|
||||||
|
checkpoint = _write_checkpoint(tmp_path)
|
||||||
|
ctx = load_for_inference(
|
||||||
|
checkpoint,
|
||||||
|
torch.device("cpu"),
|
||||||
|
"predict",
|
||||||
|
config_overrides={"stage2_model.particle_type.other_policy": "modal"},
|
||||||
|
)
|
||||||
|
assert ctx.other_policy == "modal"
|
||||||
|
|
||||||
|
|
||||||
|
def test_config_override_router_temperature_changes_router_attribute(tmp_path):
|
||||||
|
checkpoint = _write_checkpoint(tmp_path, model_cfg=_router_model_cfg())
|
||||||
|
ctx = load_for_inference(
|
||||||
|
checkpoint,
|
||||||
|
torch.device("cpu"),
|
||||||
|
"predict",
|
||||||
|
config_overrides={"stage1_model.router.temperature": 1.5},
|
||||||
|
)
|
||||||
|
assert ctx.stage1 is not None
|
||||||
|
assert ctx.stage1.trunk.router.temperature == pytest.approx(1.5)
|
||||||
|
|
||||||
|
|
||||||
|
def test_config_override_no_overrides_defaults_to_empty_dict(tmp_path):
|
||||||
|
checkpoint = _write_checkpoint(tmp_path)
|
||||||
|
ctx = load_for_inference(checkpoint, torch.device("cpu"), "predict")
|
||||||
|
assert ctx.config_overrides == {}
|
||||||
|
|
||||||
|
|
||||||
|
def test_config_override_unknown_path_raises(tmp_path):
|
||||||
|
checkpoint = _write_checkpoint(tmp_path)
|
||||||
|
with pytest.raises(CheckpointCompatibilityError, match="not an inference-safe override"):
|
||||||
|
load_for_inference(
|
||||||
|
checkpoint,
|
||||||
|
torch.device("cpu"),
|
||||||
|
"predict",
|
||||||
|
config_overrides={"stage2_model.n_sec.typo": "sample"},
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_config_override_shape_bearing_key_raises_up_front(tmp_path):
|
||||||
|
checkpoint = _write_checkpoint(tmp_path)
|
||||||
|
with pytest.raises(CheckpointCompatibilityError, match="not an inference-safe override"):
|
||||||
|
load_for_inference(
|
||||||
|
checkpoint,
|
||||||
|
torch.device("cpu"),
|
||||||
|
"predict",
|
||||||
|
config_overrides={"stage1_model.hidden_dim": 999},
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_config_override_bad_value_raises(tmp_path):
|
||||||
|
checkpoint = _write_checkpoint(tmp_path)
|
||||||
|
with pytest.raises(CheckpointCompatibilityError, match="must be one of"):
|
||||||
|
load_for_inference(
|
||||||
|
checkpoint,
|
||||||
|
torch.device("cpu"),
|
||||||
|
"predict",
|
||||||
|
config_overrides={"stage2_model.n_sec.sampling": "maybe"},
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_apply_config_overrides_no_overrides_returns_same_object():
|
||||||
|
cfg = _model_cfg()
|
||||||
|
assert apply_config_overrides(cfg, None) is cfg
|
||||||
|
assert apply_config_overrides(cfg, {}) is cfg
|
||||||
|
|
||||||
|
|
||||||
|
def test_apply_config_overrides_migrates_legacy_flat_model_config_first():
|
||||||
|
legacy_cfg = {
|
||||||
|
"pdg_vocab": len(PDG_MAP),
|
||||||
|
"mat_vocab": len(MAT_MAP),
|
||||||
|
"hidden_dim": 32,
|
||||||
|
"n_blocks": 4,
|
||||||
|
"emb_dim": 8,
|
||||||
|
"dropout": 0.1,
|
||||||
|
"k_max": 5,
|
||||||
|
}
|
||||||
|
merged = apply_config_overrides(legacy_cfg, {"stage1_model.ddpm.n_steps": 10})
|
||||||
|
assert merged["stage1_model"]["ddpm"]["n_steps"] == 10
|
||||||
|
assert merged["stage1_model"]["hidden_dim"] == 32
|
||||||
|
|||||||
@@ -172,3 +172,43 @@ def test_predict_exits_1_on_checkpoint_missing_model_config(tmp_path):
|
|||||||
|
|
||||||
assert result.exit_code == 1
|
assert result.exit_code == 1
|
||||||
assert "checkpoint has no model_config" in result.output
|
assert "checkpoint has no model_config" in result.output
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# --set (gitea #87)
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def test_predict_set_flag_without_equals_exits_1(tmp_path):
|
||||||
|
checkpoint = tmp_path / "missing.pt"
|
||||||
|
|
||||||
|
result = runner.invoke(
|
||||||
|
app,
|
||||||
|
["predict", "dummy.parquet", "--checkpoint", str(checkpoint), "--set", "sampling"],
|
||||||
|
)
|
||||||
|
|
||||||
|
assert result.exit_code == 1
|
||||||
|
assert "must be 'dotted.path=value'" in result.output
|
||||||
|
|
||||||
|
|
||||||
|
def test_predict_set_flag_disallowed_path_surfaces_compat_error(tmp_path):
|
||||||
|
checkpoint = tmp_path / "ckpt.pt"
|
||||||
|
torch.save(
|
||||||
|
{"model_config": {"stage1_model": {}, "stage2_model": {}}, "sec_decoder": {}, "normalizer": {"sec_phys": {}}},
|
||||||
|
checkpoint,
|
||||||
|
)
|
||||||
|
|
||||||
|
result = runner.invoke(
|
||||||
|
app,
|
||||||
|
[
|
||||||
|
"predict",
|
||||||
|
"dummy.parquet",
|
||||||
|
"--checkpoint",
|
||||||
|
str(checkpoint),
|
||||||
|
"--set",
|
||||||
|
"stage1_model.hidden_dim=999",
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
assert result.exit_code == 1
|
||||||
|
assert "not an inference-safe override" in result.output
|
||||||
|
|||||||
@@ -31,3 +31,28 @@ def test_rollout_exits_1_on_checkpoint_missing_model_config(tmp_path):
|
|||||||
|
|
||||||
assert result.exit_code == 1
|
assert result.exit_code == 1
|
||||||
assert "checkpoint has no model_config" in result.output
|
assert "checkpoint has no model_config" in result.output
|
||||||
|
|
||||||
|
|
||||||
|
def test_rollout_set_flag_disallowed_path_surfaces_compat_error(tmp_path):
|
||||||
|
checkpoint = tmp_path / "ckpt.pt"
|
||||||
|
torch.save(
|
||||||
|
{"model_config": {"stage1_model": {}, "stage2_model": {}}, "sec_decoder": {}, "normalizer": {"sec_phys": {}}},
|
||||||
|
checkpoint,
|
||||||
|
)
|
||||||
|
|
||||||
|
result = runner.invoke(
|
||||||
|
app,
|
||||||
|
[
|
||||||
|
"rollout",
|
||||||
|
"dummy.parquet",
|
||||||
|
"--checkpoint",
|
||||||
|
str(checkpoint),
|
||||||
|
"--geometry",
|
||||||
|
"dummy_geometry.pkl",
|
||||||
|
"--set",
|
||||||
|
"stage2_model.n_sec.typo=sample",
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
assert result.exit_code == 1
|
||||||
|
assert "not an inference-safe override" in result.output
|
||||||
|
|||||||
@@ -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):
|
def _fake_run_train_job(*, data, cfg, out_dir, **kwargs):
|
||||||
captured["cfg"] = cfg
|
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(
|
result = runner.invoke(
|
||||||
cli.app,
|
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):
|
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(
|
result = runner.invoke(
|
||||||
cli.app,
|
cli.app,
|
||||||
["train", "dummy.parquet", "--out", str(tmp_path / "run"), "--batch-size", "not-a-number"],
|
["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):
|
def _fake_run_train_job(*, data, cfg, out_dir, **kwargs):
|
||||||
captured["out_dir"] = out_dir
|
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 = tmp_path / "resumed_run"
|
||||||
resume_dir.mkdir()
|
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):
|
def _fake_run_train_job(*, data, cfg, out_dir, **kwargs):
|
||||||
captured["out_dir"] = out_dir
|
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 = tmp_path / "resumed_run"
|
||||||
resume_dir.mkdir()
|
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):
|
def _fake_run_train_job(*, data, cfg, out_dir, **kwargs):
|
||||||
captured["out_dir"] = out_dir
|
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)
|
monkeypatch.chdir(tmp_path)
|
||||||
|
|
||||||
result = runner.invoke(cli.app, ["train", "dummy.parquet"])
|
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):
|
def _fake_run_train_job(*, data, cfg, out_dir, num_workers, **kwargs):
|
||||||
captured["batch_size"] = cfg["train"]["batch_size"]
|
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)
|
monkeypatch.setattr(cli.gconfig, "estimate_batch_size", lambda hidden_dim, n_blocks, device: 123)
|
||||||
|
|
||||||
result = runner.invoke(
|
result = runner.invoke(
|
||||||
|
|||||||
@@ -12,6 +12,8 @@ from giant.cli import app
|
|||||||
from giant.materials import MATERIAL_PROPERTIES
|
from giant.materials import MATERIAL_PROPERTIES
|
||||||
from giant.model.summary import _NOT_BUILD_TIME, _built_modules, _vocab_caveats, summarize_model
|
from giant.model.summary import _NOT_BUILD_TIME, _built_modules, _vocab_caveats, summarize_model
|
||||||
|
|
||||||
|
INFERENCE_OVERRIDES = gconfig.INFERENCE_OVERRIDES
|
||||||
|
|
||||||
runner = CliRunner()
|
runner = CliRunner()
|
||||||
|
|
||||||
_PDG_VOCAB = 300
|
_PDG_VOCAB = 300
|
||||||
@@ -53,6 +55,16 @@ def test_not_build_time_allow_list_has_no_stale_entries():
|
|||||||
assert not stale, f"_NOT_BUILD_TIME entries no longer in DEFAULT_CONFIG: {sorted(stale)}"
|
assert not stale, f"_NOT_BUILD_TIME entries no longer in DEFAULT_CONFIG: {sorted(stale)}"
|
||||||
|
|
||||||
|
|
||||||
|
def test_inference_overrides_allow_list_has_no_stale_entries():
|
||||||
|
in_scope = set(gconfig.leaf_paths(gconfig.DEFAULT_CONFIG))
|
||||||
|
stale = set(INFERENCE_OVERRIDES) - in_scope
|
||||||
|
assert not stale, f"INFERENCE_OVERRIDES entries no longer in DEFAULT_CONFIG: {sorted(stale)}"
|
||||||
|
|
||||||
|
|
||||||
|
def test_default_config_overridable_lists_every_allowlisted_path(default_summary):
|
||||||
|
assert set(default_summary.overridable) == set(INFERENCE_OVERRIDES)
|
||||||
|
|
||||||
|
|
||||||
def test_router_disabled_by_default_so_its_fields_are_inert(default_summary):
|
def test_router_disabled_by_default_so_its_fields_are_inert(default_summary):
|
||||||
assert "stage1_model.router.n_experts" in default_summary.inert
|
assert "stage1_model.router.n_experts" in default_summary.inert
|
||||||
assert "stage1_model.router.temperature" in default_summary.inert
|
assert "stage1_model.router.temperature" in default_summary.inert
|
||||||
@@ -135,3 +147,5 @@ def test_cli_default_smoke():
|
|||||||
assert "parameters" in result.output
|
assert "parameters" in result.output
|
||||||
assert "trunk" in result.output
|
assert "trunk" in result.output
|
||||||
assert "inert under this config" in result.output
|
assert "inert under this config" in result.output
|
||||||
|
assert "inference-overridable without retraining" in result.output
|
||||||
|
assert "stage2_model.n_sec.sampling" in result.output
|
||||||
|
|||||||
Reference in New Issue
Block a user