diff --git a/.bumpversion.toml b/.bumpversion.toml index 006100e..b4fbb41 100644 --- a/.bumpversion.toml +++ b/.bumpversion.toml @@ -1,5 +1,5 @@ [tool.bumpversion] -current_version = "0.3.10" +current_version = "0.3.15" parse = "(?P\\d+)\\.(?P\\d+)\\.(?P\\d+)" serialize = ["{major}.{minor}.{patch}"] search = "{current_version}" @@ -8,7 +8,7 @@ regex = false allow_dirty = false commit = true 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"] [[tool.bumpversion.files]] diff --git a/.gitea/workflows/ci.yml b/.gitea/workflows/ci.yml index 351101c..8658980 100644 --- a/.gitea/workflows/ci.yml +++ b/.gitea/workflows/ci.yml @@ -2,10 +2,9 @@ name: CI "on": push: - branches: ["**"] + branches: ["master"] tags: ["**"] - pull_request: - branches: [master] + pull_request: {} env: UV_CACHE_DIR: /uv-cache @@ -162,7 +161,7 @@ jobs: uv run git-cliff --tag "$TAG" --unreleased --prepend CHANGELOG.md git add CHANGELOG.md 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 git restore --staged CHANGELOG.md fi @@ -203,7 +202,7 @@ jobs: git config user.name "gitea-actions" git config user.email "actions@git.larsbogner.de" 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 ":refs/tags/${GITHUB_REF_NAME}" git tag -f "${GITHUB_REF_NAME}" HEAD @@ -212,3 +211,33 @@ jobs: echo "Tag version matches project version ($CURRENT_VERSION)" fi - run: uv cache prune --ci + + 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 }}" diff --git a/CHANGELOG.md b/CHANGELOG.md index b4fe538..7fee54c 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,5 +1,37 @@ # 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 + +### Added + +- Add sampled n_sec under n_sec.mode = 'head' [gitea #86](https://git.larsbogner.de/lars/giant/issues/86) + +## [0.3.11] - 2026-08-26 + +### Changed + +- Feat(analysis): per-step secondary multiplicity plots + ## [0.3.10] - 2026-08-26 ### Changed diff --git a/CLAUDE.md b/CLAUDE.md index 7487205..9504afe 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -99,7 +99,7 @@ Secondary energies are a **stick-breaking partition of the `e_sec` budget** from **Validation** (`giant/validate.py`): step-level marginal + KL-divergence comparisons during training (`--validate-every`). -**Analysis** (`giant/analysis/`, `giant analyze` CLI): a lean, streaming rollout-vs-reference plotting pipeline that compares one or more autoregressive `giant rollout` runs against a single held-out miniCaloSim reference steps file shared by all of them, and produces publication-styled PDFs assembled into an HTML gallery — one distinctly colored series per rollout, one reference line/panel. It exploits the fact that rollout output and a raw reference file share a world-frame physical column subset under identical names (`pre_*`/`post_*`/`edep`/`step_length`/`pdg`/`material`/`event_id`), so no ALR/local-frame decode is needed — everything is world-frame mm/MeV. Structure: `sources.py` (canonical LazyFrames + `RolloutSpec`/`Side` — a rollout's opened frames + per-checkpoint diagnostic inputs — + synthetic-termination-row filtering + the secondary view, which is `generation>0 & step_no==0` rollout tracks vs exploded `sec_*_list` reference columns), `variables.py` (the per-step value expressions shared by range sizing and the plot registry), `reduce.py` (the streaming primitives — a single `hist1d` `group_by([group,bin]).len()` pass, per-event scalars, edep-weighted depth/transverse profiles, species share, leakage), `grouping.py`/`context.py` (fixed bin edges + energy-quantile/pdg/material group sets resolved once by `prep` into `shared.json` over the union of the reference and every rollout, so every compute job is one pass with no range scan), `reduced.py` (`Partial`/`Reduced` — the compact self-describing JSON a compute job emits), `catalog.py` (the declarative `PlotSpec` registry — marginals × {overall,energy,pdg,material}, per-event totals, shower profiles/containment, species/leakage, secondaries, distance/confusion summaries, router and type-embedding diagnostics; `giant analyze list` prints every id), `runtime_estimate.py` (per-(plot, chunk) walltime estimates for the submit description), and `render.py` (the only module importing ETPlot's `plotstyle`/LaTeX; dispatches on `Reduced.kind`, writes PDFs + `metadata.yaml`; each rollout gets a stable `ps.get_color(i)` slot by its position in `series`, the reference always draws in one fixed dashed-ink style). `Bundle.rollouts` is a name-keyed dict of `Side`, and every `compute_partial`/`finalize` builds a `Reduced.payload["series"]` dict keyed the same way, with `payload["reference"]` as the one distinguished non-rollout entry. The heatmap-shaped specs (`marginal_distance_summary`, `n_sec_confusion`) and the checkpoint-bound diagnostics (`router_gating.py`, `type_embedding_distance.py`) are inherently one-matrix/one-checkpoint per rollout, so they render as one panel per rollout instead of one line/bar per rollout. +**Analysis** (`giant/analysis/`, `giant analyze` CLI): a lean, streaming rollout-vs-reference plotting pipeline that compares one or more autoregressive `giant rollout` runs against a single held-out miniCaloSim reference steps file shared by all of them, and produces publication-styled PDFs assembled into an HTML gallery — one distinctly colored series per rollout, one reference line/panel. It exploits the fact that rollout output and a raw reference file share a world-frame physical column subset under identical names (`pre_*`/`post_*`/`edep`/`step_length`/`pdg`/`material`/`event_id`), so no ALR/local-frame decode is needed — everything is world-frame mm/MeV. Structure: `sources.py` (canonical LazyFrames + `RolloutSpec`/`Side` — a rollout's opened frames + per-checkpoint diagnostic inputs — + synthetic-termination-row filtering + the secondary view, which is `generation>0 & step_no==0` rollout tracks vs exploded `sec_*_list` reference columns), `variables.py` (the per-step value expressions shared by range sizing and the plot registry), `reduce.py` (the streaming primitives — a single `hist1d` `group_by([group,bin]).len()` pass, per-event scalars, edep-weighted depth/transverse profiles, species share, leakage), `grouping.py`/`context.py` (fixed bin edges + energy-quantile/pdg/material group sets resolved once by `prep` into `shared.json` over the union of the reference and every rollout, so every compute job is one pass with no range scan), `reduced.py` (`Partial`/`Reduced` — the compact self-describing JSON a compute job emits), `catalog.py` (the declarative `PlotSpec` registry — marginals × {overall,energy,pdg,material}, per-event totals, shower profiles/containment, species/leakage, secondaries, distance summaries, router and type-embedding diagnostics; `giant analyze list` prints every id), `runtime_estimate.py` (per-(plot, chunk) walltime estimates for the submit description), and `render.py` (the only module importing ETPlot's `plotstyle`/LaTeX; dispatches on `Reduced.kind`, writes PDFs + `metadata.yaml`; each rollout gets a stable `ps.get_color(i)` slot by its position in `series`, the reference always draws in one fixed dashed-ink style). `Bundle.rollouts` is a name-keyed dict of `Side`, and every `compute_partial`/`finalize` builds a `Reduced.payload["series"]` dict keyed the same way, with `payload["reference"]` as the one distinguished non-rollout entry. The heatmap-shaped specs (`marginal_distance_summary`, `sec_count_per_step_by_species` — the latter also drawing the reference as its own panel) and the checkpoint-bound diagnostics (`router_gating.py`, `type_embedding_distance.py`) are inherently one-matrix/one-checkpoint per rollout, so they render as one panel per rollout instead of one line/bar per rollout. **Input is one or more `giant rollout` YAML sidecars** (`run.py:load_rollout_yamls`): each YAML's `output`/`dataset` keys name its rollout parquet and seed file (= the reference truth); every supplied YAML must resolve to the same `dataset`, checked up front with a clear error otherwise (the premise is "N candidates vs one ground truth"). Each rollout's series name comes from a repeated `--label` CLI flag, else the YAML stem (N>1), else `"rollout"` (a single YAML). `prep` creates a **run directory** (`/analysis_runs/analysis_/` by default, `--run-dir` to override) holding `shared.json`, `run_meta.json` (`RunMeta.rollouts: list[{name,path,plot_meta}]`, insertion order = CLI order = every plot's series order), `reduced_partial/`, `reduced/`, `plots/`. **Compute/merge/render split:** `giant analyze prep a.yaml [b.yaml ...] --chunks N` records `N` in `run_meta.json`, and the workflow's `AnalysisComputeTask` runs one HTCondor job per (plot, chunk) pair (`compute-one --id --chunk --run-dir`, polars/numpy only — no LaTeX on workers), each streaming over an `event_id`-disjoint slice (`event_id % N == chunk`) of the reference **and every rollout** and writing a small `reduced_partial/__.json`; every `PlotSpec` splits into a `compute_partial`/`finalize` pair so chunks can be summed/concatenated back per rollout (`chunkable=False` specs — the checkpoint-bound diagnostics, already bounded/subsampled — always run as a single chunk). The local `giant analyze render ` first joins every plot's chunk partials into `reduced/.json` (`merge_all`, a no-op join when `N=1`; `merge-one` does a single plot for debugging), then turns those into the styled PDF/gallery tree. `giant analyze metrics ` is a separate, unrelated entry point: training-progress plots straight from a run's `metrics.csv`. @@ -117,6 +117,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. +**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. **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. diff --git a/cliff.toml b/cliff.toml index bca7fc8..7207e63 100644 --- a/cliff.toml +++ b/cliff.toml @@ -37,7 +37,7 @@ commit_preprocessors = [ protect_breaking_commits = false commit_parsers = [ { message = "^Merge ", skip = true }, - { message = "\\[skip ci\\]", skip = true }, + { message = "^chore: (bump version|update changelog|sync project version)", skip = true }, { message = "^Add", group = "Added" }, { message = "^(Fix|Clamp|Clip)", group = "Fixed" }, { message = "^(Remove|Drop|Deprecate)", group = "Removed" }, diff --git a/configs/baseline.toml b/configs/baseline.toml index 3bd095e..0f0a044 100644 --- a/configs/baseline.toml +++ b/configs/baseline.toml @@ -29,11 +29,21 @@ # capacity overfitting is not the binding constraint, and every recent # run used 0.0. # -# Known weak spots this baseline is expected to *exhibit* (they are the -# reason for the comparisons, not a reason to retune this file): every model -# on record under-produces steps per event by ~2x (rollout ~7e4 vs Geant4 -# ~1.4e5) and secondaries per event by 2-3.5x (~2-3e4 vs 7.2e4), and n_sec -# head accuracy sits at 0.863-0.867 regardless of size or objective. +# Known weak spots, now measured against this exact config rather than +# extrapolated from the pre-v0.3 field (analysis_341dfb14, best.pt @ epoch +# 50/50, full writeup: knowledge-base/experiments/ +# giant-baseline-flow-ar-rollout-validation.md). Unlike every pre-v0.3 +# 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] # REQUIRED. Without it config.migrate_config reads this file as v0.2 and diff --git a/giant/analysis/catalog.py b/giant/analysis/catalog.py index 18fbc1b..eaa81eb 100644 --- a/giant/analysis/catalog.py +++ b/giant/analysis/catalog.py @@ -22,9 +22,9 @@ which is the order rollouts were given on the CLI) plus the single reference. ``finalize`` merges each rollout's chunks independently and assembles a ``Reduced.payload`` keyed the same way: ``"series": {name: ...}`` for the rollouts, ``"reference": ...`` as one distinguished entry (omitted on -rollout-only plots like ``leakage_fraction``). The two heatmap-shaped specs -(``marginal_distance_summary``, ``n_sec_confusion``) and the router -diagnostics are inherently one-matrix/one-checkpoint per rollout, so their +rollout-only plots like ``leakage_fraction``). The heatmap-shaped specs +(``marginal_distance_summary``, ``sec_count_per_step_by_species``) and the +router diagnostics are inherently one-matrix/one-checkpoint per rollout, so their ``"series"`` entries are whole per-rollout artifacts (a matrix, a gating dict) rather than a single number/array — ``render.py`` draws those as one panel per rollout instead of one line/bar per rollout. @@ -60,7 +60,6 @@ from giant.analysis.reduce import ( leakage_fraction, profile_finalize, profile_partial, - sec_count_by_event, species_share, sum_merge, transverse_expr, @@ -72,7 +71,15 @@ from giant.analysis.router_gating import ( compute_router_share_by_process, compute_router_specialization, ) -from giant.analysis.sources import RolloutSide, RolloutSpec, Side, open_side, physical_steps, secondaries +from giant.analysis.sources import ( + RolloutSide, + RolloutSpec, + Side, + open_side, + physical_steps, + secondaries, + secondaries_by_step, +) from giant.analysis.type_embedding_distance import compute_type_embedding_l1_distance from giant.analysis.variables import RANGED_VARS, cos_scatter_expr @@ -211,31 +218,6 @@ def _ks_statistic(r_counts, t_counts) -> float: return float(np.max(np.abs(r_cdf - t_cdf))) -def _integer_confusion( - t: np.ndarray, r: np.ndarray, max_bins: int = 21, cap: int | None = None -) -> tuple[list[str], np.ndarray]: - """Confusion matrix of two paired small-integer arrays (e.g. secondary counts). - - Bins are consecutive integers ``0..cap``, with the last bin an overflow - ``"cap+"`` bucket, so an occasional pathological count doesn't blow up the - heatmap. Returns ``(labels, matrix)`` with ``matrix[i, j]`` counting pairs - with ``t == i`` and ``r == j`` (both clipped into ``[0, cap]``). - - ``cap``, if given, is used as-is instead of being derived from ``t``/``r`` - — lets a multi-rollout caller fix one shared cap (and so one shared label - set) across every rollout's matrix rather than each panel picking its own. - """ - if cap is None: - cap = min(max(int(t.max()) if len(t) else 0, int(r.max()) if len(r) else 0, 1), max_bins - 1) - t_c = np.clip(t.astype(np.int64), 0, cap) - r_c = np.clip(r.astype(np.int64), 0, cap) - n = cap + 1 - mat = np.zeros((n, n), dtype=np.int64) - np.add.at(mat, (t_c, r_c), 1) - labels = [str(i) for i in range(cap)] + [f"{cap}+"] - return labels, mat - - def _containment_depths(mat: np.ndarray, edges: np.ndarray, quantile: float) -> np.ndarray: """Per-event depth containing ``quantile`` of that event's deposited energy. @@ -802,6 +784,139 @@ def _sec_count_per_species_finalize(parts: list[dict], ctx: Context) -> Reduced: ) +# Per-step secondary multiplicity. Fixed integer edges (bin i == exactly i +# secondaries, the top bin an overflow bucket) keep both plots sum-mergeable +# across chunks — no shared-range pass needed. The species heatmap gets a +# shorter row axis because a single step rarely emits many of *one* species. +_N_SEC_STEP_CAP = 20 +_N_SEC_SPECIES_CAP = 10 +_OTHER_KEY = "other" + + +def _n_sec_edges(cap: int) -> np.ndarray: + return np.arange(-0.5, cap + 1.5) + + +def _sec_step_key_lf(lf: pl.LazyFrame, side: Side) -> pl.LazyFrame: + """Secondaries with their emitting-step key. + + The rollout side reads *all* rows, not just physical ones: a secondary + whose very first row is a synthetic termination row (born, then immediately + escaped or cut) was still produced by its parent step, and dropping it would + undercount that step's multiplicity. + """ + return secondaries_by_step(lf, side) + + +def _n_steps(lf: pl.LazyFrame) -> int: + """Number of (physical) step rows — the denominator the zero rows come from.""" + return int(lf.select(pl.len()).collect(engine="streaming").item()) + + +def _sec_count_per_step_partial(b: Bundle) -> dict: + edges = _n_sec_edges(_N_SEC_STEP_CAP) + + def _side(sec_lf: pl.LazyFrame, steps_lf: pl.LazyFrame) -> dict: + per_step = sec_lf.group_by("step_key").agg(pl.len().alias("n")) + return { + "h": _partial_hist(per_step, pl.col("n").clip(0, _N_SEC_STEP_CAP), edges), + "n_steps": _n_steps(steps_lf), + } + + return { + "r": _per_rollout(b, lambda rs: _side(_sec_step_key_lf(rs.all, Side.rollout), rs.phys)), + "t": _side(_sec_step_key_lf(b.t_all, Side.reference), b.t_phys), + } + + +def _zero_filled(part_hists: list[dict], n_steps: int, key, nbins: int) -> list[int]: + """Merged counts for one series, with bin 0 (= steps that emitted none) filled in. + + The reduction only ever sees steps that produced at least one secondary, so + the empty ones are recovered by subtraction from the total step count. + """ + counts = _finalize_counts(sum_merge(part_hists), key, nbins) + counts[0] = max(n_steps - int(sum(counts)), 0) + return [int(c) for c in counts] + + +def _sec_count_per_step_finalize(parts: list[dict], ctx: Context) -> Reduced: + edges = _n_sec_edges(_N_SEC_STEP_CAP) + nb = len(edges) - 1 + names = list(parts[0]["r"]) + series = { + name: _zero_filled([p["r"][name]["h"] for p in parts], sum(p["r"][name]["n_steps"] for p in parts), 0, nb) + for name in names + } + return Reduced( + id="sec_count_per_step", + family="secondaries", + kind="overlay_hist", + title="Number of secondaries per step", + xlabel="secondaries per step", + payload={ + "edges": edges.tolist(), + "series": series, + "reference": _zero_filled([p["t"]["h"] for p in parts], sum(p["t"]["n_steps"] for p in parts), 0, nb), + "log_y": True, + }, + ) + + +def _species_key_expr(top_pdgs: list[int]) -> pl.Expr: + """``pdg`` bucketed into the shared top-K columns plus one ``other`` bin.""" + return pl.when(pl.col("pdg").is_in(list(top_pdgs))).then(pl.col("pdg").cast(pl.Utf8)).otherwise(pl.lit(_OTHER_KEY)) + + +def _sec_count_per_step_by_species_partial(b: Bundle) -> dict: + edges = _n_sec_edges(_N_SEC_SPECIES_CAP) + group = _species_key_expr(b.ctx.top_pdgs) + + def _side(sec_lf: pl.LazyFrame, steps_lf: pl.LazyFrame) -> dict: + per_step_species = sec_lf.group_by("step_key", "pdg").agg(pl.len().alias("n")) + return { + "h": _partial_hist(per_step_species, pl.col("n").clip(0, _N_SEC_SPECIES_CAP), edges, group=group), + "n_steps": _n_steps(steps_lf), + } + + return { + "r": _per_rollout(b, lambda rs: _side(_sec_step_key_lf(rs.all, Side.rollout), rs.phys)), + "t": _side(_sec_step_key_lf(b.t_all, Side.reference), b.t_phys), + } + + +def _sec_count_per_step_by_species_finalize(parts: list[dict], ctx: Context) -> Reduced: + edges = _n_sec_edges(_N_SEC_SPECIES_CAP) + nb = len(edges) - 1 + names = list(parts[0]["r"]) + keys = [str(p) for p in ctx.top_pdgs] + [_OTHER_KEY] + + def _matrix(hists: list[dict], n_steps: int) -> list[list[int]]: + # columns = species, rows = multiplicity; every species gets its own + # zero row (steps that produced none of *that* species). + cols = [_zero_filled(hists, n_steps, k, nb) for k in keys] + return [[cols[j][i] for j in range(len(keys))] for i in range(nb)] + + return Reduced( + id="sec_count_per_step_by_species", + family="secondaries", + kind="heatmap", + title="Per-step secondary multiplicity by species", + xlabel="species", + payload={ + "series": { + n: _matrix([p["r"][n]["h"] for p in parts], sum(p["r"][n]["n_steps"] for p in parts)) for n in names + }, + "reference": _matrix([p["t"]["h"] for p in parts], sum(p["t"]["n_steps"] for p in parts)), + "row_labels": [str(i) for i in range(_N_SEC_SPECIES_CAP)] + [f"{_N_SEC_SPECIES_CAP}+"], + "col_labels": [pdg_label(k) for k in ctx.top_pdgs] + [_OTHER_KEY], + "ylabel": "secondaries of this species per step", + "cbar_label": "step count", + "log_color": True, + }, + ) + + def _sec_energy_partial(b: Bundle) -> dict: edges = np.linspace(*b.ctx.sec_energy_range, b.ctx.n_sec_bins + 1) return { @@ -858,62 +973,6 @@ def _sec_cos_angle_finalize(parts: list[dict], ctx: Context) -> Reduced: ) -def _n_sec_confusion_partial(b: Bundle) -> dict: - t_ids, t_n = sec_count_by_event(b.t_all, _t_sec(b)) - - def _r(rs: RolloutSide) -> dict: - ids, n = sec_count_by_event(rs.phys, _r_sec(rs)) - return {"ids": ids.tolist(), "n": n.tolist()} - - return {"r": _per_rollout(b, _r), "t": {"ids": t_ids.tolist(), "n": t_n.tolist()}} - - -def _n_sec_confusion_finalize(parts: list[dict], ctx: Context) -> Reduced: - names = list(parts[0]["r"]) - # event-disjoint chunking (see Bundle.open) means each event_id appears in - # exactly one part on each side, so a plain dict build is a safe merge. - t_ids = np.concatenate([np.asarray(p["t"]["ids"], dtype=np.int64) for p in parts]) - t_n = np.concatenate([np.asarray(p["t"]["n"], dtype=np.int64) for p in parts]) - t_map = dict(zip(t_ids.tolist(), t_n.tolist())) - - pairs: dict[str, tuple[np.ndarray, np.ndarray]] = {} - max_val = 0 - for name in names: - r_ids = np.concatenate([np.asarray(p["r"][name]["ids"], dtype=np.int64) for p in parts]) - r_n = np.concatenate([np.asarray(p["r"][name]["n"], dtype=np.int64) for p in parts]) - r_map = dict(zip(r_ids.tolist(), r_n.tolist())) - common = sorted(set(r_map) & set(t_map)) - true_n = np.array([t_map[e] for e in common], dtype=np.int64) - pred_n = np.array([r_map[e] for e in common], dtype=np.int64) - pairs[name] = (true_n, pred_n) - if len(true_n): - max_val = max(max_val, int(true_n.max()), int(pred_n.max())) - - cap = min(max(max_val, 1), 20) - matrices: dict[str, list[list[int]]] = {} - labels: list[str] = [] - for name in names: - true_n, pred_n = pairs[name] - labels, mat = _integer_confusion(true_n, pred_n, cap=cap) - matrices[name] = mat.tolist() - - return Reduced( - id="n_sec_confusion", - family="secondaries", - kind="heatmap", - title="Predicted vs true secondary count per event", - xlabel="predicted secondaries (rollout)", - payload={ - "series": matrices, - "row_labels": labels, - "col_labels": labels, - "ylabel": "true secondaries (reference)", - "cbar_label": "event count", - "vmin": 0.0, - }, - ) - - # --------------------------------------------------------------------------- # router diagnostics (not chunked — already bounded/subsampled) # --------------------------------------------------------------------------- @@ -1077,6 +1136,18 @@ def build_catalog() -> list[PlotSpec]: compute_partial=_sec_count_per_species_partial, finalize=_sec_count_per_species_finalize, ), + PlotSpec( + "sec_count_per_step", + "secondaries", + compute_partial=_sec_count_per_step_partial, + finalize=_sec_count_per_step_finalize, + ), + PlotSpec( + "sec_count_per_step_by_species", + "secondaries", + compute_partial=_sec_count_per_step_by_species_partial, + finalize=_sec_count_per_step_by_species_finalize, + ), PlotSpec( "sec_energy", "secondaries", @@ -1089,12 +1160,6 @@ def build_catalog() -> list[PlotSpec]: compute_partial=_sec_cos_angle_partial, finalize=_sec_cos_angle_finalize, ), - PlotSpec( - "n_sec_confusion", - "secondaries", - compute_partial=_n_sec_confusion_partial, - finalize=_n_sec_confusion_finalize, - ), PlotSpec( "router_gating", "model", diff --git a/giant/analysis/reduce.py b/giant/analysis/reduce.py index 8e1d48e..ce65c89 100644 --- a/giant/analysis/reduce.py +++ b/giant/analysis/reduce.py @@ -271,20 +271,3 @@ def leakage_fraction(lf: pl.LazyFrame) -> np.ndarray: escaped = per_event["escaped"].fill_null(0.0).to_numpy() total = deposited + escaped return np.where(total > 0, escaped / total, 0.0) - - -def sec_count_by_event(lf_all: pl.LazyFrame, sec_lf: pl.LazyFrame) -> tuple[np.ndarray, np.ndarray]: - """Per-event secondary count, zero-filled for events that produced none. - - Two bounded per-event ``group_by``s — the full event set (from ``lf_all``) - and the secondary counts (from ``sec_lf``, see ``sources.secondaries``) — - merged in Python via a dict. Both results are event-granularity (not - per-row), so this stays in the same bounded-memory budget as - ``event_scalars``; a plain ``group_by`` on ``sec_lf`` alone would silently - drop zero-secondary events instead of zero-filling them. - """ - ev = lf_all.select("event_id").unique().collect(engine="streaming")["event_id"].to_numpy() - cnt_df = sec_lf.group_by("event_id").agg(pl.len().alias("n")).collect(engine="streaming") - cnt = dict(zip(cnt_df["event_id"].to_list(), cnt_df["n"].to_list())) - counts = np.array([cnt.get(int(e), 0) for e in ev], dtype=np.int64) - return ev, counts diff --git a/giant/analysis/reduced.py b/giant/analysis/reduced.py index bbf79ac..f697cbc 100644 --- a/giant/analysis/reduced.py +++ b/giant/analysis/reduced.py @@ -27,7 +27,7 @@ from pathlib import Path # "router_specialization" max gate weight vs energy (one scalar trend line # summarizing "router_gating"), per rollout with an enabled router # "heatmap" row x col matrix + colorbar, one panel per rollout (a -# distance scorecard or a predicted-vs-true confusion matrix) +# distance scorecard) # "unavailable" plot not applicable to this run (e.g. no MoE checkpoint) diff --git a/giant/analysis/render.py b/giant/analysis/render.py index 46a684d..e447bc8 100644 --- a/giant/analysis/render.py +++ b/giant/analysis/render.py @@ -27,6 +27,7 @@ from pathlib import Path import numpy as np import plotstyle as ps +from matplotlib.colors import LogNorm import yaml from giant.analysis.reduced import Reduced @@ -376,10 +377,16 @@ def _render_router_specialization(r: Reduced, params: dict): def _render_heatmap(r: Reduced, params: dict): - series = r.payload["series"] + series = dict(r.payload["series"]) row_labels = r.payload["row_labels"] col_labels = r.payload["col_labels"] + # A heatmap-shaped plot is one matrix per rollout, so the reference (when the + # comparison has one — the distance scorecard doesn't) becomes one more panel + # rather than another line. + if r.payload.get("reference") is not None: + series["reference"] = r.payload["reference"] names = list(series) + norm = LogNorm(vmin=1) if r.payload.get("log_color") else None fig, axes = ps.new_figure( "slide-16x9" if len(names) > 1 else "thesis-single", title=r.title, @@ -397,8 +404,9 @@ def _render_heatmap(r: Reduced, params: dict): origin="upper", aspect="auto", cmap=r.payload.get("cmap", "viridis"), - vmin=r.payload.get("vmin"), - vmax=r.payload.get("vmax"), + norm=norm, + vmin=None if norm else r.payload.get("vmin"), + vmax=None if norm else r.payload.get("vmax"), ) ax.set_xticks(range(len(col_labels))) ax.set_xticklabels(col_labels, rotation=45, ha="right") diff --git a/giant/analysis/sources.py b/giant/analysis/sources.py index 15b496a..42dccc1 100644 --- a/giant/analysis/sources.py +++ b/giant/analysis/sources.py @@ -236,3 +236,34 @@ def secondaries(lf: pl.LazyFrame, side: Side) -> pl.LazyFrame: pl.col("sec_dz_list").alias("sdz"), ) ) + + +def secondaries_by_step(lf: pl.LazyFrame, side: Side) -> pl.LazyFrame: + """One row per produced secondary, tagged with the step that produced it. + + Canonical columns: ``step_key`` (an opaque struct identifying the emitting + step) and ``pdg``. ``secondaries`` deliberately drops that link; the + per-step multiplicity plots need it, so this is a separate view rather than + extra columns every other consumer would pay for. + + - rollout: a secondary's birth row carries ``parent_id`` and a birth + position copied verbatim from the parent step's ``post_pos``, so + ``(event_id, parent_id, pre_pos)`` identifies the emitting step exactly — + no join against the (large) step frame is needed. + - reference: secondaries already live on their parent step's row, so the + row index *is* the step key. It is only ever used as a group key inside + one chunk's own aggregation, so indices repeating across chunks is + harmless. + """ + if side is Side.rollout: + return lf.filter((pl.col("generation") > 0) & (pl.col("step_no") == 0)).select( + pl.struct("event_id", "parent_id", "pre_x", "pre_y", "pre_z").alias("step_key"), + "pdg", + ) + return ( + lf.select("sec_pdg_list") + .with_row_index("_row") + .explode("sec_pdg_list") + .drop_nulls("sec_pdg_list") + .select(pl.struct("_row").alias("step_key"), pl.col("sec_pdg_list").cast(pl.Int64).alias("pdg")) + ) diff --git a/giant/checkpoint_io.py b/giant/checkpoint_io.py index fe6ce63..258d2a1 100644 --- a/giant/checkpoint_io.py +++ b/giant/checkpoint_io.py @@ -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 `CheckpointCompatibilityError` with the same wording the CLI has always 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 dataclasses import dataclass +import copy +from dataclasses import dataclass, field from pathlib import Path import torch @@ -29,13 +36,45 @@ from giant.constants import K_MAX from giant.data.loader import TopNMap from giant.data.setup_cache import topnmap_from_json 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): """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]: """(particle_conditioning, material_conditioning) for `giant.data.transforms.build_cond_features`/`build_features` — from @@ -128,6 +167,7 @@ class InferenceContext: model_config: dict epoch: int | None best_val_loss: float | None + config_overrides: dict[str, object] = field(default_factory=dict) def load_for_inference( @@ -136,6 +176,7 @@ def load_for_inference( command_name: str, weights: str = "raw", require_stage2: bool = True, + config_overrides: dict[str, object] | None = None, ) -> InferenceContext: """Load *checkpoint* and reconstruct everything `predict`/`rollout` need 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 parameter since `stage{1,2}_model.active` is a real, if currently 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) for key in ("model_config", "sec_decoder"): @@ -159,7 +207,7 @@ def load_for_inference( 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) pdg_topn_map = load_pdg_topn_map(ckpt) mat_topn_map = load_mat_topn_map(ckpt) @@ -231,4 +279,5 @@ def load_for_inference( model_config=model_cfg, epoch=ckpt.get("epoch"), best_val_loss=ckpt.get("best_val_loss"), + config_overrides=dict(config_overrides) if config_overrides else {}, ) diff --git a/giant/cli.py b/giant/cli.py index db1a2d0..49afbe2 100644 --- a/giant/cli.py +++ b/giant/cli.py @@ -1,3 +1,5 @@ +from __future__ import annotations + from collections import Counter from datetime import datetime, timezone from enum import Enum @@ -5,18 +7,14 @@ import math from pathlib import Path import re import sys -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 ( @@ -26,30 +24,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) @@ -142,6 +121,23 @@ def _parse_router_axis_flags(specs: list[str]) -> dict[str, object]: 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( router: bool | None, router_type: str | None, @@ -203,6 +199,9 @@ def _write_prediction_ref( is kept, so ad-hoc runs and the ``/ceph`` predictions convention are unaffected. """ + import yaml + + ref = { "prediction_id": pred_uuid, "output": str(out), @@ -632,6 +631,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: @@ -1030,8 +1033,36 @@ def predict( help="Free-text note recorded in the prediction's YAML sidecar", ), ] = 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: """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": @@ -1050,8 +1081,11 @@ def predict( typer.echo(f"device: {_device}") # --- Load checkpoint --- + config_overrides = _parse_set_flags(set_) 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: typer.echo(f"error: {exc}", err=True) raise typer.Exit(1) @@ -1325,6 +1359,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): @@ -1418,8 +1456,28 @@ def rollout( Optional[int], typer.Option("--seed", help="Torch/numpy seed for reproducibility"), ] = 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: """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) @@ -1427,8 +1485,11 @@ def rollout( _device = torch.device(device) if device else gconfig.auto_device() typer.echo(f"device: {_device}") + config_overrides = _parse_set_flags(set_) 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: typer.echo(f"error: {exc}", err=True) raise typer.Exit(1) @@ -1544,6 +1605,7 @@ def rollout( # model knob (router type/n_experts, noise_dim, vocab sizes, ...) # is available downstream without touching this command again. "model_config": dict(model_cfg), + "config_overrides": dict(ctx.config_overrides), "training_epoch": ctx.epoch, "best_val_loss": ctx.best_val_loss, # [train]/[meta] from the sibling config.toml (giant.config.save_config) diff --git a/giant/config.py b/giant/config.py index a7d7fe7..30aba58 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): @@ -404,6 +406,16 @@ class Stage2RouterConfig(RouterConfig): 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) class NSecConfig: # "head": a classifier over {0..k_max} on the condition encoding alone @@ -424,11 +436,15 @@ class NSecConfig: # n_sec head was trained against Stage 1's own ConditionEncoder output and so has # to stay attached there, not just be labeled as such). owner: str = "stage2" - # mode="stop_token" only: how sample_secondaries_ar turns a slot's stop logit into a - # stop/continue decision. "greedy": sigmoid(logit) >= 0.5 (deterministic). "sample": - # a Bernoulli draw at sigmoid(logit) (a real sample from the learned length - # distribution, at the cost of an extra RNG draw per slot). - stop_sampling: str = "greedy" + # How resolve_n_sec/sample_secondaries_ar turn a count-bearing head's output into an + # actual n_sec decision. mode="head": "greedy" is argmax over the classifier logits + # (deterministic — the conditional mode, not a sample); "sample" is a categorical draw + # from softmax(logits) (a real sample from the learned count distribution). mode= + # "stop_token": "greedy" is sigmoid(stop_logit) >= 0.5 per slot (deterministic); + # "sample" is a Bernoulli draw at sigmoid(stop_logit) per slot. Renamed from + # "stop_sampling" (gitea #86), which is still accepted as a deprecated alias since it + # appears in existing checkpoints' model_config. + sampling: str = "greedy" @classmethod def from_dict(cls, d: dict | None) -> "NSecConfig": @@ -437,7 +453,7 @@ class NSecConfig: mode=d.get("mode", "head"), lambda_weight=d.get("lambda", 0.1), owner=d.get("owner", "stage2"), - stop_sampling=d.get("stop_sampling", "greedy"), + sampling=d.get("sampling", d.get("stop_sampling", "greedy")), ) def to_dict(self) -> dict: @@ -445,7 +461,7 @@ class NSecConfig: "mode": self.mode, "lambda": self.lambda_weight, "owner": self.owner, - "stop_sampling": self.stop_sampling, + "sampling": self.sampling, } @@ -927,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(): @@ -972,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() @@ -1075,6 +1095,27 @@ def _set_path(d: dict, dotted: str, value) -> None: cur[parts[-1]] = value +def _pop_path(d: dict, dotted: str) -> None: + """Remove a dotted path from a nested dict, if present. No-op if any + component along the path is missing.""" + parts = dotted.split(".") + cur = d + for part in parts[:-1]: + if not isinstance(cur, dict) or part not in cur: + return + cur = cur[part] + if isinstance(cur, dict): + cur.pop(parts[-1], None) + + +# Config keys renamed within v0.3 itself (not part of the v0.2->v0.3 migration +# above) — normalized by migrate_config so a config.toml still using an older +# v0.3 key name keeps passing validate_config_keys. +_RENAMED_KEYS = { + "stage2_model.n_sec.stop_sampling": "stage2_model.n_sec.sampling", # gitea #86 +} + + def _deep_merge(base: dict, override: dict) -> dict: """Recursively merge `override` onto a copy of `base`. @@ -1093,6 +1134,72 @@ def _deep_merge(base: dict, override: dict) -> dict: 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) class FlagSpec: """One CLI flag's mapping into the config-overrides tree. @@ -1271,11 +1378,21 @@ def migrate_config(cfg: dict) -> dict: (which additionally carries n_sec_head ownership and needs `network.build_models`'s cooperation) is a separate migration surface, deferred to the network.py refactor. - """ - if _get_path(cfg, "meta.config_version") == CONFIG_VERSION: - return copy.deepcopy(cfg) + Independently of the v0.2/v0.3 branch below, `_RENAMED_KEYS` normalizes + keys renamed within v0.3 itself (e.g. `stop_sampling` -> `sampling`, + gitea #86) so a config.toml written against an older v0.3 key name still + passes `validate_config_keys`. + """ cfg = copy.deepcopy(cfg) + for old_path, new_path in _RENAMED_KEYS.items(): + if _get_path(cfg, old_path) is not None and _get_path(cfg, new_path) is None: + _set_path(cfg, new_path, _get_path(cfg, old_path)) + _pop_path(cfg, old_path) + + if _get_path(cfg, "meta.config_version") == CONFIG_VERSION: + return cfg + old_train = cfg.pop("train", {}) old_model = cfg.pop("model", {}) old_router = dict(old_model.pop("router", {})) @@ -1512,9 +1629,9 @@ def validate_config(cfg: dict, *, resume: bool = False) -> None: "conditioning to hang an EOS decision off" ) - stop_sampling = _get_path(cfg, "stage2_model.n_sec.stop_sampling") - if stop_sampling not in ("greedy", "sample"): - raise ValueError(f"stage2_model.n_sec.stop_sampling = {stop_sampling!r} — must be 'greedy' or 'sample'") + n_sec_sampling = _get_path(cfg, "stage2_model.n_sec.sampling") + 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'") precision = _get_path(cfg, "train.precision") if precision not in ("fp32", "bf16"): @@ -1569,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( @@ -1783,6 +1902,9 @@ def epoch_seed(seed: int, epoch: int) -> int: def seed_everything(seed: int) -> None: + import numpy as np + import torch + random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) @@ -1836,6 +1958,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/model/builders.py b/giant/model/builders.py index 0849db2..6c1a8b3 100644 --- a/giant/model/builders.py +++ b/giant/model/builders.py @@ -138,7 +138,7 @@ def build_models(model_config: dict) -> dict[str, nn.Module | None]: n_sec_head_cfg=s2_spec.heads.n_sec.to_dict(), type_head_cfg=s2_spec.heads.type.to_dict(), build_stop_head=stop_token, - stop_sampling=s2_spec.n_sec.stop_sampling, + n_sec_sampling=s2_spec.n_sec.sampling, stop_head_cfg=s2_spec.heads.n_sec.to_dict(), ) else: @@ -168,6 +168,7 @@ def build_models(model_config: dict) -> dict[str, nn.Module | None]: cond_enc=shared_cond_enc, n_sec_head_cfg=s2_spec.heads.n_sec.to_dict(), type_head_cfg=s2_spec.heads.type.to_dict(), + n_sec_sampling=s2_spec.n_sec.sampling, ) return result diff --git a/giant/model/models.py b/giant/model/models.py index 8fb91c3..0e26323 100644 --- a/giant/model/models.py +++ b/giant/model/models.py @@ -370,6 +370,7 @@ class Stage2OneShot(StageModel): cond_enc: ConditionEncoder | None = None, n_sec_head_cfg: dict | None = None, type_head_cfg: dict | None = None, + n_sec_sampling: str = "greedy", ) -> None: super().__init__( pdg_vocab, @@ -383,6 +384,7 @@ class Stage2OneShot(StageModel): particle_type_cfg=particle_type_cfg, cond_enc=cond_enc, ) + self.n_sec_sampling = n_sec_sampling self._build_context_fusion(x_dim, context_dim, cond_out_dim) target = self.particle_type_cfg.target type_head_out_dim = None if target == "physical" else k_max * self.type_dim @@ -498,7 +500,7 @@ class Stage2Autoregressive(StageModel): n_sec_head_cfg: dict | None = None, type_head_cfg: dict | None = None, build_stop_head: bool = False, - stop_sampling: str = "greedy", + n_sec_sampling: str = "greedy", stop_head_cfg: dict | None = None, ) -> None: super().__init__( @@ -514,7 +516,7 @@ class Stage2Autoregressive(StageModel): cond_enc=cond_enc, ) self.history_kind = history - self.stop_sampling = stop_sampling + self.n_sec_sampling = n_sec_sampling self.context_adapter = ContextAdapter(x_dim, context_dim) self.base_fuse = nn.Sequential( nn.Linear(cond_out_dim + context_dim, cond_out_dim), diff --git a/giant/model/summary.py b/giant/model/summary.py index 36a34d3..a00f524 100644 --- a/giant/model/summary.py +++ b/giant/model/summary.py @@ -37,7 +37,7 @@ from dataclasses import dataclass, field 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.trunks import RoutedTrunk @@ -111,6 +111,7 @@ class ModelSummary: pdg_vocab: int mat_vocab: int 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: @@ -140,7 +141,7 @@ def _fingerprint(modules: dict[str, nn.Module]) -> list: exist, every parameter's/buffer's shape+dtype (never values — those are randomly initialized and irrelevant to *structure*), and every plain scalar attribute any module stores on itself (e.g. `Stage2Autoregressive - .stop_sampling`, `EnergyRouter.temperature`) — this is what makes a + .n_sec_sampling`, `EnergyRouter.temperature`) — this is what makes a non-parametric key's effect on construction observable.""" sig = [] for stage_name, module in modules.items(): @@ -234,6 +235,8 @@ def summarize_model(cfg: dict, pdg_vocab: int, mat_vocab: int) -> ModelSummary: else: inert.append(path) + overridable = sorted(p for p in in_scope if p in INFERENCE_OVERRIDES) + return ModelSummary( modules=modules, consumed=sorted(consumed), @@ -242,6 +245,7 @@ def summarize_model(cfg: dict, pdg_vocab: int, mat_vocab: int) -> ModelSummary: pdg_vocab=pdg_vocab, mat_vocab=mat_vocab, vocab_caveats=_vocab_caveats(cfg), + overridable=overridable, ) @@ -309,6 +313,12 @@ def render_summary(summary: ModelSummary) -> str: else: 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: lines.append("") lines.append("vocab placeholder caveats:") diff --git a/giant/sample.py b/giant/sample.py index c1d3f15..c0ac973 100644 --- a/giant/sample.py +++ b/giant/sample.py @@ -261,7 +261,7 @@ def sample_secondaries_ar( slot's own stop logit (`predict_stop`, evaluated on the same prefix conditioning as the token itself — see `predict_type`'s docstring for why this needs no extra state) decides whether generation should have - already stopped, per `sec_decoder.stop_sampling` ("greedy": threshold at + already stopped, per `sec_decoder.n_sec_sampling` ("greedy": threshold at 0; "sample": a Bernoulli draw at `sigmoid(logit)`). A row's own `n_sec_pred` is the first slot index where this fires; once every row in the batch has fired, the loop breaks before spending a model call on the @@ -345,7 +345,7 @@ def sample_secondaries_ar( slot_idx, hist=hist, ).squeeze(1) - if sec_decoder.stop_sampling == "sample": + if sec_decoder.n_sec_sampling == "sample": stop_now = torch.rand(B, device=device) < torch.sigmoid(stop_logit) else: stop_now = stop_logit >= 0.0 @@ -496,7 +496,12 @@ def resolve_n_sec( Raises if neither stage owns any n_sec mechanism at all — the only way that happens is `stage2_model.n_sec.mode = "truth"`, which is not a valid - rollout-/predict-capable checkpoint.""" + rollout-/predict-capable checkpoint. + + `n_sec.mode = "head"` resolves the classifier logits per + `sec_decoder.n_sec_sampling`: "greedy" (default) takes the conditional + mode via argmax; "sample" draws a real sample from the learned count + distribution via `torch.multinomial` on the softmax — see gitea #86.""" if n_sec_pred is not None: return n_sec_pred if getattr(sec_decoder, "stop_head", None) is not None: @@ -508,4 +513,6 @@ def resolve_n_sec( "'truth' is standalone-evaluation-only" ) logits = sec_decoder.predict_n_sec(cond_cont, cond_cat, stage1_out) + if sec_decoder.n_sec_sampling == "sample": + return torch.multinomial(logits.softmax(dim=-1), 1).squeeze(-1) return logits.argmax(dim=-1) 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/pyproject.toml b/pyproject.toml index 391505f..293389a 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "giant" -version = "0.3.10" +version = "0.3.15" description = "Geant4 step-function surrogate via conditional flow matching" readme = "README.md" requires-python = ">=3.12" diff --git a/tests/test_analysis_reduce.py b/tests/test_analysis_reduce.py index 71cb4a8..0a456f0 100644 --- a/tests/test_analysis_reduce.py +++ b/tests/test_analysis_reduce.py @@ -13,6 +13,7 @@ from giant.analysis.sources import ( open_side, physical_steps, secondaries, + secondaries_by_step, ) from giant.data.loader import EVENT_ID_FILE_STRIDE @@ -159,18 +160,17 @@ def test_secondaries_rollout_vs_reference_align(): assert t["pdg"].to_list() == [22, 22] -def test_sec_count_by_event_zero_fills_events_with_no_secondaries(): - r_phys = physical_steps(_rollout_frame(), Side.rollout) - r_sec = secondaries(_rollout_frame(), Side.rollout) - ev, n = R.sec_count_by_event(r_phys, r_sec) - # event 1 has one secondary track; event 2 has none and must still appear (as 0), - # not silently drop out of a plain group_by on the secondaries frame alone. - assert dict(zip(ev.tolist(), n.tolist())) == {1: 1, 2: 0} +def test_secondaries_by_step_keys_each_secondary_to_its_emitting_step(): + r = secondaries_by_step(_rollout_frame(), Side.rollout).collect() + assert r["pdg"].to_list() == [22] + # the rollout key is (event_id, parent_id, birth position) — the parent + # step's post_pos, copied verbatim onto the child's birth row. + assert r["step_key"][0] == {"event_id": 1, "parent_id": 0, "pre_x": 0.0, "pre_y": 0.0, "pre_z": 1.0} - t_all = _reference_frame() - t_sec = secondaries(t_all, Side.reference) - ev, n = R.sec_count_by_event(t_all, t_sec) - assert dict(zip(ev.tolist(), n.tolist())) == {1: 1, 2: 1} + t = secondaries_by_step(_reference_frame(), Side.reference).collect() + assert t["pdg"].to_list() == [22, 22] + # one row per emitting step; the empty-list step drops out entirely + assert [k["_row"] for k in t["step_key"]] == [0, 2] def test_leakage_fraction(): diff --git a/tests/test_catalog.py b/tests/test_catalog.py index 377ac05..be50f27 100644 --- a/tests/test_catalog.py +++ b/tests/test_catalog.py @@ -10,10 +10,10 @@ from giant.analysis.catalog import ( Bundle, PlotSpec, _containment_depths, - _integer_confusion, _ks_statistic, ) from giant.analysis.context import Context, build_context +from giant.analysis.grouping import pdg_label from giant.analysis.sources import RolloutSpec from tests.test_analysis_reduce import _reference_frame, _rollout_frame @@ -160,8 +160,9 @@ def _validate_payload(r, names: list[str]) -> None: # data-dependent edges (event_total_edep), concat-then-mean/std (shower_ # longitudinal), concat-then-max-edge (leakage_fraction), pdg-keyed sum with a # ratio (species_edep_share), a chunkable=False passthrough (router_gating), -# nested sum-merge into a scorecard (marginal_distance_summary), concat-then- -# event-id-join (n_sec_confusion), and concat-then-per-event-derived-quantity +# sum-mergeable-with-a-zero-fill-denominator (sec_count_per_step{,_by_species}), +# nested sum-merge into a scorecard (marginal_distance_summary), and +# concat-then-per-event-derived-quantity # (shower_containment_depth_90, reusing the profile matrix's own merge shape). _CHUNK_EQUIVALENCE_IDS = [ "marginal_edep", @@ -170,9 +171,10 @@ _CHUNK_EQUIVALENCE_IDS = [ "shower_longitudinal", "leakage_fraction", "sec_count_per_species", + "sec_count_per_step", + "sec_count_per_step_by_species", "router_gating", "marginal_distance_summary", - "n_sec_confusion", "shower_containment_depth_90", ] @@ -217,7 +219,7 @@ def test_chunked_matches_unchunked(two_ctx: Context, spec_id: str): # --------------------------------------------------------------------------- -# new (gitea #76) reductions: KS distance, confusion matrix, containment depth +# new (gitea #76) reductions: KS distance and containment depth # --------------------------------------------------------------------------- @@ -228,28 +230,6 @@ def test_ks_statistic(): assert _ks_statistic([10, 0], [0, 0]) == 1.0 # one side empty, other isn't -> maximal mismatch -def test_integer_confusion_matches_event_pairing(): - # true (reference) n_sec = [1, 1]; predicted (rollout) n_sec = [1, 0] - labels, mat = _integer_confusion(np.array([1, 1]), np.array([1, 0])) - assert labels == ["0", "1+"] - assert mat.tolist() == [[0, 0], [1, 1]] # row=true, col=pred - - -def test_integer_confusion_caps_pathological_outliers(): - labels, mat = _integer_confusion(np.array([0, 500]), np.array([0, 0]), max_bins=5) - assert labels[-1] == "4+" - assert mat.shape == (5, 5) - assert mat.sum() == 2 - - -def test_integer_confusion_explicit_cap_overrides_local_range(): - # Even though this pair's own max is 1, an explicit shared cap forces a - # wider (and so cross-rollout-consistent) label set. - labels, mat = _integer_confusion(np.array([1, 1]), np.array([0, 1]), cap=3) - assert labels == ["0", "1", "2", "3+"] - assert mat.shape == (4, 4) - - def test_containment_depths_simple_ramp(): # one event, edep concentrated in the first bin -> 90%/95% containment # depth is the first bin's right edge; a zero-energy event is dropped. @@ -259,17 +239,25 @@ def test_containment_depths_simple_ramp(): assert depths.tolist() == [1.0] -def test_n_sec_confusion_spec(bundle): - spec = get_spec("n_sec_confusion") +def test_sec_count_per_step_counts_empty_steps(bundle): + spec = get_spec("sec_count_per_step") r = spec.finalize([spec.compute_partial(bundle)], bundle.ctx) - assert r.payload["row_labels"] == r.payload["col_labels"] == ["0", "1+"] - assert r.payload["series"]["rollout"] == [[0, 0], [1, 1]] + # reference: 3 steps, two of which emit exactly one secondary + assert r.payload["reference"][:2] == [1, 2] + # rollout: 4 physical steps, one of which emits a single secondary + assert r.payload["series"]["rollout"][:2] == [3, 1] + assert sum(r.payload["reference"]) == 3 -def test_n_sec_confusion_shares_one_cap_across_rollouts(two_bundle): - spec = get_spec("n_sec_confusion") - r = spec.finalize([spec.compute_partial(two_bundle)], two_bundle.ctx) - assert list(r.payload["series"]) == ["flow", "wgan"] - # both rollouts share the same fixture data here, so their matrices (and - # the shared label set) must be identical. - assert r.payload["series"]["flow"] == r.payload["series"]["wgan"] +def test_sec_count_per_step_by_species_zero_row_is_per_species(bundle): + spec = get_spec("sec_count_per_step_by_species") + r = spec.finalize([spec.compute_partial(bundle)], bundle.ctx) + cols = r.payload["col_labels"] + ref = r.payload["reference"] + g = cols.index(pdg_label(22)) + # two reference steps emit one photon each; the third emits none + assert [row[g] for row in ref][:2] == [1, 2] + # every other species column is "no such secondary" on all 3 steps + for j, _ in enumerate(cols): + if j != g: + assert ref[0][j] == 3 and sum(row[j] for row in ref[1:]) == 0 diff --git a/tests/test_checkpoint_io.py b/tests/test_checkpoint_io.py index fcceaf8..71a47cd 100644 --- a/tests/test_checkpoint_io.py +++ b/tests/test_checkpoint_io.py @@ -14,6 +14,7 @@ from giant import config as gconfig from giant.checkpoint_io import ( CheckpointCompatibilityError, InferenceContext, + apply_config_overrides, conditioning_axes, load_for_inference, stage_cfg, @@ -278,3 +279,122 @@ def test_stage_cfg_new_shape_returns_subdict(): def test_stage_cfg_v02_flat_shape_returns_empty_dict(): model_cfg = {"hidden_dim": 32, "n_blocks": 4} 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 diff --git a/tests/test_cli_predict.py b/tests/test_cli_predict.py index 7c8e237..657f5ab 100644 --- a/tests/test_cli_predict.py +++ b/tests/test_cli_predict.py @@ -172,3 +172,43 @@ def test_predict_exits_1_on_checkpoint_missing_model_config(tmp_path): assert result.exit_code == 1 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 diff --git a/tests/test_cli_rollout.py b/tests/test_cli_rollout.py index b5bd2a6..06857fe 100644 --- a/tests/test_cli_rollout.py +++ b/tests/test_cli_rollout.py @@ -31,3 +31,28 @@ def test_rollout_exits_1_on_checkpoint_missing_model_config(tmp_path): assert result.exit_code == 1 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 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( diff --git a/tests/test_config.py b/tests/test_config.py index dabfbd0..acc8fe1 100644 --- a/tests/test_config.py +++ b/tests/test_config.py @@ -171,17 +171,40 @@ def test_n_sec_config_owner_defaults_to_stage2(): def test_n_sec_config_owner_round_trips(): n_sec = gconfig.NSecConfig.from_dict({"mode": "head", "owner": "stage1"}) assert n_sec.owner == "stage1" - assert n_sec.to_dict() == {"mode": "head", "lambda": 0.1, "owner": "stage1", "stop_sampling": "greedy"} + assert n_sec.to_dict() == {"mode": "head", "lambda": 0.1, "owner": "stage1", "sampling": "greedy"} -def test_n_sec_config_stop_sampling_defaults_to_greedy(): - assert gconfig.NSecConfig().stop_sampling == "greedy" +def test_n_sec_config_sampling_defaults_to_greedy(): + assert gconfig.NSecConfig().sampling == "greedy" -def test_n_sec_config_stop_sampling_round_trips(): +def test_n_sec_config_sampling_round_trips(): + n_sec = gconfig.NSecConfig.from_dict({"mode": "stop_token", "sampling": "sample"}) + assert n_sec.sampling == "sample" + assert n_sec.to_dict()["sampling"] == "sample" + + +def test_n_sec_config_stop_sampling_alias_still_honored(): + """gitea #86: stop_sampling was renamed to sampling; old checkpoints' + model_config still carries the old key and must keep working.""" n_sec = gconfig.NSecConfig.from_dict({"mode": "stop_token", "stop_sampling": "sample"}) - assert n_sec.stop_sampling == "sample" - assert n_sec.to_dict()["stop_sampling"] == "sample" + assert n_sec.sampling == "sample" + assert "stop_sampling" not in n_sec.to_dict() + + +def test_n_sec_config_sampling_key_wins_over_stop_sampling_alias(): + n_sec = gconfig.NSecConfig.from_dict({"sampling": "sample", "stop_sampling": "greedy"}) + assert n_sec.sampling == "sample" + + +def test_migrate_config_renames_stop_sampling_key(): + cfg = { + "meta": {"config_version": gconfig.CONFIG_VERSION}, + "stage2_model": {"n_sec": {"stop_sampling": "sample"}}, + } + migrated = gconfig.migrate_config(cfg) + assert gconfig._get_path(migrated, "stage2_model.n_sec.sampling") == "sample" + assert gconfig._get_path(migrated, "stage2_model.n_sec.stop_sampling") is None # --------------------------------------------------------------------------- @@ -851,13 +874,13 @@ def test_validate_config_stop_token_rejected_for_stage1_owner(): assert "stop_token" in str(e) and "owner" in str(e) -def test_validate_config_bad_stop_sampling_rejected(): - cfg = _cfg_with(**{"stage2_model.n_sec.stop_sampling": "bogus"}) +def test_validate_config_bad_n_sec_sampling_rejected(): + cfg = _cfg_with(**{"stage2_model.n_sec.sampling": "bogus"}) try: gconfig.validate_config(cfg) assert False, "expected ValueError" except ValueError as e: - assert "stop_sampling" in str(e) + assert "sampling" in str(e) def test_validate_config_default_precision_is_fp32(): diff --git a/tests/test_model_summary.py b/tests/test_model_summary.py index 5685858..dac918b 100644 --- a/tests/test_model_summary.py +++ b/tests/test_model_summary.py @@ -12,6 +12,8 @@ from giant.cli import app from giant.materials import MATERIAL_PROPERTIES from giant.model.summary import _NOT_BUILD_TIME, _built_modules, _vocab_caveats, summarize_model +INFERENCE_OVERRIDES = gconfig.INFERENCE_OVERRIDES + runner = CliRunner() _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)}" +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): assert "stage1_model.router.n_experts" 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 "trunk" 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 diff --git a/tests/test_render.py b/tests/test_render.py index af77b3c..0082cde 100644 --- a/tests/test_render.py +++ b/tests/test_render.py @@ -304,13 +304,15 @@ def test_render_one_of_each_kind(tmp_path: Path): "hm1", "secondaries", "heatmap", - "Confusion (single rollout)", + "Heatmap (single rollout)", "predicted", { "series": {"flow": [[1, 0], [0, 1]]}, + "reference": [[2, 0], [0, 1]], "row_labels": ["0", "1+"], "col_labels": ["0", "1+"], "cbar_label": "count", + "log_color": True, }, ), ] diff --git a/tests/test_sample.py b/tests/test_sample.py index eb31f89..345a656 100644 --- a/tests/test_sample.py +++ b/tests/test_sample.py @@ -14,6 +14,7 @@ from giant.model.network import ( stage2_trunk_sec_dim, ) from giant.sample import ( + resolve_n_sec, sample_flow, sample_secondaries, sample_secondaries_ar, @@ -72,6 +73,7 @@ def _stage2_ar( mat: int = 2, k_max: int = 5, history: str = "markov", + n_sec_sampling: str = "greedy", ) -> Stage2Autoregressive: particle_cfg, material_cfg = _particle_material_cfg(_conditioning_for(target), emb_dim) return Stage2Autoregressive( @@ -89,6 +91,7 @@ def _stage2_ar( history=history, attn_n_heads=2, attn_n_layers=1, + n_sec_sampling=n_sec_sampling, ).eval() @@ -99,7 +102,7 @@ def _expected_type_dim(target: str, emb_dim: int) -> int: def _stage2_ar_stop_token( target: str, generator: str, - stop_sampling: str = "greedy", + n_sec_sampling: str = "greedy", emb_dim: int = 6, pdg: int = 3, mat: int = 2, @@ -120,7 +123,7 @@ def _stage2_ar_stop_token( particle_type_cfg=ParticleTypeConfig(target=target), build_n_sec_head=False, build_stop_head=True, - stop_sampling=stop_sampling, + n_sec_sampling=n_sec_sampling, ).eval() @@ -267,14 +270,14 @@ def test_sample_secondaries_ar_first_slot_has_no_history(): # ── Stage2Autoregressive: n_sec.mode = "stop_token" ───────────────────────── -@pytest.mark.parametrize("stop_sampling", ["greedy", "sample"]) -def test_sample_secondaries_ar_stop_token_forced_stop_gives_zero_secondaries(stop_sampling): +@pytest.mark.parametrize("n_sec_sampling", ["greedy", "sample"]) +def test_sample_secondaries_ar_stop_token_forced_stop_gives_zero_secondaries(n_sec_sampling): """A stop_head pinned to a large positive logit fires at slot 0 for every row under both policies (greedy: sigmoid(logit) >= 0.5; sample: a Bernoulli draw at sigmoid(logit) ~= 1) — the loop should break before generating any token.""" B, k_max = 4, 5 - decoder = _stage2_ar_stop_token("physical", "flow", stop_sampling=stop_sampling, k_max=k_max) + decoder = _stage2_ar_stop_token("physical", "flow", n_sec_sampling=n_sec_sampling, k_max=k_max) _force_stop_head_logit(decoder, 50.0) cond_cont, cond_cat = _cond(B) stage1_out = torch.randn(B, X_DIM) @@ -283,13 +286,13 @@ def test_sample_secondaries_ar_stop_token_forced_stop_gives_zero_secondaries(sto assert not sec_valid.any() -@pytest.mark.parametrize("stop_sampling", ["greedy", "sample"]) -def test_sample_secondaries_ar_stop_token_forced_never_stop_runs_to_k_max(stop_sampling): +@pytest.mark.parametrize("n_sec_sampling", ["greedy", "sample"]) +def test_sample_secondaries_ar_stop_token_forced_never_stop_runs_to_k_max(n_sec_sampling): """A stop_head pinned to a large negative logit never fires under either policy, so every row is capped at k_max (the safety cap, not a modeling ceiling).""" B, k_max = 4, 5 - decoder = _stage2_ar_stop_token("physical", "flow", stop_sampling=stop_sampling, k_max=k_max) + decoder = _stage2_ar_stop_token("physical", "flow", n_sec_sampling=n_sec_sampling, k_max=k_max) _force_stop_head_logit(decoder, -50.0) cond_cont, cond_cat = _cond(B) stage1_out = torch.randn(B, X_DIM) @@ -336,3 +339,60 @@ def test_sample_secondaries_ar_none_n_sec_pred_without_stop_head_raises(): stage1_out = torch.randn(3, X_DIM) with pytest.raises(AssertionError): sample_secondaries_ar(decoder, cond_cont, cond_cat, stage1_out, None, steps=2) + + +# ── resolve_n_sec: n_sec.mode = "head" sampling policy (gitea #86) ────────── + + +def _force_n_sec_head_bias(decoder: Stage2Autoregressive, bias: torch.Tensor) -> None: + """Zeroes n_sec_head's weights and pins its bias, so predict_n_sec + returns `bias` (broadcast over the batch) as logits regardless of + conditioning — mirrors `_force_stop_head_logit`.""" + assert decoder.n_sec_head is not None + last_linear = decoder.n_sec_head[-1] + with torch.no_grad(): + last_linear.weight.zero_() + last_linear.bias.copy_(bias) + + +@pytest.mark.parametrize("n_sec_sampling", ["greedy", "sample"]) +def test_resolve_n_sec_head_mode_sharply_peaked_logits_pick_dominant_class(n_sec_sampling): + """A logit vector overwhelmingly favoring one class gives the same + answer under both policies — greedy because it's the argmax, sample + because softmax puts ~all mass on it.""" + B, k_max = 8, 5 + decoder = _stage2_ar("physical", "flow", k_max=k_max, n_sec_sampling=n_sec_sampling) + bias = torch.full((k_max + 1,), -50.0) + bias[2] = 50.0 + _force_n_sec_head_bias(decoder, bias) + cond_cont, cond_cat = _cond(B) + stage1_out = torch.randn(B, X_DIM) + n_sec = resolve_n_sec(decoder, decoder, cond_cont, cond_cat, stage1_out, None) + assert n_sec is not None + assert torch.equal(n_sec, torch.full((B,), 2, dtype=torch.long)) + + +def test_resolve_n_sec_head_mode_greedy_is_deterministic_under_flat_logits(): + B, k_max = 32, 5 + decoder = _stage2_ar("physical", "flow", k_max=k_max, n_sec_sampling="greedy") + _force_n_sec_head_bias(decoder, torch.zeros(k_max + 1)) + cond_cont, cond_cat = _cond(B) + stage1_out = torch.randn(B, X_DIM) + n_sec = resolve_n_sec(decoder, decoder, cond_cont, cond_cat, stage1_out, None) + assert n_sec is not None + assert n_sec.unique().numel() == 1 + + +def test_resolve_n_sec_head_mode_sample_varies_under_flat_logits(): + """Under a flat logit vector, a categorical draw across a large batch + should hit more than one class — the whole point of gitea #86: greedy + always collapses to one, sample should not.""" + torch.manual_seed(0) + B, k_max = 256, 5 + decoder = _stage2_ar("physical", "flow", k_max=k_max, n_sec_sampling="sample") + _force_n_sec_head_bias(decoder, torch.zeros(k_max + 1)) + cond_cont, cond_cat = _cond(B) + stage1_out = torch.randn(B, X_DIM) + n_sec = resolve_n_sec(decoder, decoder, cond_cont, cond_cat, stage1_out, None) + assert n_sec is not None + assert n_sec.unique().numel() > 1 diff --git a/uv.lock b/uv.lock index 91760e9..3d05dce 100644 --- a/uv.lock +++ b/uv.lock @@ -713,7 +713,7 @@ wheels = [ [[package]] name = "giant" -version = "0.3.10" +version = "0.3.15" source = { editable = "." } dependencies = [ { name = "numpy" },