9 Commits
Author SHA1 Message Date
gitea-actions 3864e5249b chore: release v0.3.23
CI / Tests (push) Successful in 4m24s
CI / Lint (ruff check) (push) Successful in 1m38s
CI / Format (ruff format) (push) Successful in 1m37s
CI / Type check (ty) (push) Successful in 1m37s
CI / Sync project version with tag (hand-pushed tags only) (push) Skipped
CI / Publish package to Gitea package registry (push) Skipped
CI / Release (bump, changelog, badges, tag) on merge to master (push) Successful in 25s
2026-09-07 15:11:10 +00:00
lars d1fb54fd09 Merge pull request 'feat(ci): add scheduled dependency-bump workflows (Renovate-lite)' (#99) from feature/scheduled-dependency-bump into master
CI / Sync project version with tag (hand-pushed tags only) (push) Skipped
CI / Publish package to Gitea package registry (push) Skipped
CI / Lint (ruff check) (push) Successful in 1m17s
CI / Type check (ty) (push) Successful in 1m16s
CI / Format (ruff format) (push) Successful in 1m21s
CI / Tests (push) Successful in 3m6s
CI / Release (bump, changelog, badges, tag) on merge to master (push) Successful in 1m24s
Reviewed-on: #99
2026-09-07 17:05:39 +02:00
larsandClaude Sonnet 5 acd2350f51 feat(ci): add scheduled dependency-bump workflows (Renovate-lite)
CI / Sync project version with tag (hand-pushed tags only) (pull_request) Skipped
CI / Publish package to Gitea package registry (pull_request) Skipped
CI / Lint (ruff check) (pull_request) Successful in 1m28s
CI / Type check (ty) (pull_request) Successful in 1m28s
CI / Format (ruff format) (pull_request) Successful in 1m28s
CI / Tests (pull_request) Successful in 3m34s
CI / Release (bump, changelog, badges, tag) on merge to master (pull_request) Skipped
Adds two Gitea Actions scheduled workflows, each opening/updating one
standing pull request rather than pushing to master or opening an
issue, so the existing pull_request CI (lint/format/type-check/tests)
gates every change before a human merges:

- deps-lock.yml (weekly, Mondays): `uv lock --upgrade` within the
  existing pyproject.toml constraints, PR'd on chore/uv-lock-upgrade.
- deps-bounds.yml (monthly): raises pyproject.toml upper bounds that
  have fallen behind the latest PyPI release, re-locks, PR'd on
  chore/dep-bounds. torch (pinned <2.4 for portal-machine driver
  support), plotstyle (private index), and the giant[...]
  self-references are permanently excluded.

Both branches are force-pushed fresh from master each run (no history
accumulation, at most one open PR per job), sharing PR-upsert
mechanics in deps-pr.sh. Both support workflow_dispatch for manual
testing. No changes to ci.yml, release-commit.sh, .bumpversion.toml,
or cliff.toml — merging either PR flows through the existing release
job unchanged.

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01KhEer5EWzFU6NJi8hTnDe1
2026-09-07 16:59:41 +02:00
larsandClaude Sonnet 5 e24907862f fix(deps): silence polars explode() empty_as_null deprecation warnings
CI / Sync project version with tag (hand-pushed tags only) (push) Skipped
CI / Publish package to Gitea package registry (push) Skipped
CI / Lint (ruff check) (push) Successful in 1m7s
CI / Format (ruff format) (push) Successful in 1m8s
CI / Type check (ty) (push) Successful in 1m12s
CI / Tests (push) Successful in 2m41s
CI / Release (bump, changelog, badges, tag) on merge to master (push) Successful in 8s
pytest emitted 113 DeprecationWarnings, all from the same source: polars
2.0 changes explode()'s default handling of empty lists from "explode to
null" to "drop the row". Every explode() call site in this repo already
follows the explode with drop_nulls() (or otherwise excludes empty
lists), so the new behavior is what we always wanted — pass
empty_as_null=False explicitly rather than suppressing the warning.

Bump the polars floor from >=1.0 to >=1.43, since the empty_as_null
kwarg doesn't exist before ~1.35.

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_012dE2r7cNo9Lbthh1JUW1NF
2026-09-07 16:25:19 +02:00
gitea-actions 8cdeba088e chore: release v0.3.22
CI / Type check (ty) (push) Successful in 2m54s
CI / Format (ruff format) (push) Successful in 2m56s
CI / Lint (ruff check) (push) Successful in 3m0s
CI / Sync project version with tag (hand-pushed tags only) (push) Has been skipped
CI / Publish package to Gitea package registry (push) Has been skipped
CI / Tests (push) Successful in 3m29s
CI / Release (bump, changelog, badges, tag) on merge to master (push) Successful in 14s
2026-09-07 10:02:33 +00:00
lars 48208e6d18 Merge pull request 'feat(analyze): add paired truth/pred plots from giant predict' (#98) from feature/predict-analysis-plots into master
CI / Sync project version with tag (hand-pushed tags only) (push) Has been skipped
CI / Publish package to Gitea package registry (push) Has been skipped
CI / Format (ruff format) (push) Successful in 42s
CI / Type check (ty) (push) Successful in 50s
CI / Lint (ruff check) (push) Successful in 52s
CI / Tests (push) Successful in 3m9s
CI / Release (bump, changelog, badges, tag) on merge to master (push) Successful in 1m47s
Reviewed-on: #98
2026-09-07 11:57:34 +02:00
larsandClaude Sonnet 5 51790d3e0a feat(predict): enrich YAML sidecar with provenance and timing
CI / Sync project version with tag (hand-pushed tags only) (pull_request) Has been skipped
CI / Publish package to Gitea package registry (pull_request) Has been skipped
CI / Type check (ty) (pull_request) Successful in 48s
CI / Lint (ruff check) (pull_request) Successful in 49s
CI / Format (ruff format) (pull_request) Successful in 49s
CI / Tests (pull_request) Successful in 3m14s
CI / Release (bump, changelog, badges, tag) on merge to master (pull_request) Has been skipped
`giant predict`'s sidecar previously stopped at kind/prediction_id/
output/dataset/checkpoint/timestamp, unlike `giant rollout`'s, which
carries full run provenance (model_config, training_epoch,
training_config, timing, ...) that flows into analysis gallery
metadata. `analyze --prediction` consumed the same thin sidecar, so a
prediction series in an analysis run was nearly unlabeled compared to
its rollout counterparts.

- `_write_prediction_ref` takes an `extra: dict | None` merged into
  the sidecar; `giant rollout` now uses it instead of a
  load/update/rewrite round trip (identical output).
- New `_build_predict_timing`, key-compatible with
  `_build_rollout_timing`, from timers now wrapping predict's setup/
  sample/write phases.
- `giant predict` writes coord, has_truth, schema_version, steps,
  weights, device, batch_size(+auto), row/skip/unknown-pdg counts,
  timing, and the checkpoint's model_config/config_overrides/
  training_epoch/best_val_loss/training_config/training_meta.
- `giant/analysis/condor.py`'s `_PLOT_META_KEYS` forwards the new
  predict-only keys (plus rollout's previously-unforwarded
  config_overrides) into each plot's gallery metadata.yaml.
- Fixes a `ty` regression from the prior commit in
  tests/test_cli_predict.py (Command has no static `.commands`).

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01PpxE9nij3ujg9XcuzvQ26q
2026-09-07 11:52:45 +02:00
larsandClaude Sonnet 5 51f9dad3b0 fix(tests): make predict --truth flag test robust to terminal rendering
CI / Sync project version with tag (hand-pushed tags only) (pull_request) Has been skipped
CI / Publish package to Gitea package registry (pull_request) Has been skipped
CI / Lint (ruff check) (pull_request) Successful in 1m12s
CI / Type check (ty) (pull_request) Failing after 1m12s
CI / Format (ruff format) (pull_request) Successful in 1m12s
CI / Tests (pull_request) Has been skipped
CI / Release (bump, changelog, badges, tag) on merge to master (pull_request) Has been skipped
The --help-text assertion was brittle to CI's terminal width/color
settings (rich can wrap or re-color the flag name mid-word), causing a
false CI failure even though the flag itself is fine. Inspect the
click command's registered option directly instead of parsing rendered
--help output.

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01PpxE9nij3ujg9XcuzvQ26q
2026-09-07 11:35:34 +02:00
larsandClaude Sonnet 5 ac01966a1f feat(analyze): add paired truth/pred plots from giant predict
CI / Sync project version with tag (hand-pushed tags only) (pull_request) Has been skipped
CI / Publish package to Gitea package registry (pull_request) Has been skipped
CI / Lint (ruff check) (pull_request) Successful in 29s
CI / Format (ruff format) (pull_request) Successful in 48s
CI / Type check (ty) (pull_request) Successful in 49s
CI / Tests (pull_request) Failing after 3m5s
CI / Release (bump, changelog, badges, tag) on merge to master (pull_request) Has been skipped
Adds a `prediction` plot family to `giant analyze`, alongside the existing
rollout-vs-reference comparison, and extends `giant predict` to make it
possible:

- `giant predict --coord global` gains schema v3 (`--truth/--no-truth`,
  default on): writes true_* physical columns and true secondary lists
  alongside the predictions, so the output is fully paired.
- New `giant/analysis/prediction.py` builds one canonical true/pred frame
  (`paired_frame`) from either predict coord mode.
- `catalog.py` gains 35 `pred_*` specs: marginals, 2D truth-vs-pred scatter
  (new `heatmap2d` kind), residuals/relative-residuals/calibration profiles,
  KS/bias/RMSE scorecards, n_sec + secondary-species confusion matrices,
  direction-alignment and constraint-violation checks, and a correlation
  delta. Two new Reduced kinds (`paired_hist`, `heatmap2d`) get renderers.
  Every spec degrades to kind="unavailable" with no --prediction given.
- `condor.py`/`cli.py`: `--prediction`/`--prediction-label` on
  `analyze prep`/`submit`, threaded through RunMeta and every compute job.

Full test suite (1162 tests), ruff, and ty all pass.

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01PpxE9nij3ujg9XcuzvQ26q
2026-09-07 11:19:19 +02:00
30 changed files with 3033 additions and 82 deletions
+1 -1
View File
@@ -1,5 +1,5 @@
[tool.bumpversion]
current_version = "0.3.21"
current_version = "0.3.23"
parse = "(?P<major>\\d+)\\.(?P<minor>\\d+)\\.(?P<patch>\\d+)"
serialize = ["{major}.{minor}.{patch}"]
search = "{current_version}"
+184
View File
@@ -0,0 +1,184 @@
#!/usr/bin/env python3
"""Check (and optionally raise) pyproject.toml's dependency upper bounds
against the latest release on PyPI.
Used by the monthly deps-bounds workflow (see deps-bounds-pr.sh): a lockfile
refresh (deps-lock.yml, weekly) only ever moves within the existing
`>=..,<..` constraints, so a stale upper bound never gets touched by that
job. This script closes that gap by proposing (or applying) a raised
ceiling, always as its own reviewable PR — a deliberately excluded package
(torch, pinned to a driver-compatible range; plotstyle, not on PyPI at all)
never gets edited.
Standalone stdlib + `packaging` script (no project deps needed to run it):
uv run --no-project --with packaging python check_dep_bounds.py --report -
uv run --no-project --with packaging python check_dep_bounds.py --apply --report bounds.md
"""
from __future__ import annotations
import argparse
import json
import re
import sys
import tomllib
import urllib.request
from dataclasses import dataclass
from pathlib import Path
from packaging.requirements import Requirement
from packaging.version import Version
ROOT = Path(__file__).resolve().parents[2]
PYPROJECT = ROOT / "pyproject.toml"
# Deliberately never touched by this script, with the reason surfaced in the
# report footer.
SKIP_REASONS = {
"giant": "self-reference (extras-of-extras), not a real upper bound to check",
"torch": "pinned <2.4 deliberately — newer torch needs newer NVIDIA drivers"
" than the shared portal machines have (see CLAUDE.md)",
"plotstyle": "served from the private `larsbogner` index, not PyPI",
}
@dataclass
class Finding:
name: str
old_specifier: str
new_specifier: str
latest: str
def iter_requirement_strings(pyproject: dict) -> list[str]:
project = pyproject["project"]
reqs = list(project.get("dependencies", []))
for group_reqs in project.get("optional-dependencies", {}).values():
reqs.extend(group_reqs)
return reqs
def canonical_requirements(pyproject: dict) -> dict[str, Requirement]:
"""One Requirement per distinct package name (specifiers are expected to
agree across groups — that's true today; a future mismatch would just
mean the last-seen group's specifier gets checked, which is fine for a
monthly advisory script)."""
out: dict[str, Requirement] = {}
for req_str in iter_requirement_strings(pyproject):
req = Requirement(req_str)
out[req.name] = req
return out
def fetch_latest_version(name: str) -> str | None:
url = f"https://pypi.org/pypi/{name}/json"
try:
with urllib.request.urlopen(url, timeout=10) as resp:
data = json.load(resp)
return data["info"]["version"]
except Exception as exc: # noqa: BLE001 - network hiccup: degrade, don't fail the job
print(f"warning: could not fetch latest version for {name}: {exc}", file=sys.stderr)
return None
def next_ceiling(latest: Version) -> str:
"""The upper-bound scheme already used by this repo's own pins:
`<{major+1}` once a package is past 0.x (e.g. numpy>=1.26,<3), but
`<0.{minor+1}` while still pre-1.0 (e.g. ruff>=0.15,<1 -> next would be
<0.17 style if ruff were still 0.x; ty>=0.0.50,<0.1 is the same idea one
level deeper). Only the major (or, pre-1.0, the minor) component of the
latest release matters here — the point is "next breaking-change
boundary", not "exactly latest.patch + epsilon".
"""
if latest.major == 0:
return f"<0.{latest.minor + 1}"
return f"<{latest.major + 1}"
def find_findings(requirements: dict[str, Requirement], skip: dict[str, str]) -> list[Finding]:
findings = []
for name, req in sorted(requirements.items()):
if name in skip:
continue
latest_str = fetch_latest_version(name)
if latest_str is None:
continue
latest = Version(latest_str)
if latest in req.specifier:
continue
new_upper = next_ceiling(latest)
lower_clauses = [str(s) for s in req.specifier if s.operator != "<"]
new_specifier = ",".join([*lower_clauses, new_upper])
findings.append(
Finding(name=name, old_specifier=str(req.specifier), new_specifier=new_specifier, latest=latest_str)
)
return findings
def apply_findings(text: str, findings: list[Finding]) -> str:
"""Textual, in-place replacement of each finding's specifier substring —
deliberately not a TOML round-trip, so comments/formatting/key order in
pyproject.toml survive untouched. Every occurrence is replaced (a name
like polars appears in three separate dependency groups with identical
specifiers, and all of them must move together)."""
for finding in findings:
# Matched by package name rather than the old specifier string
# verbatim: packaging.requirements.Requirement's str(specifier)
# doesn't preserve clause order (e.g. "numpy>=1.26,<3" round-trips
# as "<3,>=1.26"), so an exact-string match on the old requirement
# would rarely hit. subn with no count replaces every occurrence in
# one pass, which is what a multi-group dependency (e.g. polars)
# needs.
new = f'"{finding.name}{finding.new_specifier}"'
# Negative lookahead guards against matching a longer package name
# sharing this one as a prefix (e.g. "numpy" must not match
# "numpydoc>=...").
pattern = re.compile(rf'"{re.escape(finding.name)}(?![\w.-])[^"]*"')
text, n = pattern.subn(new, text)
if n == 0:
print(f"warning: could not locate {finding.name!r} requirement string to rewrite", file=sys.stderr)
return text
def render_report(findings: list[Finding], skip: dict[str, str]) -> str:
lines = []
if findings:
lines.append("| package | old constraint | new constraint | latest on PyPI |")
lines.append("| --- | --- | --- | --- |")
for f in findings:
lines.append(f"| {f.name} | `{f.old_specifier}` | `{f.new_specifier}` | {f.latest} |")
else:
lines.append("All checked dependency upper bounds already cover the latest PyPI release.")
lines.append("")
lines.append("Skipped (never auto-raised):")
for name, reason in sorted(skip.items()):
lines.append(f"- `{name}` — {reason}")
return "\n".join(lines) + "\n"
def main() -> int:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--apply", action="store_true", help="rewrite pyproject.toml in place")
parser.add_argument("--report", default="-", help="write the markdown report here ('-' for stdout, default: -)")
args = parser.parse_args()
pyproject_text = PYPROJECT.read_text()
pyproject = tomllib.loads(pyproject_text)
requirements = canonical_requirements(pyproject)
findings = find_findings(requirements, SKIP_REASONS)
report = render_report(findings, SKIP_REASONS)
if args.report == "-":
print(report, end="")
else:
Path(args.report).write_text(report)
if args.apply and findings:
new_text = apply_findings(pyproject_text, findings)
PYPROJECT.write_text(new_text)
return 0
if __name__ == "__main__":
raise SystemExit(main())
+44
View File
@@ -0,0 +1,44 @@
#!/usr/bin/env bash
# Monthly job: raise pyproject.toml upper bounds that have fallen behind the
# latest PyPI release (see check_dep_bounds.py — torch/plotstyle/giant are
# deliberately excluded there), re-lock, and open/update a PR with the
# result. Separate from the weekly uv.lock-only refresh (deps-lock.yml)
# because this one can legitimately break CI (a new major version), which
# should never block the routine weekly lockfile bump.
#
# Preconditions: repo checked out on master with fetch-depth: 0, uv synced
# (./.gitea/actions/setup), CI_TOKEN/GITHUB_* env set by Gitea Actions.
set -euo pipefail
BRANCH="chore/dep-bounds"
git config user.name "gitea-actions"
git config user.email "actions@git.larsbogner.de"
git checkout -B "$BRANCH" origin/master
BODY_FILE=$(mktemp)
uv run --no-project --with packaging python .gitea/scripts/check_dep_bounds.py \
--apply --report "$BODY_FILE"
if git diff --quiet -- pyproject.toml; then
echo "No upper bounds out of date; nothing to propose"
.gitea/scripts/deps-pr.sh close "$BRANCH"
exit 0
fi
# Re-lock so the PR carries a pyproject + uv.lock that agree.
uv lock
N_RAISED=$(grep -c '^| ' "$BODY_FILE" || true)
# Subtract the header + separator row from the markdown table, if present.
if [ "$N_RAISED" -ge 2 ]; then
N_RAISED=$((N_RAISED - 2))
else
N_RAISED=0
fi
git add pyproject.toml uv.lock
.gitea/scripts/deps-pr.sh open "$BRANCH" \
"chore(deps): raise dependency upper bounds" \
"chore(deps): raise dependency upper bounds ($N_RAISED packages)" \
"$BODY_FILE"
+48
View File
@@ -0,0 +1,48 @@
#!/usr/bin/env bash
# Weekly job: refresh uv.lock against the pinned constraints in
# pyproject.toml (no constraint edits — see check_dep_bounds.py for the
# separate monthly job that raises upper bounds) and open/update a PR with
# the result. See deps-pr.sh for the commit/push/PR-upsert mechanics.
#
# Preconditions: repo checked out on master with fetch-depth: 0, uv synced
# (./.gitea/actions/setup), CI_TOKEN/GITHUB_* env set by Gitea Actions.
set -euo pipefail
BRANCH="chore/uv-lock-upgrade"
git config user.name "gitea-actions"
git config user.email "actions@git.larsbogner.de"
git checkout -B "$BRANCH" origin/master
UPDATES_FILE=$(mktemp)
uv lock --upgrade 2>&1 | tee "$UPDATES_FILE"
if git diff --quiet -- uv.lock; then
echo "uv.lock already up to date; nothing to propose"
.gitea/scripts/deps-pr.sh close "$BRANCH"
exit 0
fi
N_PACKAGES=$(grep -c '^Update ' "$UPDATES_FILE" || true)
BODY_FILE=$(mktemp)
{
if [ "$N_PACKAGES" -gt 0 ]; then
echo "Weekly automated \`uv lock --upgrade\` — updates within the existing"
echo "\`pyproject.toml\` constraints:"
echo
grep '^Update ' "$UPDATES_FILE" | sed 's/^/- /'
else
echo "Weekly automated \`uv lock --upgrade\` refreshed the lockfile (e.g. hash"
echo "or metadata changes) without a visible version bump."
fi
echo
echo "CI on this PR (lint/format/type-check/tests) is the gate; merge normally"
echo "once green, which triggers the usual patch release."
} > "$BODY_FILE"
git add uv.lock
.gitea/scripts/deps-pr.sh open "$BRANCH" \
"chore(deps): weekly uv.lock refresh" \
"chore(deps): weekly uv.lock refresh ($N_PACKAGES packages)" \
"$BODY_FILE"
+84
View File
@@ -0,0 +1,84 @@
#!/usr/bin/env bash
# Shared commit/push/PR-upsert mechanics for the scheduled dependency-bump
# workflows (deps-lock.yml, deps-bounds.yml). Both jobs stage some changes,
# then delegate here to turn them into a standing pull request against
# master — one stable branch per job, force-pushed every run, so the PR
# stays a single commit and a single open proposal across weeks/months
# instead of accumulating history or duplicate PRs.
#
# Usage:
# deps-pr.sh open <branch> <commit-subject> <pr-title> <body-file>
# Commit the currently staged changes, force-push <branch>, and
# create-or-update an open PR from <branch> onto master.
# deps-pr.sh close <branch>
# Close any open PR from <branch> onto master (if one exists) and
# delete the remote branch. Used when a run finds nothing to change.
#
# Preconditions: repo checked out with fetch-depth: 0, git user.name/email
# already configured, GITHUB_SERVER_URL/GITHUB_REPOSITORY/CI_TOKEN set (all
# provided by Gitea Actions), and — for "open" — the changes to publish are
# already `git add`-ed.
set -euo pipefail
MODE="${1:?usage: deps-pr.sh open|close ...}"
BRANCH="${2:?branch name required}"
API="${GITHUB_SERVER_URL}/api/v1/repos/${GITHUB_REPOSITORY}"
AUTH_HEADER="Authorization: token ${CI_TOKEN}"
# Look up the currently open PR (if any) from $BRANCH onto master.
find_open_pr() {
curl -sf -H "$AUTH_HEADER" "${API}/pulls?state=open&base=master" \
| jq -r --arg ref "$BRANCH" '.[] | select(.head.ref == $ref) | .number' \
| head -n1
}
case "$MODE" in
open)
SUBJECT="${3:?commit subject required}"
TITLE="${4:?PR title required}"
BODY_FILE="${5:?PR body file required}"
git commit -F - <<EOF
$SUBJECT
$(cat "$BODY_FILE")
EOF
git push --force origin "HEAD:refs/heads/$BRANCH"
BODY_JSON=$(jq -Rs '.' < "$BODY_FILE")
PR_NUMBER=$(find_open_pr || true)
if [ -n "$PR_NUMBER" ]; then
echo "Updating existing PR #$PR_NUMBER from $BRANCH"
curl -sf -X PATCH -H "$AUTH_HEADER" -H "Content-Type: application/json" \
"${API}/pulls/${PR_NUMBER}" \
-d "{\"title\": $(jq -Rs '.' <<<"$TITLE"), \"body\": ${BODY_JSON}}" \
> /dev/null
else
echo "Opening new PR from $BRANCH"
curl -sf -X POST -H "$AUTH_HEADER" -H "Content-Type: application/json" \
"${API}/pulls" \
-d "{\"head\": \"${BRANCH}\", \"base\": \"master\", \"title\": $(jq -Rs '.' <<<"$TITLE"), \"body\": ${BODY_JSON}}" \
> /dev/null
fi
;;
close)
PR_NUMBER=$(find_open_pr || true)
if [ -n "$PR_NUMBER" ]; then
echo "Closing PR #$PR_NUMBER from $BRANCH (nothing to update)"
curl -sf -X PATCH -H "$AUTH_HEADER" -H "Content-Type: application/json" \
"${API}/pulls/${PR_NUMBER}" -d '{"state": "closed"}' > /dev/null
else
echo "No open PR from $BRANCH to close"
fi
if git ls-remote --exit-code --heads origin "$BRANCH" > /dev/null 2>&1; then
git push origin --delete "$BRANCH" || echo "Could not delete remote branch $BRANCH (already gone?)"
fi
;;
*)
echo "Unknown mode: $MODE (expected 'open' or 'close')" >&2
exit 1
;;
esac
+29
View File
@@ -0,0 +1,29 @@
name: Monthly dependency upper-bound raise
"on":
schedule:
- cron: "0 5 1 * *"
workflow_dispatch: {}
env:
UV_CACHE_DIR: /uv-cache
jobs:
bounds-upgrade:
name: raise stale pyproject ceilings -> PR
runs-on: ubuntu-latest
container:
image: docker.gitea.com/runner-images:ubuntu-latest
volumes:
- /srv/act-runner-cache/uv:/uv-cache
steps:
# CI_TOKEN needs write:repository (same scope the release job uses)
# to push the refresh branch and open/update its PR.
- uses: actions/checkout@v4
with:
token: ${{ secrets.CI_TOKEN }}
fetch-depth: 0
- uses: ./.gitea/actions/setup
- run: .gitea/scripts/deps-bounds-pr.sh
env:
CI_TOKEN: ${{ secrets.CI_TOKEN }}
+29
View File
@@ -0,0 +1,29 @@
name: Weekly dependency lock refresh
"on":
schedule:
- cron: "0 4 * * 1"
workflow_dispatch: {}
env:
UV_CACHE_DIR: /uv-cache
jobs:
lock-upgrade:
name: uv lock --upgrade -> PR
runs-on: ubuntu-latest
container:
image: docker.gitea.com/runner-images:ubuntu-latest
volumes:
- /srv/act-runner-cache/uv:/uv-cache
steps:
# CI_TOKEN needs write:repository (same scope the release job uses)
# to push the refresh branch and open/update its PR.
- uses: actions/checkout@v4
with:
token: ${{ secrets.CI_TOKEN }}
fetch-depth: 0
- uses: ./.gitea/actions/setup
- run: .gitea/scripts/deps-lock-pr.sh
env:
CI_TOKEN: ${{ secrets.CI_TOKEN }}
+18
View File
@@ -1,5 +1,23 @@
# Changelog
## [0.3.23] - 2026-09-07
### Changed
- Fix(deps): silence polars explode() empty_as_null deprecation warnings
- Feat(ci): add scheduled dependency-bump workflows (Renovate-lite)
## [0.3.22] - 2026-09-07
### Changed
- Feat(analyze): add paired truth/pred plots from giant predict
- Fix(tests): make predict --truth flag test robust to terminal rendering
- Feat(predict): enrich YAML sidecar with provenance and timing
## [0.3.21] - 2026-09-07
### Changed
+9 -1
View File
@@ -18,7 +18,7 @@ giant train path/to/steps.parquet --router --router-type energy # MoE routing t
giant model summary --config config.toml # build-only: parameter counts + which config keys actually bite
giant predict path/to/steps.parquet --checkpoint ckpt/best.pt # per-step predictions
giant rollout path/to/steps.parquet --checkpoint ckpt/best.pt --geometry oracle.pkl # full showers
giant analyze submit rollout.yaml --accounting-group cms # parallel rollout-vs-reference analysis on HTCondor
giant analyze submit rollout.yaml --prediction pred.yaml --accounting-group cms # + paired truth/pred plots
giant analyze render <run_dir> --gallery # render PDFs + HTML gallery (run_dir from prep/submit)
giant analyze metrics <train_run_dir> # training-progress plots from metrics.csv
dwarf --help # dataset/tooling CLI: convert, migrate, bump-gen,
@@ -45,6 +45,12 @@ Part of the `dev` extra. Run these periodically (not just at commit time) to cat
Merges to `master` auto-bump the patch version, tag, and update `CHANGELOG.md` via the Gitea workflow in `.gitea/workflows/ci.yml` (bump-my-version + git-cliff). Don't hand-edit the version in `pyproject.toml` or write changelog entries by hand.
Two scheduled Gitea workflows keep dependencies current, each opening/updating one standing pull request (never editing `master` directly, never opening a Gitea issue) so the normal `pull_request` CI (lint/format/type-check/tests) gates every change before a human merges it:
- **`deps-lock.yml`** (weekly, Mondays) — `uv lock --upgrade` within the existing `pyproject.toml` constraints, PR'd on branch `chore/uv-lock-upgrade` (`.gitea/scripts/deps-lock-pr.sh`).
- **`deps-bounds.yml`** (monthly) — raises `pyproject.toml` upper bounds that have fallen behind the latest PyPI release, re-locks, and PRs on branch `chore/dep-bounds` (`.gitea/scripts/deps-bounds-pr.sh`, driving `.gitea/scripts/check_dep_bounds.py`). `torch` (pinned `<2.4` for portal-machine driver support), `plotstyle` (private index, not on PyPI), and the `giant[...]` self-references are never auto-raised — see `check_dep_bounds.py`'s `SKIP_REASONS`.
Both scripts share PR-upsert mechanics in `.gitea/scripts/deps-pr.sh`; both branches are force-pushed fresh from `master` each run rather than accumulated, so at most one open PR exists per job at a time. Both can be triggered by hand via `workflow_dispatch` in the Gitea Actions UI.
## Compute environment
Work on this repo happens across three kinds of machine:
@@ -101,6 +107,8 @@ Secondary energies are a **stick-breaking partition of the `e_sec` budget** from
**Input is one or more `giant rollout` YAML sidecars** (`condor.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** (`<cwd>/analysis_runs/analysis_<id>/` 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 submit a.yaml [b.yaml ...] --chunks N` runs `prep` (recording `N` in `run_meta.json`) then submits 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/<id>__<chunk>.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 <run_dir>` first joins every plot's chunk partials into `reduced/<id>.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 <train_run_dir>` is a separate, unrelated entry point: training-progress plots straight from a run's `metrics.csv`.
**`prediction` family (paired truth/pred, `giant/analysis/prediction.py`):** an optional add-on to the rollout comparison, driven by `--prediction`/`--prediction-label` on `analyze prep`/`submit` (repeatable, same convention as `--label`/rollout YAMLs; series name defaults to the YAML stem for N>1 or `"prediction"` for one). Unlike a rollout (freely generated, no row-level correspondence to truth), a `giant predict` output has a matching truth row for every prediction — a paired, not distributional, comparison. `giant predict --coord global` (schema v3, `--truth` on by default) writes both `pred_*` and `true_*` physical columns plus truth/predicted secondary lists; `--coord local` is the older, always-paired 9D model-space output (`pred_{name}`/`true_{name}` for `LOCAL_TARGET_NAMES`, no secondaries — stage 2 doesn't run there). `paired_frame()` normalizes either coord into one canonical `true_<var>`/`pred_<var>` frame over `PAIRED_VARS` (`step_length`, `edep`, `delta_e`, `post_E`, `cos_scatter`, `cos_travel`), decoding local coord's ALR energy logits the same way `energy_simplex_decode` does. Every prediction in one run must share one `--coord` and the rollouts' `dataset` (`condor.load_prediction_yamls`). The catalog's `prediction` family (`catalog.py`, ids prefixed `pred_`) covers per-variable marginals (new `paired_hist` kind: true dashed / pred solid) and truth-vs-pred 2D histograms (new `heatmap2d` kind, with a y=x guide), residuals/relative-residuals/residual-vs-truth profiles, KS/bias/RMSE scorecards (reusing `heatmap`), `n_sec` and secondary-species confusion matrices (row-normalised `heatmap`), direction-alignment and physical-constraint-violation checks, and a pred/true correlation-matrix delta. Every spec degrades to `kind="unavailable"` when no `--prediction` was given, so a rollout-only run is unaffected. `giant predict` also writes a YAML sidecar next to the checkpoint (`cli.py:_write_prediction_ref`, mirroring `giant rollout`'s) carrying the run's provenance and timing — coord/weights/steps/batch size, row/skip/unknown-PDG counts, a `timing` block, and the checkpoint's `model_config`/`training_epoch`/`training_config` — which `--prediction` consumes the same way `--label` rollout YAMLs are consumed, surfacing those keys into each plot's gallery `metadata.yaml` (`condor.py:_PLOT_META_KEYS`).
**Shower rollout** (`giant/rollout.py`, `giant rollout` CLI): autoregressively steps the two-stage model into a full shower, advancing tracks breadth-first (every sweep steps all active tracks once, in `batch_size` chunks, so many tracks share each forward pass). Each primary post-step becomes the next pre-step, secondaries are pushed as new tracks, and per-step `material`/`layer_id` come from a `GeometryOracle` (`giant/geometry.py`, built via `dwarf build-geometry-oracle`) that learns position → (material, layer_id) from data and flags detector escape by nearest-neighbour distance. Tracks terminate on one of the `TERM_*` reasons in `constants.py` (energy cutoff, max steps, escape, natural end, unknown pdg, max tracks); energy is deposited locally on every stop except escape (leakage), so showers conserve energy by construction. `giant/checkpoint_io.py` is the shared checkpoint → ready-to-run-models path used by both `predict` and `rollout`.
## Roadmap
+2 -2
View File
@@ -10,8 +10,8 @@ calorimeter showers.
[![python](https://img.shields.io/badge/python-3.12%2B-3776AB?logo=python&logoColor=white)](pyproject.toml)
[![torch](https://img.shields.io/badge/torch-2.3.x-EE4C2C?logo=pytorch&logoColor=white)](pyproject.toml)
[![version](https://img.shields.io/badge/version-0.3.21-informational)](CHANGELOG.md)
[![tests](https://img.shields.io/badge/tests-1138%20passing-brightgreen)](tests/)
[![version](https://img.shields.io/badge/version-0.3.23-informational)](CHANGELOG.md)
[![tests](https://img.shields.io/badge/tests-1176%20passing-brightgreen)](tests/)
[![CI](https://git.larsbogner.de/lars/giant/actions/workflows/ci.yml/badge.svg?branch=master)](https://git.larsbogner.de/lars/giant/actions)
[![license](https://img.shields.io/badge/license-unlicensed-lightgrey)](#license)
+8
View File
@@ -13,12 +13,15 @@ re-exported here is plotstyle-free so it runs on a compute worker. Import
from giant.analysis.catalog import build_catalog, catalog_ids, get_spec
from giant.analysis.condor import (
LoadedPrediction,
LoadedRollout,
RunMeta,
SubmitConfig,
compute_one,
compute_reduced,
derive_run_dir,
load_prediction_yaml,
load_prediction_yamls,
load_rollout_yaml,
load_rollout_yamls,
merge_all,
@@ -27,6 +30,7 @@ from giant.analysis.condor import (
write_submit,
)
from giant.analysis.context import Context, build_context
from giant.analysis.prediction import PredictionSpec
from giant.analysis.reduced import Partial, Reduced
from giant.analysis.runtime_estimate import RUNTIME_SAFETY_MARGIN, estimate_runtime_s
from giant.analysis.sources import RolloutSpec, Side
@@ -34,8 +38,10 @@ from giant.analysis.sources import RolloutSpec, Side
__all__ = [
"RUNTIME_SAFETY_MARGIN",
"Context",
"LoadedPrediction",
"LoadedRollout",
"Partial",
"PredictionSpec",
"Reduced",
"RolloutSpec",
"RunMeta",
@@ -49,6 +55,8 @@ __all__ = [
"derive_run_dir",
"estimate_runtime_s",
"get_spec",
"load_prediction_yaml",
"load_prediction_yamls",
"load_rollout_yaml",
"load_rollout_yamls",
"merge_all",
+943 -4
View File
@@ -39,7 +39,8 @@ variable x grouping, secondaries, ...) into concrete specs.
from __future__ import annotations
from collections.abc import Callable
from dataclasses import asdict, dataclass
from dataclasses import asdict, dataclass, field
from typing import cast
import numpy as np
import polars as pl
@@ -52,12 +53,25 @@ from giant.analysis.grouping import (
material_label,
pdg_label,
)
from giant.analysis.prediction import (
PAIRED_SCALARS,
PAIRED_VARS,
PredictionSide,
PredictionSpec,
dir_alignment_expr,
open_prediction,
paired_frame,
paired_secondaries,
paired_vars_for_coord,
)
from giant.analysis.reduce import (
attach_entry_axis,
binned_moments,
depth_expr,
entry_axis,
event_scalars,
hist1d,
hist2d,
leakage_fraction,
profile_finalize,
profile_partial,
@@ -93,6 +107,7 @@ class Bundle:
rollouts: dict[str, RolloutSide] # name -> frames, insertion order = CLI order
t_all: pl.LazyFrame # reference, all rows
t_phys: pl.LazyFrame # reference, physical steps only
predictions: dict[str, PredictionSide] = field(default_factory=dict) # name -> paired frames, CLI order
@classmethod
def open(
@@ -101,12 +116,14 @@ class Bundle:
reference,
ctx: Context,
chunk: tuple[int, int] | None = None,
predictions: list[PredictionSpec] | None = None,
) -> Bundle:
"""Open the reference + every rollout, optionally restricted to one event-disjoint chunk.
"""Open the reference + every rollout + every prediction, optionally
restricted to one event-disjoint chunk.
``chunk = (chunk_index, n_chunks)`` filters every side to
``event_id % n_chunks == chunk_index`` *before* deriving the physical/
secondary views, so every downstream reduction (which is either
secondary/paired views, so every downstream reduction (which is either
row-local or a ``group_by("event_id")``) sees a self-contained,
event-disjoint slice — no cross-chunk lookups are ever needed.
"""
@@ -128,7 +145,24 @@ class Bundle:
type_embedding_l1_dist=rs.type_embedding_l1_dist,
timing=rs.timing,
)
return cls(ctx=ctx, rollouts=sides, t_all=t_all, t_phys=physical_steps(t_all, Side.reference))
pred_sides: dict[str, PredictionSide] = {}
for ps in predictions or []:
opened = open_prediction(ps.source)
lf = opened.lf if pred is None else opened.lf.filter(pred)
pred_sides[ps.name] = PredictionSide(
lf=lf,
paired=paired_frame(lf, opened.coord, opened.has_truth),
coord=opened.coord,
has_truth=opened.has_truth,
checkpoint=ps.checkpoint,
)
return cls(
ctx=ctx,
rollouts=sides,
t_all=t_all,
t_phys=physical_steps(t_all, Side.reference),
predictions=pred_sides,
)
@dataclass
@@ -1076,6 +1110,807 @@ MARGINAL_VARS = ["step_length", "edep", "delta_e", "post_E", "cos_scatter"]
GROUPING_AXES = ["energy", "pdg", "material"]
# ---------------------------------------------------------------------------
# giant predict: paired truth/pred family ("prediction")
#
# Unlike every spec above (rollout series vs one shared reference — an
# *unpaired* distribution comparison), a prediction has a truth row for every
# output row. These specs compare true_<var> against pred_<var> from
# `giant.analysis.prediction.paired_frame` directly — no reference series,
# and (unlike a rollout) a prediction whose --coord is "local" never has
# secondaries or a material/pdg breakdown. `b.predictions` is empty on a run
# with no --prediction given, in which case every spec here degrades to
# kind="unavailable" rather than raising.
# ---------------------------------------------------------------------------
_PRED_UNAVAILABLE_NOTE = "no --prediction given to `analyze prep`/`submit`"
_PRED_TITLE_NAMES = {
"step_length": "Step length",
"edep": "Deposited energy per step",
"delta_e": "Energy loss per step",
"post_E": "Post-step energy",
"cos_scatter": "Scattering cosine (pre_dir . post_dir)",
"cos_travel": "Travel-direction cosine (pre_dir . (post_pos - pre_pos))",
}
_PRED_VAR_LABELS = {
"step_length": "step length [mm]",
"edep": "deposited energy [MeV]",
"delta_e": "energy loss [MeV]",
"post_E": "post-step energy [MeV]",
"cos_scatter": "cos(scattering angle)",
"cos_travel": "cos(travel-direction angle)",
}
def _per_prediction(b: Bundle, fn: Callable[[PredictionSide], object]) -> dict[str, object]:
"""``{name: fn(prediction_side)}`` over every prediction, preserving CLI order."""
return {name: fn(ps) for name, ps in b.predictions.items()}
def _pred_unavailable(
spec_id: str, family: str, title: str, xlabel: str, note: str = _PRED_UNAVAILABLE_NOTE
) -> Reduced:
return Reduced(id=spec_id, family=family, kind="unavailable", title=title, xlabel=xlabel, payload={"note": note})
# ---- marginals + 2D truth-vs-pred -----------------------------------------
def _pred_marginal_partial(b: Bundle, var: str) -> dict:
if not b.predictions or var not in b.ctx.pred_var_ranges:
return {"available": False}
edges = b.ctx.pred_marginal_edges(var)
nb = len(edges) - 1
def _one(ps: PredictionSide) -> dict | None:
if var not in paired_vars_for_coord(ps.coord):
return None
entry = {"pred": _counts(hist1d(ps.paired, pl.col(f"pred_{var}"), edges), 0, nb)}
if ps.has_truth:
entry["true"] = _counts(hist1d(ps.paired, pl.col(f"true_{var}"), edges), 0, nb)
return entry
return {"available": True, "p": _per_prediction(b, _one)}
def _pred_marginal_finalize(parts: list[dict], ctx: Context, var: str) -> Reduced:
if not parts[0]["available"]:
return _pred_unavailable(f"pred_marginal_{var}", "prediction", _PRED_TITLE_NAMES[var], _PRED_VAR_LABELS[var])
edges = ctx.pred_marginal_edges(var)
names = list(parts[0]["p"])
series: dict[str, dict] = {}
for name in names:
entries = [p["p"][name] for p in parts]
if entries[0] is None:
continue
merged = sum_merge(entries)
series[name] = {k: [int(x) for x in v] for k, v in merged.items()}
return Reduced(
id=f"pred_marginal_{var}",
family="prediction",
kind="paired_hist",
title=_PRED_TITLE_NAMES[var],
xlabel=_PRED_VAR_LABELS[var],
payload={"edges": edges.tolist(), "series": series, "log_y": var in PAIRED_SCALARS},
)
def _pred_scatter_partial(b: Bundle, var: str) -> dict:
if not b.predictions or var not in b.ctx.pred_var_ranges:
return {"available": False}
lo, hi = b.ctx.pred_var_ranges[var]
edges = np.linspace(lo, hi, b.ctx.n_marginal_bins + 1)
def _one(ps: PredictionSide) -> list[list[int]] | None:
if var not in paired_vars_for_coord(ps.coord) or not ps.has_truth:
return None
return hist2d(ps.paired, pl.col(f"true_{var}"), pl.col(f"pred_{var}"), edges, edges).tolist()
return {"available": True, "edges": edges.tolist(), "p": _per_prediction(b, _one)}
def _pred_scatter_finalize(parts: list[dict], ctx: Context, var: str) -> Reduced:
if not parts[0]["available"]:
return _pred_unavailable(
f"pred_scatter_{var}", "prediction", f"{_PRED_TITLE_NAMES[var]}: truth vs prediction", _PRED_VAR_LABELS[var]
)
edges = np.asarray(parts[0]["edges"])
names = list(parts[0]["p"])
series: dict[str, list] = {}
for name in names:
mats = [p["p"][name] for p in parts]
if mats[0] is None:
continue
series[name] = np.sum([np.asarray(m, dtype=np.int64) for m in mats], axis=0).tolist()
return Reduced(
id=f"pred_scatter_{var}",
family="prediction",
kind="heatmap2d",
title=f"{_PRED_TITLE_NAMES[var]}: truth vs prediction",
xlabel=f"true {_PRED_VAR_LABELS[var]}",
payload={
"x_edges": edges.tolist(),
"y_edges": edges.tolist(),
"series": series,
"ylabel": f"predicted {_PRED_VAR_LABELS[var]}",
"cbar_label": "step count",
"log_color": True,
"diagonal": True,
},
)
# ---- residuals + calibration -----------------------------------------------
def _pred_residual_partial(b: Bundle, var: str) -> dict:
if not b.predictions or var not in b.ctx.pred_residual_ranges:
return {"available": False}
edges = b.ctx.pred_residual_edges(var)
nb = len(edges) - 1
def _one(ps: PredictionSide) -> list[int] | None:
if var not in paired_vars_for_coord(ps.coord) or not ps.has_truth:
return None
resid = pl.col(f"pred_{var}") - pl.col(f"true_{var}")
return _counts(hist1d(ps.paired, resid, edges), 0, nb)
return {"available": True, "p": _per_prediction(b, _one)}
def _pred_residual_finalize(parts: list[dict], ctx: Context, var: str) -> Reduced:
title, xlabel = f"{_PRED_TITLE_NAMES[var]} residual", f"pred - true {_PRED_VAR_LABELS[var]}"
if not parts[0]["available"]:
return _pred_unavailable(f"pred_residual_{var}", "prediction", title, xlabel)
edges = ctx.pred_residual_edges(var)
names = list(parts[0]["p"])
series: dict[str, list[int]] = {}
for name in names:
entries = [p["p"][name] for p in parts]
if entries[0] is None:
continue
series[name] = [int(x) for x in sum_merge([{"c": e} for e in entries])["c"]]
return Reduced(
id=f"pred_residual_{var}",
family="prediction",
kind="single_hist",
title=title,
xlabel=xlabel,
payload={"edges": edges.tolist(), "series": series, "log_y": True},
)
_REL_RESIDUAL_LO, _REL_RESIDUAL_HI = -2.0, 2.0
def _pred_relative_residual_partial(b: Bundle, var: str) -> dict:
if not b.predictions:
return {"available": False}
edges = np.linspace(_REL_RESIDUAL_LO, _REL_RESIDUAL_HI, b.ctx.n_marginal_bins + 1)
nb = len(edges) - 1
def _one(ps: PredictionSide) -> list[int] | None:
if var not in paired_vars_for_coord(ps.coord) or not ps.has_truth:
return None
rel = (pl.col(f"pred_{var}") - pl.col(f"true_{var}")) / pl.col(f"true_{var}")
return _counts(hist1d(ps.paired, rel, edges), 0, nb)
return {"available": True, "edges": edges.tolist(), "p": _per_prediction(b, _one)}
def _pred_relative_residual_finalize(parts: list[dict], ctx: Context, var: str) -> Reduced:
title, xlabel = f"{_PRED_TITLE_NAMES[var]} relative residual", "(pred - true) / true"
if not parts[0]["available"]:
return _pred_unavailable(f"pred_relative_residual_{var}", "prediction", title, xlabel)
edges = np.asarray(parts[0]["edges"])
names = list(parts[0]["p"])
series: dict[str, list[int]] = {}
for name in names:
entries = [p["p"][name] for p in parts]
if entries[0] is None:
continue
series[name] = [int(x) for x in sum_merge([{"c": e} for e in entries])["c"]]
return Reduced(
id=f"pred_relative_residual_{var}",
family="prediction",
kind="single_hist",
title=title,
xlabel=xlabel,
payload={"edges": edges.tolist(), "series": series, "log_y": True},
)
def _pred_residual_profile_partial(b: Bundle, var: str) -> dict:
if not b.predictions or var not in b.ctx.pred_var_ranges:
return {"available": False}
edges = b.ctx.pred_marginal_edges(var) # bin by truth value
def _one(ps: PredictionSide) -> dict | None:
if var not in paired_vars_for_coord(ps.coord) or not ps.has_truth:
return None
resid = pl.col(f"pred_{var}") - pl.col(f"true_{var}")
return binned_moments(ps.paired, pl.col(f"true_{var}"), resid, edges)
return {"available": True, "p": _per_prediction(b, _one)}
def _moments_to_mean_std(merged: dict[str, list]) -> tuple[np.ndarray, np.ndarray]:
n = np.asarray(merged["n"], dtype=np.float64)
s = np.asarray(merged["sum"], dtype=np.float64)
ss = np.asarray(merged["sumsq"], dtype=np.float64)
with np.errstate(invalid="ignore", divide="ignore"):
mean = np.where(n > 0, s / n, 0.0)
var = np.where(n > 0, ss / n - mean**2, 0.0)
return mean, np.sqrt(np.clip(var, 0.0, None))
def _pred_residual_profile_finalize(parts: list[dict], ctx: Context, var: str) -> Reduced:
title, xlabel = f"{_PRED_TITLE_NAMES[var]} residual vs truth", f"true {_PRED_VAR_LABELS[var]}"
if not parts[0]["available"]:
return _pred_unavailable(f"pred_residual_profile_{var}", "prediction", title, xlabel)
edges = ctx.pred_marginal_edges(var)
names = list(parts[0]["p"])
series: dict[str, dict] = {}
for name in names:
entries = [p["p"][name] for p in parts]
if entries[0] is None:
continue
mean, std = _moments_to_mean_std(sum_merge(entries))
series[name] = {"mean": mean.tolist(), "std": std.tolist()}
return Reduced(
id=f"pred_residual_profile_{var}",
family="prediction",
kind="profile",
title=title,
xlabel=xlabel,
payload={"edges": edges.tolist(), "series": series},
)
_PRED_GROUP_AXES = ("overall", "energy", "pdg", "material")
def _pred_group_expr(lf: pl.LazyFrame, axis: str, ctx: Context) -> pl.Expr:
if axis == "pdg":
return pl.col("pdg")
if axis == "material":
return pl.col("material")
if axis == "energy":
return _energy_group_expr(lf, np.asarray(ctx.energy_edges))
return pl.lit(0, dtype=pl.Int64)
def _pred_group_keys(ctx: Context, axis: str) -> list:
if axis == "pdg":
return list(ctx.top_pdgs)
if axis == "material":
return list(ctx.materials)
if axis == "energy":
return list(range(len(ctx.energy_edges) - 1))
return [0]
def _grouped_moments(lf: pl.LazyFrame, group: pl.Expr, value: pl.Expr) -> dict[str, dict[str, float]]:
"""``{str(group_key): {"n", "sum", "sumsq"}}`` — one streaming pass, sum-mergeable."""
res = (
lf.select(group.alias("_g"), value.alias("_v"))
.drop_nulls(["_g", "_v"])
.group_by("_g")
.agg(pl.len().alias("_n"), pl.col("_v").sum().alias("_s"), (pl.col("_v") ** 2).sum().alias("_ss"))
.collect(engine="streaming")
)
return {str(g): {"n": float(n), "sum": float(s), "sumsq": float(ss)} for g, n, s, ss in res.iter_rows()}
def _pred_scorecard_partial(b: Bundle) -> dict:
"""Per-(var, prediction, axis) group-keyed truth/pred histograms + residual
moments — the shared input to the KS/bias/RMSE scorecards below (three
separate specs, each doing its own finalize math over this one compute)."""
if not b.predictions:
return {"available": False}
out: dict[str, dict] = {}
for var in PAIRED_SCALARS:
if var not in b.ctx.pred_var_ranges:
continue
edges = b.ctx.pred_marginal_edges(var)
def _one(ps: PredictionSide, var=var, edges=edges) -> dict | None:
if var not in paired_vars_for_coord(ps.coord) or not ps.has_truth:
return None
resid = pl.col(f"pred_{var}") - pl.col(f"true_{var}")
per_axis = {}
for axis in _PRED_GROUP_AXES:
group = _pred_group_expr(ps.paired, axis, b.ctx)
per_axis[axis] = {
"true_hist": {
str(k): v.tolist() for k, v in hist1d(ps.paired, pl.col(f"true_{var}"), edges, group).items()
},
"pred_hist": {
str(k): v.tolist() for k, v in hist1d(ps.paired, pl.col(f"pred_{var}"), edges, group).items()
},
"moments": _grouped_moments(ps.paired, group, resid),
"true_moments": _grouped_moments(ps.paired, group, pl.col(f"true_{var}")),
}
return per_axis
out[var] = _per_prediction(b, _one)
return {"available": True, "vars": out}
def _pred_scorecard_matrix(
parts: list[dict], ctx: Context, cell: Callable[[dict, dict, str, str], float]
) -> dict[str, list[list[float]]]:
"""Shared finalize skeleton for the three scorecards: rows = ``PAIRED_SCALARS``,
cols = ``_PRED_GROUP_AXES``, one matrix per prediction. ``cell(true_h_or_moments,
pred_h_or_moments, axis, group_key)`` computes one entry from that axis'
merged group-keyed dict pair (weighted-averaged over the axis' groups)."""
names = list(parts[0]["vars"][next(iter(parts[0]["vars"]))])
matrices: dict[str, list[list[float]]] = {name: [] for name in names}
for var in PAIRED_SCALARS:
var_parts = [p["vars"].get(var) for p in parts]
for name in names:
row: list[float] = []
for axis in _PRED_GROUP_AXES:
raw_entries = [vp[name][axis] if vp is not None and vp[name] is not None else None for vp in var_parts]
if raw_entries[0] is None:
row.append(float("nan"))
continue
# non-None for every chunk: (var, name, axis) availability is
# a static fact of the prediction's coord, not chunk-dependent.
entries = cast("list[dict]", raw_entries)
merged_true_hist = sum_merge([{k: v for k, v in e["true_hist"].items()} for e in entries])
merged_pred_hist = sum_merge([{k: v for k, v in e["pred_hist"].items()} for e in entries])
merged_moments: dict[str, dict[str, float]] = {}
for e in entries:
for k, m in e["moments"].items():
acc = merged_moments.setdefault(k, {"n": 0.0, "sum": 0.0, "sumsq": 0.0})
for f in ("n", "sum", "sumsq"):
acc[f] += m[f]
merged_true_moments: dict[str, dict[str, float]] = {}
for e in entries:
for k, m in e["true_moments"].items():
acc = merged_true_moments.setdefault(k, {"n": 0.0, "sum": 0.0, "sumsq": 0.0})
for f in ("n", "sum", "sumsq"):
acc[f] += m[f]
dists, weights = [], []
for k in _pred_group_keys(ctx, axis):
key = str(k)
if key not in merged_moments:
continue
val = cell(
{"hist": merged_true_hist.get(key), "moments": merged_true_moments.get(key)},
{"hist": merged_pred_hist.get(key), "moments": merged_moments.get(key)},
axis,
key,
)
w = merged_moments[key]["n"]
if w <= 0 or not np.isfinite(val):
continue
dists.append(val)
weights.append(w)
row.append(float(np.average(dists, weights=weights)) if dists else float("nan"))
matrices[name].append(row)
return matrices
def _pred_ks_finalize(parts: list[dict], ctx: Context) -> Reduced:
if not parts[0]["available"] or not parts[0]["vars"]:
return _pred_unavailable(
"pred_ks_summary", "prediction", "Truth/pred distance summary (KS statistic)", "grouping axis"
)
def _cell(true_side: dict, pred_side: dict, axis: str, key: str) -> float:
if pred_side["hist"] is None or true_side["hist"] is None:
return float("nan")
return _ks_statistic(pred_side["hist"], true_side["hist"])
matrices = _pred_scorecard_matrix(parts, ctx, _cell)
return Reduced(
id="pred_ks_summary",
family="prediction",
kind="heatmap",
title="Truth/pred distance summary (KS statistic)",
xlabel="grouping axis",
payload={
"series": matrices,
"row_labels": [_PRED_TITLE_NAMES[v] for v in PAIRED_SCALARS],
"col_labels": list(_PRED_GROUP_AXES),
"ylabel": "variable",
"cbar_label": "KS statistic (0 = identical, 1 = maximal mismatch)",
"vmin": 0.0,
"vmax": 1.0,
},
)
def _pred_bias_finalize(parts: list[dict], ctx: Context) -> Reduced:
if not parts[0]["available"] or not parts[0]["vars"]:
return _pred_unavailable("pred_bias_summary", "prediction", "Relative bias summary", "grouping axis")
def _cell(true_side: dict, pred_side: dict, axis: str, key: str) -> float:
n, s = pred_side["moments"]["n"], pred_side["moments"]["sum"]
if n <= 0:
return float("nan")
bias = s / n
tm = true_side["moments"]
denom = abs(tm["sum"] / tm["n"]) if tm and tm["n"] > 0 else 0.0
return bias / denom if denom > 1e-12 else float("nan")
matrices = _pred_scorecard_matrix(parts, ctx, _cell)
return Reduced(
id="pred_bias_summary",
family="prediction",
kind="heatmap",
title="Relative bias summary (mean(pred - true) / mean|true|)",
xlabel="grouping axis",
payload={
"series": matrices,
"row_labels": [_PRED_TITLE_NAMES[v] for v in PAIRED_SCALARS],
"col_labels": list(_PRED_GROUP_AXES),
"ylabel": "variable",
"cbar_label": "relative bias",
"cmap": "RdBu_r",
"vmin": -0.5,
"vmax": 0.5,
},
)
def _pred_rmse_finalize(parts: list[dict], ctx: Context) -> Reduced:
if not parts[0]["available"] or not parts[0]["vars"]:
return _pred_unavailable("pred_rmse_summary", "prediction", "Relative RMSE summary", "grouping axis")
def _cell(true_side: dict, pred_side: dict, axis: str, key: str) -> float:
n, s, ss = pred_side["moments"]["n"], pred_side["moments"]["sum"], pred_side["moments"]["sumsq"]
if n <= 0:
return float("nan")
rmse = float(np.sqrt(max(ss / n, 0.0)))
tm = true_side["moments"]
denom = abs(tm["sum"] / tm["n"]) if tm and tm["n"] > 0 else 0.0
del s
return rmse / denom if denom > 1e-12 else float("nan")
matrices = _pred_scorecard_matrix(parts, ctx, _cell)
return Reduced(
id="pred_rmse_summary",
family="prediction",
kind="heatmap",
title="Relative RMSE summary",
xlabel="grouping axis",
payload={
"series": matrices,
"row_labels": [_PRED_TITLE_NAMES[v] for v in PAIRED_SCALARS],
"col_labels": list(_PRED_GROUP_AXES),
"ylabel": "variable",
"cbar_label": "relative RMSE",
"vmin": 0.0,
},
)
# ---- confusion matrices -----------------------------------------------
def _n_sec_confusion_edges(cap: int) -> np.ndarray:
return np.arange(-0.5, cap + 1.5)
def _pred_n_sec_confusion_partial(b: Bundle) -> dict:
if not b.predictions:
return {"available": False}
cap = b.ctx.pred_n_sec_cap
edges = _n_sec_confusion_edges(cap)
def _one(ps: PredictionSide) -> list[list[int]] | None:
if ps.coord != "global":
return None
mat = hist2d(ps.paired, pl.col("n_sec").clip(0, cap), pl.col("n_sec_pred").clip(0, cap), edges, edges)
return mat.tolist()
return {"available": True, "p": _per_prediction(b, _one)}
def _row_normalize(mat: np.ndarray) -> np.ndarray:
row_sums = mat.sum(axis=1, keepdims=True)
return np.divide(mat, row_sums, out=np.zeros_like(mat), where=row_sums > 0)
def _pred_n_sec_confusion_finalize(parts: list[dict], ctx: Context) -> Reduced:
title, xlabel = "Secondary-count confusion matrix (row-normalised)", "predicted n_sec"
if not parts[0]["available"]:
return _pred_unavailable("pred_n_sec_confusion", "prediction", title, xlabel)
cap = ctx.pred_n_sec_cap
labels = [str(i) for i in range(cap)] + [f"{cap}+"]
names = list(parts[0]["p"])
series: dict[str, list] = {}
for name in names:
mats = [p["p"][name] for p in parts]
if mats[0] is None:
continue
mat = np.sum([np.asarray(m, dtype=np.float64) for m in mats], axis=0)
series[name] = _row_normalize(mat).tolist()
return Reduced(
id="pred_n_sec_confusion",
family="prediction",
kind="heatmap",
title=title,
xlabel=xlabel,
payload={
"series": series,
"row_labels": labels,
"col_labels": labels,
"ylabel": "true n_sec",
"cbar_label": "fraction of true-count rows",
"vmin": 0.0,
"vmax": 1.0,
},
)
_OTHER_SEC_KEY = "other"
def _sec_species_key_expr(col: str, top_pdgs: list[int]) -> pl.Expr:
return pl.when(pl.col(col).is_in(list(top_pdgs))).then(pl.col(col).cast(pl.Utf8)).otherwise(pl.lit(_OTHER_SEC_KEY))
def _pred_sec_species_confusion_partial(b: Bundle) -> dict:
if not b.predictions or not b.ctx.pred_top_sec_pdgs:
return {"available": False}
top = b.ctx.pred_top_sec_pdgs
def _one(ps: PredictionSide) -> dict | None:
if ps.coord != "global" or not ps.has_truth:
return None
pf = paired_secondaries(ps.lf)
counts = (
pf.select(
_sec_species_key_expr("true_pdg", top).alias("_t"),
_sec_species_key_expr("pred_pdg", top).alias("_p"),
)
.group_by("_t", "_p")
.agg(pl.len().alias("_n"))
.collect(engine="streaming")
)
out: dict[str, dict[str, int]] = {}
for t, p, n in counts.iter_rows():
out.setdefault(t, {})[p] = out.get(t, {}).get(p, 0) + n
return out
return {"available": True, "p": _per_prediction(b, _one)}
def _pred_sec_species_confusion_finalize(parts: list[dict], ctx: Context) -> Reduced:
title, xlabel = "Secondary-species confusion matrix (row-normalised)", "predicted species"
if not parts[0]["available"]:
return _pred_unavailable("pred_sec_species_confusion", "prediction", title, xlabel)
keys = [str(k) for k in ctx.pred_top_sec_pdgs] + [_OTHER_SEC_KEY]
labels = [pdg_label(k) for k in ctx.pred_top_sec_pdgs] + [_OTHER_SEC_KEY]
names = list(parts[0]["p"])
series: dict[str, list] = {}
for name in names:
entries = [p["p"][name] for p in parts]
if entries[0] is None:
continue
merged: dict[str, dict[str, int]] = {}
for e in entries:
for t, row in e.items():
acc = merged.setdefault(t, {})
for p_key, n in row.items():
acc[p_key] = acc.get(p_key, 0) + n
mat = []
for t in keys:
row_counts = merged.get(t, {})
row_vals = [row_counts.get(p_key, 0) for p_key in keys]
total = sum(row_vals)
mat.append([v / total if total > 0 else 0.0 for v in row_vals])
series[name] = mat
return Reduced(
id="pred_sec_species_confusion",
family="prediction",
kind="heatmap",
title=title,
xlabel=xlabel,
payload={
"series": series,
"row_labels": labels,
"col_labels": labels,
"ylabel": "true species",
"cbar_label": "fraction of true-species rows",
"vmin": 0.0,
"vmax": 1.0,
},
)
# ---- physics consistency -----------------------------------------------
def _pred_dir_alignment_partial(b: Bundle, kind: str) -> dict:
if not b.predictions:
return {"available": False}
edges = np.linspace(-1.0, 1.0, b.ctx.n_marginal_bins + 1)
nb = len(edges) - 1
def _one(ps: PredictionSide) -> list[int] | None:
if not ps.has_truth:
return None
return _counts(hist1d(ps.lf, dir_alignment_expr(ps.coord, kind), edges), 0, nb)
return {"available": True, "edges": edges.tolist(), "p": _per_prediction(b, _one)}
def _pred_dir_alignment_finalize(parts: list[dict], ctx: Context, kind: str, title: str) -> Reduced:
xlabel = f"cos(angle) between true and predicted {kind}_dir"
spec_id = f"pred_dir_alignment_{kind}"
if not parts[0]["available"]:
return _pred_unavailable(spec_id, "prediction", title, xlabel)
edges = np.asarray(parts[0]["edges"])
names = list(parts[0]["p"])
series: dict[str, list[int]] = {}
for name in names:
entries = [p["p"][name] for p in parts]
if entries[0] is None:
continue
series[name] = [int(x) for x in sum_merge([{"c": e} for e in entries])["c"]]
return Reduced(
id=spec_id,
family="prediction",
kind="single_hist",
title=title,
xlabel=xlabel,
payload={"edges": edges.tolist(), "series": series, "log_y": True},
)
def _rate_partial(lf: pl.LazyFrame, cond: pl.Expr) -> tuple[int, int]:
res = lf.select(cond.cast(pl.Int64).sum().alias("v"), pl.len().alias("n")).collect(engine="streaming")
return int(res["v"][0]), int(res["n"][0])
_CONSTRAINT_LABELS = ["post_dir_norm", "travel_dir_norm", "step_length_neg", "edep_neg", "delta_e_neg"]
def _pred_constraint_partial(b: Bundle) -> dict:
"""Rate of physical-constraint violations in the *predicted* values.
Direction unit-norm is only a meaningful check for `--coord local`
(`giant predict --coord global` already renormalises before writing, so
it is compliant by construction there — reported as 0/1 rather than
skipped, since "no violations" is still the correct answer).
"""
if not b.predictions:
return {"available": False}
def _one(ps: PredictionSide) -> dict:
out: dict[str, tuple[int, int]] = {}
if ps.coord == "local":
for label, cols in (
("post_dir_norm", ["pred_post_dx", "pred_post_dy", "pred_post_dz"]),
("travel_dir_norm", ["pred_travel_dx", "pred_travel_dy", "pred_travel_dz"]),
):
norm = pl.sum_horizontal([pl.col(c) ** 2 for c in cols]).sqrt()
out[label] = _rate_partial(ps.lf, (norm - 1).abs() > 0.05)
else:
out["post_dir_norm"] = (0, 1)
out["travel_dir_norm"] = (0, 1)
for label, col in (
("step_length_neg", "pred_step_length"),
("edep_neg", "pred_edep"),
("delta_e_neg", "pred_delta_e"),
):
out[label] = _rate_partial(ps.paired, pl.col(col) < 0)
return {k: list(v) for k, v in out.items()}
return {"available": True, "p": _per_prediction(b, _one)}
def _pred_constraint_finalize(parts: list[dict], ctx: Context) -> Reduced:
title, xlabel = "Physical-constraint violation rate", "check"
if not parts[0]["available"]:
return _pred_unavailable("pred_constraint_violations", "prediction", title, xlabel)
names = list(parts[0]["p"])
series: dict[str, list[float]] = {}
for name in names:
entries = [p["p"][name] for p in parts]
rates = []
for check in _CONSTRAINT_LABELS:
v = sum(e[check][0] for e in entries)
n = sum(e[check][1] for e in entries)
rates.append(v / n if n > 0 else 0.0)
series[name] = rates
return Reduced(
id="pred_constraint_violations",
family="prediction",
kind="bar",
title=title,
xlabel=xlabel,
payload={"labels": _CONSTRAINT_LABELS, "series": series, "ylabel": "violation rate"},
)
_CORR_PAIRS = [(a, b) for i, a in enumerate(PAIRED_SCALARS) for b in PAIRED_SCALARS[i + 1 :]]
def _pred_correlation_partial(b: Bundle) -> dict:
if not b.predictions:
return {"available": False}
def _one(ps: PredictionSide) -> dict | None:
if not ps.has_truth:
return None
vs = [v for v in PAIRED_SCALARS if v in paired_vars_for_coord(ps.coord)]
pairs = [(a, c) for a, c in _CORR_PAIRS if a in vs and c in vs]
agg = [pl.len().alias("_n")]
for prefix in ("true", "pred"):
for v in vs:
agg.append(pl.col(f"{prefix}_{v}").sum().alias(f"s_{prefix}_{v}"))
agg.append((pl.col(f"{prefix}_{v}") ** 2).sum().alias(f"ss_{prefix}_{v}"))
for a, c in pairs:
agg.append((pl.col(f"{prefix}_{a}") * pl.col(f"{prefix}_{c}")).sum().alias(f"sxy_{prefix}_{a}_{c}"))
res = ps.paired.select(agg).collect(engine="streaming")
return {"vars": vs, "pairs": pairs, "row": {k: float(res[k][0]) for k in res.columns}}
return {"available": True, "p": _per_prediction(b, _one)}
def _corr_matrix(vs: list[str], pairs: list[tuple[str, str]], row: dict[str, float], prefix: str) -> np.ndarray:
n = row["_n"]
mat = np.eye(len(vs))
idx = {v: i for i, v in enumerate(vs)}
for a, c in pairs:
sa, sc = row[f"s_{prefix}_{a}"], row[f"s_{prefix}_{c}"]
ssa, ssc = row[f"ss_{prefix}_{a}"], row[f"ss_{prefix}_{c}"]
sxy = row[f"sxy_{prefix}_{a}_{c}"]
cov = sxy / n - (sa / n) * (sc / n)
var_a = ssa / n - (sa / n) ** 2
var_c = ssc / n - (sc / n) ** 2
denom = np.sqrt(max(var_a, 0.0) * max(var_c, 0.0))
corr = cov / denom if denom > 1e-12 else float("nan")
mat[idx[a], idx[c]] = mat[idx[c], idx[a]] = corr
return mat
def _pred_correlation_finalize(parts: list[dict], ctx: Context) -> Reduced:
title, xlabel = "Correlation delta (corr(pred) - corr(true))", "variable"
if not parts[0]["available"]:
return _pred_unavailable("pred_correlation_delta", "prediction", title, xlabel)
names = list(parts[0]["p"])
series: dict[str, list] = {}
labels: list[str] = []
for name in names:
entries = [p["p"][name] for p in parts]
if entries[0] is None:
continue
vs, pairs = entries[0]["vars"], [tuple(p) for p in entries[0]["pairs"]]
labels = [_PRED_TITLE_NAMES[v] for v in vs] # same for every prediction: all share one coord (see condor.py)
merged_row: dict[str, float] = {}
for e in entries:
for k, v in e["row"].items():
merged_row[k] = merged_row.get(k, 0.0) + v
delta = _corr_matrix(vs, pairs, merged_row, "pred") - _corr_matrix(vs, pairs, merged_row, "true")
series[name] = delta.tolist()
return Reduced(
id="pred_correlation_delta",
family="prediction",
kind="heatmap",
title=title,
xlabel=xlabel,
payload={
"series": series,
"row_labels": labels,
"col_labels": labels,
"ylabel": "variable",
"cbar_label": "corr(pred) - corr(true)",
"cmap": "RdBu_r",
"vmin": -1.0,
"vmax": 1.0,
},
)
def build_catalog() -> list[PlotSpec]:
"""All concrete plot specs, each with a unique id."""
specs: list[PlotSpec] = []
@@ -1279,6 +2114,110 @@ def build_catalog() -> list[PlotSpec]:
chunkable=False,
),
]
for var in PAIRED_VARS:
specs.append(
PlotSpec(
f"pred_marginal_{var}",
"prediction",
compute_partial=lambda b, v=var: _pred_marginal_partial(b, v),
finalize=lambda parts, ctx, v=var: _pred_marginal_finalize(parts, ctx, v),
)
)
specs.append(
PlotSpec(
f"pred_scatter_{var}",
"prediction",
compute_partial=lambda b, v=var: _pred_scatter_partial(b, v),
finalize=lambda parts, ctx, v=var: _pred_scatter_finalize(parts, ctx, v),
)
)
specs.append(
PlotSpec(
f"pred_residual_{var}",
"prediction",
compute_partial=lambda b, v=var: _pred_residual_partial(b, v),
finalize=lambda parts, ctx, v=var: _pred_residual_finalize(parts, ctx, v),
)
)
for var in PAIRED_SCALARS:
specs.append(
PlotSpec(
f"pred_relative_residual_{var}",
"prediction",
compute_partial=lambda b, v=var: _pred_relative_residual_partial(b, v),
finalize=lambda parts, ctx, v=var: _pred_relative_residual_finalize(parts, ctx, v),
)
)
specs.append(
PlotSpec(
f"pred_residual_profile_{var}",
"prediction",
compute_partial=lambda b, v=var: _pred_residual_profile_partial(b, v),
finalize=lambda parts, ctx, v=var: _pred_residual_profile_finalize(parts, ctx, v),
)
)
specs += [
PlotSpec(
"pred_ks_summary",
"prediction",
compute_partial=_pred_scorecard_partial,
finalize=_pred_ks_finalize,
),
PlotSpec(
"pred_bias_summary",
"prediction",
compute_partial=_pred_scorecard_partial,
finalize=_pred_bias_finalize,
),
PlotSpec(
"pred_rmse_summary",
"prediction",
compute_partial=_pred_scorecard_partial,
finalize=_pred_rmse_finalize,
),
PlotSpec(
"pred_n_sec_confusion",
"prediction",
compute_partial=_pred_n_sec_confusion_partial,
finalize=_pred_n_sec_confusion_finalize,
),
PlotSpec(
"pred_sec_species_confusion",
"prediction",
compute_partial=_pred_sec_species_confusion_partial,
finalize=_pred_sec_species_confusion_finalize,
),
PlotSpec(
"pred_dir_alignment_post",
"prediction",
compute_partial=lambda b: _pred_dir_alignment_partial(b, "post"),
finalize=lambda parts, ctx: _pred_dir_alignment_finalize(
parts, ctx, "post", "Post-direction alignment (true vs predicted)"
),
),
PlotSpec(
"pred_dir_alignment_travel",
"prediction",
compute_partial=lambda b: _pred_dir_alignment_partial(b, "travel"),
finalize=lambda parts, ctx: _pred_dir_alignment_finalize(
parts, ctx, "travel", "Travel-direction alignment (true vs predicted)"
),
),
PlotSpec(
"pred_constraint_violations",
"prediction",
compute_partial=_pred_constraint_partial,
finalize=_pred_constraint_finalize,
),
PlotSpec(
"pred_correlation_delta",
"prediction",
compute_partial=_pred_correlation_partial,
finalize=_pred_correlation_finalize,
),
]
return specs
+125 -13
View File
@@ -56,11 +56,14 @@ import yaml
from giant.analysis.catalog import Bundle, catalog_ids, get_spec
from giant.analysis.context import Context, build_context
from giant.analysis.prediction import PredictionSpec, open_prediction
from giant.analysis.reduced import Partial
from giant.analysis.runtime_estimate import estimate_runtime_s
from giant.analysis.sources import RolloutSpec, Side, open_side
# Keys copied verbatim from a rollout YAML into each plot's gallery metadata.
# Keys copied verbatim from a rollout or prediction YAML into each plot's
# gallery metadata. Rollout-only and predict-only keys both live here —
# `_plot_meta` copies only whichever of these are present in a given YAML.
_PLOT_META_KEYS = (
"prediction_id",
"checkpoint",
@@ -84,10 +87,20 @@ _PLOT_META_KEYS = (
"termination_reason_counts",
"timing",
"model_config",
"config_overrides",
"training_epoch",
"best_val_loss",
"training_config",
"training_meta",
# giant predict only (giant/cli.py's predict command).
"coord",
"has_truth",
"schema_version",
"n_input_rows",
"n_files",
"n_skipped_rows",
"unknown_pdg_counts",
"batch_size_auto",
# Diagnostic — only present when giant rollout ran under
# stage2_model.particle_type.target="embedding" (see giant/cli.py's
# rollout command and giant.rollout.L1DistCollector); absent otherwise,
@@ -163,6 +176,72 @@ def load_rollout_yamls(
return [LoadedRollout(name=n, yaml=y) for n, y in zip(names, yamls)], yamls[0]["dataset"]
def load_prediction_yaml(path: str | Path) -> dict:
"""Load a `giant predict` YAML sidecar, requiring the two file paths."""
d = yaml.safe_load(Path(path).read_text())
for key in ("output", "dataset"):
if key not in d:
raise ValueError(
f"{path} is not a prediction YAML (missing {key!r}); expected the "
"sidecar `giant predict` writes next to the checkpoint"
)
if d.get("kind") not in (None, "prediction"):
raise ValueError(f"{path} has kind={d.get('kind')!r}, not a prediction YAML")
return d
@dataclass
class LoadedPrediction:
"""One prediction YAML plus its resolved series ``name`` and predict ``coord``."""
name: str
yaml: dict
coord: str
def load_prediction_yamls(
paths: Sequence[str | Path], reference: str, labels: Sequence[str] | None = None
) -> list[LoadedPrediction]:
"""Load every prediction YAML, resolve each one's series name, and verify
they're seeded from the same ``reference`` as the rollout(s) and all share
one predict ``--coord`` (direction components mean different things in
the two coords — see ``giant.analysis.prediction``'s module docstring).
Names follow the same convention as ``load_rollout_yamls``: an explicit
``labels[i]`` if given, else the YAML stem for N>1, or ``"prediction"``
for the single-YAML case.
"""
if labels and len(labels) != len(paths):
raise ValueError(
f"--prediction-label given {len(labels)} time(s) but {len(paths)} --prediction YAML(s) were passed"
)
yamls = [load_prediction_yaml(p) for p in paths]
if labels:
names = list(labels)
elif len(paths) == 1:
names = ["prediction"]
else:
names = [Path(p).stem for p in paths]
if len(set(names)) != len(names):
dupes = sorted({n for n in names if names.count(n) > 1})
raise ValueError(f"prediction series names collide: {dupes} — pass --prediction-label to disambiguate")
bad_ref = [(p, y) for p, y in zip(paths, yamls) if str(y["dataset"]) != str(reference)]
if bad_ref:
detail = "\n".join(f" {p}: dataset={y['dataset']!r}" for p, y in bad_ref)
raise ValueError(
f"every --prediction must be seeded from the same reference as the rollout(s) "
f"({reference!r}) — mismatched:\n{detail}"
)
coords = {str(p): open_prediction(y["output"]).coord for p, y in zip(paths, yamls)}
if len(set(coords.values())) > 1:
detail = "\n".join(f" {p}: coord={c!r}" for p, c in coords.items())
raise ValueError(f"every --prediction in one run must share one --coord — got:\n{detail}")
return [LoadedPrediction(name=n, yaml=y, coord=coords[str(p)]) for n, y, p in zip(names, yamls, paths)]
def _run_tag(y: dict) -> str:
rollout = Path(y["output"])
return str(y.get("prediction_id") or rollout.stem)[:8]
@@ -223,6 +302,10 @@ class RunMeta:
# Empty/0 on run directories written before this field existed.
rows_per_chunk: list[int] = field(default_factory=list)
total_rows: int = 0
# `giant predict` inputs (the paired-truth "prediction" family) — same
# shape as `rollouts`. Empty on a run with no --prediction, so old
# run_meta.json files still load.
predictions: list[dict] = field(default_factory=list)
def save(self, path: str | Path) -> None:
Path(path).write_text(json.dumps(self.__dict__, indent=2))
@@ -232,8 +315,13 @@ class RunMeta:
return cls(**json.loads(Path(path).read_text()))
def _rows_per_chunk(rollouts: list[str | Path], reference: str | Path, n_chunks: int) -> list[int]:
"""Combined rollout+reference row count of each ``event_id % n_chunks`` chunk.
def _rows_per_chunk(
rollouts: list[str | Path],
reference: str | Path,
n_chunks: int,
predictions: Sequence[str | Path] = (),
) -> list[int]:
"""Combined rollout+reference+prediction row count of each ``event_id % n_chunks`` chunk.
One cheap streaming ``group_by`` per side (just the ``event_id`` column) —
the sizing input every job's estimated walltime
@@ -249,7 +337,11 @@ def _rows_per_chunk(rollouts: list[str | Path], reference: str | Path, n_chunks:
)
out = [0] * n_chunks
sides = [open_side(reference, Side.reference)] + [open_side(r, Side.rollout) for r in rollouts]
sides = (
[open_side(reference, Side.reference)]
+ [open_side(r, Side.rollout) for r in rollouts]
+ [open_prediction(p).lf for p in predictions]
)
for lf in sides:
df = counts(lf)
for c, n in zip(df["_c"].to_list(), df["n"].to_list()):
@@ -263,26 +355,33 @@ def prep(
n_chunks: int = 1,
default_base: str | Path | None = None,
labels: Sequence[str] | None = None,
prediction_yamls: Sequence[str | Path] = (),
prediction_labels: Sequence[str] | None = None,
**ctx_kwargs,
) -> Path:
"""Read the rollout YAML(s), build the shared context, and lay out the run dir.
"""Read the rollout (+ optional prediction) YAML(s), build the shared
context, and lay out the run dir.
Writes ``shared.json`` + ``run_meta.json`` and returns the run directory.
``n_chunks`` is the run-level chunk count every ``compute-one``/``merge-one``
job reads back out of ``run_meta.json`` (via ``RunMeta.n_chunks``), so it is
resolved once here rather than re-passed (and risking disagreement) at every
later step. See ``derive_run_dir`` for how ``run_dir``/``default_base``
resolve the actual directory, and ``load_rollout_yamls`` for how
``labels``/YAML stems resolve each rollout's series name.
resolve the actual directory, ``load_rollout_yamls`` for how
``labels``/YAML stems resolve each rollout's series name, and
``load_prediction_yamls`` for the same on ``prediction_yamls`` (which,
unlike rollouts, is optional — the ``prediction`` plot family degrades to
``kind="unavailable"`` when it's empty).
Clears any existing ``reduced_partial/``/``reduced/`` from a prior prep of
this same ``run_dir``: partial files carry no record of what context
(``n_chunks``, bin edges, group sets) they were computed under, so
re-prepping with a different ``n_chunks``/``**ctx_kwargs`` (or after the
rollout/reference files changed) would otherwise let ``merge_one`` silently
merge stale partials against the new ``shared.json``.
rollout/reference/prediction files changed) would otherwise let
``merge_one`` silently merge stale partials against the new ``shared.json``.
"""
loaded, reference = load_rollout_yamls(list(rollout_yamls), labels)
loaded_preds = load_prediction_yamls(list(prediction_yamls), reference, prediction_labels)
run_path = derive_run_dir([lr.yaml for lr in loaded], run_dir, default_base=default_base)
run_path.mkdir(parents=True, exist_ok=True)
@@ -292,14 +391,20 @@ def prep(
shutil.rmtree(stale_dir)
rollout_specs = [RolloutSpec(name=lr.name, source=lr.yaml["output"]) for lr in loaded]
ctx = build_context(rollout_specs, reference, **ctx_kwargs)
pred_specs = [PredictionSpec(name=lp.name, source=lp.yaml["output"]) for lp in loaded_preds]
ctx = build_context(rollout_specs, reference, predictions=pred_specs, **ctx_kwargs)
ctx.save(run_path / "shared.json")
rows_per_chunk = _rows_per_chunk([lr.yaml["output"] for lr in loaded], reference, n_chunks)
rows_per_chunk = _rows_per_chunk(
[lr.yaml["output"] for lr in loaded], reference, n_chunks, [lp.yaml["output"] for lp in loaded_preds]
)
rollouts_meta = [
{"name": lr.name, "path": str(lr.yaml["output"]), "plot_meta": _plot_meta(lr.yaml)} for lr in loaded
]
predictions_meta = [
{"name": lp.name, "path": str(lp.yaml["output"]), "plot_meta": _plot_meta(lp.yaml)} for lp in loaded_preds
]
ckpts = ", ".join(Path(lr.yaml.get("checkpoint", "")).name or "rollout" for lr in loaded)
RunMeta(
@@ -310,6 +415,7 @@ def prep(
n_chunks=n_chunks,
rows_per_chunk=rows_per_chunk,
total_rows=sum(rows_per_chunk),
predictions=predictions_meta,
).save(run_path / "run_meta.json")
return run_path
@@ -327,12 +433,15 @@ def compute_reduced(
out: str | Path,
chunk_index: int = 0,
n_chunks: int = 1,
predictions: Sequence[dict] = (),
) -> Path:
"""Core: run one (plot, chunk)'s partial reduction against explicit paths.
``rollouts``: ``[{"name", "path", "checkpoint"?, "type_embedding_l1_dist"?,
"timing"?}, ...]``, one per rollout series (insertion order preserved
through to every plot's ``Reduced.payload["series"]``).
through to every plot's ``Reduced.payload["series"]``). ``predictions``:
``[{"name", "path"}, ...]``, one per ``giant predict`` series (the
``prediction`` family; empty on a run with no ``--prediction``).
Writes a ``Partial`` JSON — the raw, not-yet-merged output of
``PlotSpec.compute_partial`` — never a finished ``Reduced``; ``merge_one``
@@ -357,7 +466,8 @@ def compute_reduced(
)
for r in rollouts
]
bundle = Bundle.open(rollout_specs, reference, ctx, chunk=(chunk_index, effective_n))
pred_specs = [PredictionSpec(name=p["name"], source=p["path"]) for p in predictions]
bundle = Bundle.open(rollout_specs, reference, ctx, chunk=(chunk_index, effective_n), predictions=pred_specs)
partial = Partial(
id=spec_id,
family=spec.family,
@@ -383,6 +493,7 @@ def compute_one(spec_id: str, run_dir: str | Path, chunk_index: int = 0) -> Path
}
for ro in meta.rollouts
]
predictions = [{"name": p["name"], "path": p["path"]} for p in meta.predictions]
return compute_reduced(
spec_id,
rollouts,
@@ -391,6 +502,7 @@ def compute_one(spec_id: str, run_dir: str | Path, chunk_index: int = 0) -> Path
run_path / "reduced_partial" / f"{spec_id}__{chunk_index}.json",
chunk_index=chunk_index,
n_chunks=meta.n_chunks,
predictions=predictions,
)
+64
View File
@@ -20,6 +20,7 @@ import numpy as np
import polars as pl
from giant.analysis.grouping import energy_bin_edges
from giant.analysis.prediction import PredictionSpec, open_prediction, paired_vars_for_coord, prediction_secondaries
from giant.analysis.reduce import (
attach_entry_axis,
depth_expr,
@@ -44,6 +45,12 @@ class Context:
sec_energy_range: tuple[float, float]
n_sec_bins: int
n_events: dict[str, int] = field(default_factory=dict)
# -- giant predict (paired truth/pred comparison) — empty when no
# --prediction was given to `prep`, so old shared.json files still load.
pred_var_ranges: dict[str, tuple[float, float]] = field(default_factory=dict)
pred_residual_ranges: dict[str, tuple[float, float]] = field(default_factory=dict)
pred_n_sec_cap: int = 10
pred_top_sec_pdgs: list[int] = field(default_factory=list)
# -- (de)serialization -------------------------------------------------
def save(self, path: str | Path) -> None:
@@ -54,6 +61,10 @@ class Context:
d = json.loads(Path(path).read_text())
d["var_ranges"] = {k: tuple(v) for k, v in d["var_ranges"].items()}
d["sec_energy_range"] = tuple(d["sec_energy_range"])
if "pred_var_ranges" in d:
d["pred_var_ranges"] = {k: tuple(v) for k, v in d["pred_var_ranges"].items()}
if "pred_residual_ranges" in d:
d["pred_residual_ranges"] = {k: tuple(v) for k, v in d["pred_residual_ranges"].items()}
return cls(**d)
# -- convenience -------------------------------------------------------
@@ -61,6 +72,14 @@ class Context:
lo, hi = self.var_ranges[var]
return np.linspace(lo, hi, self.n_marginal_bins + 1)
def pred_marginal_edges(self, var: str) -> np.ndarray:
lo, hi = self.pred_var_ranges[var]
return np.linspace(lo, hi, self.n_marginal_bins + 1)
def pred_residual_edges(self, var: str) -> np.ndarray:
lo, hi = self.pred_residual_ranges[var]
return np.linspace(lo, hi, self.n_marginal_bins + 1)
_LO_Q, _HI_Q = 0.001, 0.999
@@ -87,10 +106,13 @@ def build_context(
rollouts: list[RolloutSpec],
reference: str | Path | pl.LazyFrame,
*,
predictions: list[PredictionSpec] | None = None,
n_energy_bins: int = 4,
n_marginal_bins: int = 50,
n_sec_bins: int = 40,
top_k_pdg: int = 6,
pred_n_sec_cap: int = 10,
top_k_sec_pdg: int = 8,
sample_rows: int = 1_000_000,
seed: int = 0,
) -> Context:
@@ -165,6 +187,44 @@ def build_context(
}
sec_energy_range = _combined_quantiles([t_se, *r_se.values()], _LO_Q, _HI_Q)
# giant predict: paired truth/pred ranges + residual ranges + secondary
# species vocab, all over the union of every prediction's `paired` frame.
pred_var_ranges: dict[str, tuple[float, float]] = {}
pred_residual_ranges: dict[str, tuple[float, float]] = {}
top_sec_pdgs: list[int] = []
if predictions:
sides = {ps.name: open_prediction(ps.source) for ps in predictions}
present_vars = sorted(set().union(*(paired_vars_for_coord(s.coord) for s in sides.values())))
for var in present_vars:
true_samples, pred_samples, residual_samples = [], [], []
for s in sides.values():
if var not in paired_vars_for_coord(s.coord):
continue
cols = [f"pred_{var}"] + ([f"true_{var}"] if s.has_truth else [])
sample = _row_subsample(s.paired.select(cols), sample_rows, seed).collect(engine="streaming")
pred_samples.append(sample[f"pred_{var}"].to_numpy())
if s.has_truth:
true_samples.append(sample[f"true_{var}"].to_numpy())
residual_samples.append(sample[f"pred_{var}"].to_numpy() - sample[f"true_{var}"].to_numpy())
pred_var_ranges[var] = _combined_quantiles([*true_samples, *pred_samples], _LO_Q, _HI_Q)
if residual_samples:
pred_residual_ranges[var] = _combined_quantiles(residual_samples, _LO_Q, _HI_Q)
sec_pdg_counts: dict[int, int] = {}
for s in sides.values():
if s.coord != "global" or not s.has_truth:
continue
for prefix in ("true", "pred"):
counts = (
prediction_secondaries(s.lf, prefix)
.group_by("pdg")
.agg(pl.len().alias("n"))
.collect(engine="streaming")
)
for pdg, n in zip(counts["pdg"].to_list(), counts["n"].to_list()):
sec_pdg_counts[pdg] = sec_pdg_counts.get(pdg, 0) + n
top_sec_pdgs = [pdg for pdg, _ in sorted(sec_pdg_counts.items(), key=lambda kv: -kv[1])[:top_k_sec_pdg]]
return Context(
n_marginal_bins=n_marginal_bins,
var_ranges=var_ranges,
@@ -173,6 +233,10 @@ def build_context(
materials=materials,
depth_edges=[float(x) for x in depth_edges],
transverse_edges=[float(x) for x in transverse_edges],
pred_var_ranges=pred_var_ranges,
pred_residual_ranges=pred_residual_ranges,
pred_n_sec_cap=pred_n_sec_cap,
pred_top_sec_pdgs=top_sec_pdgs,
sec_energy_range=sec_energy_range,
n_sec_bins=n_sec_bins,
n_events={
+315
View File
@@ -0,0 +1,315 @@
"""Canonical paired truth/prediction LazyFrame for `giant predict` output.
Unlike a `giant rollout` (an unpaired, freely-generated shower), `giant predict`
runs the model once per real pre-step state, so every output row has a
matching truth row — a paired comparison, not a distribution comparison. This
module builds one canonical **paired** LazyFrame per prediction file, in
either coord mode `giant predict` supports, so every catalog spec in the
`prediction` family is coord-agnostic:
event_id, pdg, material, pre_E, n_sec, n_sec_pred,
true_<var>, pred_<var> for var in PAIRED_VARS
`--coord global` (v3+, `--truth` on) already carries physical `true_*`/`pred_*`-
shaped columns directly. `--coord local` carries the raw 9D `true_{name}`/
`pred_{name}` model-space target (`LOCAL_TARGET_NAMES`) instead — its two
ALR energy logits are decoded into physical `edep`/`delta_e` with the same
softmax-against-`pre_E` expressions `giant.data.transforms.energy_simplex_decode`
uses, resurrected from the pre-package-rewrite `giant/analysis.py` (see
`_edep_pl`/`_delta_e_pl`/`_raw_dim_expr` there). Direction components differ in
*meaning* between the two coords (world vs. local frame), so a run must not mix
them — `condor.load_prediction_yamls` enforces one coord across every
prediction in a run.
Secondaries only exist in `--coord global --truth` output (local mode never
samples stage 2); `paired_secondaries` is `None` otherwise, and secondary-based
specs render `kind="unavailable"` instead of raising.
plotstyle-free (runs on HTCondor workers).
"""
from __future__ import annotations
from dataclasses import dataclass
from pathlib import Path
import polars as pl
import pyarrow.parquet as pq
from giant.constants import (
LOCAL_TARGET_NAMES,
PREDICT_COORD_METADATA_KEY,
PREDICT_TRUTH_METADATA_KEY,
ROLLOUT_COORD_VALUE,
)
# The paired scalar/direction variables every coord mode can produce, in
# physical units (mm / MeV) regardless of source coord.
PAIRED_SCALARS: tuple[str, ...] = ("step_length", "edep", "delta_e", "post_E")
PAIRED_VARS: tuple[str, ...] = (*PAIRED_SCALARS, "cos_scatter", "cos_travel")
_LOG_EPS = 1e-6
@dataclass
class PredictionSpec:
"""One named prediction input, as fed to `build_context`/`Bundle.open`.
Mirrors `sources.RolloutSpec`: `name` is the series identity carried
through `payload["series"]` keys, legend labels, and color assignment.
"""
name: str
source: str | Path | pl.LazyFrame
checkpoint: str | None = None
@dataclass
class PredictionSide:
"""One prediction's opened frame + its coord/truth-availability facts."""
lf: pl.LazyFrame # raw scan, chunk-filtered
paired: pl.LazyFrame # canonical paired frame (see module docstring)
coord: str # "global" | "local"
has_truth: bool
checkpoint: str | None = None
def _check_predict_metadata(path: Path) -> tuple[str, bool]:
"""Return `(coord, has_truth)`, raising if `path` isn't predict output.
Distinguishes a predict file from a rollout file (both are tagged with
`PREDICT_COORD_METADATA_KEY`, but a rollout's value is `ROLLOUT_COORD_VALUE`
rather than `"global"`/`"local"`).
"""
metadata = pq.read_schema(path).metadata or {}
coord = metadata.get(PREDICT_COORD_METADATA_KEY.encode())
if coord is None:
raise ValueError(
f"{path} has no {PREDICT_COORD_METADATA_KEY!r} parquet metadata — it wasn't "
"written by `giant predict` (or predates schema tagging)"
)
coord = coord.decode()
if coord == ROLLOUT_COORD_VALUE:
raise ValueError(f"{path} is a `giant rollout` file, not `giant predict` output")
if coord not in ("global", "local"):
raise ValueError(f"{path} has unrecognised predict coord {coord!r}")
# A v1 file predates truth tagging; only --coord local was paired then.
truth_raw = metadata.get(PREDICT_TRUTH_METADATA_KEY.encode())
has_truth = truth_raw.decode() == "1" if truth_raw is not None else coord == "local"
return coord, has_truth
def _edep_pl(prefix: str) -> pl.Expr:
"""Physical edep from `{prefix}_edep_logit`/`{prefix}_sec_logit` + `pre_E`.
Polars equivalent of `energy_simplex_decode(...)[0]` (the deposit
component): a softmax over `[z_edep, z_sec, 0]` times `pre_E`.
"""
z1, z2 = pl.col(f"{prefix}_edep_logit"), pl.col(f"{prefix}_sec_logit")
m = pl.max_horizontal(z1, z2, pl.lit(0.0))
e1, e2, e3 = (z1 - m).exp(), (z2 - m).exp(), (pl.lit(0.0) - m).exp()
return (e1 / (e1 + e2 + e3)) * pl.col("pre_E")
def _delta_e_pl(prefix: str) -> pl.Expr:
"""Physical delta_e (= edep + e_sec = pre_E - post_E) from the ALR logits + pre_E."""
z1, z2 = pl.col(f"{prefix}_edep_logit"), pl.col(f"{prefix}_sec_logit")
m = pl.max_horizontal(z1, z2, pl.lit(0.0))
e1, e2, e3 = (z1 - m).exp(), (z2 - m).exp(), (pl.lit(0.0) - m).exp()
return ((e1 + e2) / (e1 + e2 + e3)) * pl.col("pre_E")
def _local_var_expr(prefix: str, var: str) -> pl.Expr:
"""Physical value of one `PAIRED_VARS` entry from a `--coord local` file."""
if var == "step_length":
return pl.col(f"{prefix}_log_step_length").exp() - _LOG_EPS
if var == "edep":
return _edep_pl(prefix)
if var == "delta_e":
return _delta_e_pl(prefix)
if var == "post_E":
return pl.col("pre_E") - _delta_e_pl(prefix)
if var == "cos_scatter":
dot = pl.sum_horizontal([pl.col(f"{prefix}_post_d{ax}") * pl.col(f"{prefix}_travel_d{ax}") for ax in "xyz"])
return dot.clip(-1.0, 1.0)
raise ValueError(f"{var!r} has no direction-alignment meaning as a solo local-frame variable")
def _g(prefix: str, name: str) -> str:
"""Global-coord column name for `name` under `prefix`.
`giant predict --coord global` writes the *predicted* value under its bare
name (`step_length`, `edep`, `post_dx`, ...) and the truth under a
`true_` prefix (`true_step_length`, ...) — asymmetric, unlike the `local`
coord's symmetric `pred_*`/`true_*` naming.
"""
return name if prefix == "pred" else f"true_{name}"
def _global_var_expr(prefix: str, var: str) -> pl.Expr:
"""Physical value of one `PAIRED_VARS` entry from a `--coord global` file."""
if var == "post_E":
# Not written directly for the prediction (it's implied by energy
# conservation: post_E = pre_E - delta_e); truth carries it as
# true_post_E.
return pl.col("pre_E") - pl.col(_g(prefix, "delta_e")) if prefix == "pred" else pl.col(_g(prefix, "post_E"))
if var == "cos_scatter":
dot = pl.sum_horizontal([pl.col(f"pre_d{ax}") * pl.col(_g(prefix, f"post_d{ax}")) for ax in "xyz"])
return dot.clip(-1.0, 1.0)
if var == "cos_travel":
# travel_dir isn't written by predict (only rollout reconstructs
# post_pos from it) — approximate with the post_pos - pre_pos
# direction instead, which is exactly what travel_dir encodes.
dx = pl.col(_g(prefix, "post_x")) - pl.col("pre_x")
dy = pl.col(_g(prefix, "post_y")) - pl.col("pre_y")
dz = pl.col(_g(prefix, "post_z")) - pl.col("pre_z")
norm = (dx**2 + dy**2 + dz**2).sqrt()
dot = (
pl.col("pre_dx") * dx / (norm + 1e-8)
+ pl.col("pre_dy") * dy / (norm + 1e-8)
+ pl.col("pre_dz") * dz / (norm + 1e-8)
)
return dot.clip(-1.0, 1.0)
return pl.col(_g(prefix, var))
def _var_expr(coord: str, prefix: str, var: str) -> pl.Expr:
# `cos_travel` is excluded for `coord == "local"` by `paired_vars_for_coord`
# (predict never reconstructs post_pos/travel_dir there), so this only
# ever sees local-representable vars on that path.
if coord == "local":
return _local_var_expr(prefix, var)
return _global_var_expr(prefix, var)
def paired_vars_for_coord(coord: str) -> tuple[str, ...]:
"""The `PAIRED_VARS` a given coord mode can actually produce.
`cos_travel` needs a reconstructed `travel_dir`/`post_pos`, which
`--coord local` predict output never has (stage 2 doesn't run there) —
so local-coord predictions drop it rather than emit a meaningless value.
"""
if coord == "local":
return PAIRED_SCALARS + ("cos_scatter",)
return PAIRED_VARS
def dir_alignment_expr(coord: str, kind: str) -> pl.Expr:
"""cos angle between the true and predicted direction vector (raw, not paired).
`kind="post"` compares `post_dir`; `kind="travel"` compares the
post_pos-implied travel direction. Reads the *raw* opened frame
(`PredictionSide.lf`), not `paired` — direction components aren't part of
`PAIRED_VARS` (only their two scattering cosines are), so this stays a
separate helper.
"""
if coord == "local":
prefix_dim = "post_d" if kind == "post" else "travel_d"
true_v = [pl.col(f"true_{prefix_dim}{ax}") for ax in "xyz"]
pred_v = [pl.col(f"pred_{prefix_dim}{ax}") for ax in "xyz"]
elif kind == "post":
true_v = [pl.col(f"true_post_d{ax}") for ax in "xyz"]
pred_v = [pl.col(f"post_d{ax}") for ax in "xyz"] # unprefixed: see paired_frame's _g
else:
true_v = [pl.col(f"true_post_{ax}") - pl.col(f"pre_{ax}") for ax in "xyz"]
pred_v = [pl.col(f"post_{ax}") - pl.col(f"pre_{ax}") for ax in "xyz"]
dot = pl.sum_horizontal([t * p for t, p in zip(true_v, pred_v)])
true_norm = pl.sum_horizontal([t**2 for t in true_v]).sqrt()
pred_norm = pl.sum_horizontal([p**2 for p in pred_v]).sqrt()
return (dot / (true_norm * pred_norm + 1e-8)).clip(-1.0, 1.0)
def paired_frame(lf: pl.LazyFrame, coord: str, has_truth: bool) -> pl.LazyFrame:
"""Canonical `event_id, pdg, material, pre_E, n_sec, n_sec_pred, true_*, pred_*` frame."""
schema = lf.collect_schema().names()
cols = [
"event_id",
"pdg",
"pre_E",
"material",
"n_sec",
pl.col("n_sec_pred") if "n_sec_pred" in schema else pl.lit(None, dtype=pl.Int64).alias("n_sec_pred"),
]
for var in paired_vars_for_coord(coord):
cols.append(_var_expr(coord, "pred", var).alias(f"pred_{var}"))
if has_truth:
cols.append(_var_expr(coord, "true", var).alias(f"true_{var}"))
return lf.select(cols)
def open_prediction(source: str | Path | pl.LazyFrame) -> PredictionSide:
"""Lazily scan one prediction file, verifying its predict tag."""
if isinstance(source, pl.LazyFrame):
lf = source.with_columns(pl.col("pdg").cast(pl.Int64))
schema = lf.collect_schema().names()
coord = "local" if "pred_log_step_length" in schema else "global"
has_truth = f"true_{LOCAL_TARGET_NAMES[0]}" in schema or "true_step_length" in schema
else:
path = Path(source)
coord, has_truth = _check_predict_metadata(path)
lf = pl.scan_parquet(path).with_columns(pl.col("pdg").cast(pl.Int64))
return PredictionSide(
lf=lf,
paired=paired_frame(lf, coord, has_truth),
coord=coord,
has_truth=has_truth,
)
# ---------------------------------------------------------------------------
# Secondaries (global + truth only)
# ---------------------------------------------------------------------------
def prediction_secondaries(lf: pl.LazyFrame, prefix: str) -> pl.LazyFrame:
"""One row per secondary from the true/predicted `sec_*_list` columns.
Canonical columns: `event_id, energy, pdg, sdx, sdy, sdz` — same shape as
`sources.secondaries`'s reference-side branch. `prefix` is `"true"` or
`"pred"`; matches `giant predict --coord global`'s asymmetric naming (see
`_g`) — the predicted lists are unprefixed (`sec_E_list`, ...), only the
truth ones carry `true_` (`true_sec_E_list`, ...).
"""
col_prefix = "" if prefix == "pred" else "true_"
lists = [f"{col_prefix}sec_{c}_list" for c in ("E", "pdg", "dx", "dy", "dz")]
return (
lf.select("event_id", *lists)
.explode(lists, empty_as_null=False)
.drop_nulls(lists[0])
.select(
"event_id",
pl.col(lists[0]).alias("energy"),
pl.col(lists[1]).cast(pl.Int64).alias("pdg"),
pl.col(lists[2]).alias("sdx"),
pl.col(lists[3]).alias("sdy"),
pl.col(lists[4]).alias("sdz"),
)
)
def paired_secondaries(lf: pl.LazyFrame) -> pl.LazyFrame:
"""True/predicted secondary PDG pairs, aligned by descending-energy rank.
Stage 2 emits secondaries in descending-energy order (`network.md`/
`giant/model/models.py`'s autoregressive decoder), so the natural
per-step alignment between the true and predicted secondary lists is
positional: rank `i` of one list vs. rank `i` of the other, for
`i < min(n_sec, n_sec_pred)`. Requires `--coord global --truth`.
"""
return (
lf.select("true_sec_pdg_list", "sec_pdg_list")
.with_row_index("_row")
.with_columns(
pl.col("true_sec_pdg_list").list.len().alias("_n_true"),
pl.col("sec_pdg_list").list.len().alias("_n_pred"),
)
.with_columns(pl.min_horizontal("_n_true", "_n_pred").alias("_n_paired"))
.filter(pl.col("_n_paired") > 0)
.with_columns(pl.int_ranges(0, pl.col("_n_paired")).alias("_rank"))
.explode("_rank", empty_as_null=False)
.select(
pl.col("true_sec_pdg_list").list.get(pl.col("_rank")).cast(pl.Int64).alias("true_pdg"),
pl.col("sec_pdg_list").list.get(pl.col("_rank")).cast(pl.Int64).alias("pred_pdg"),
)
)
+58
View File
@@ -70,6 +70,64 @@ def hist1d(
return out
def hist2d(
lf: pl.LazyFrame,
x: pl.Expr,
y: pl.Expr,
x_edges: np.ndarray,
y_edges: np.ndarray,
) -> np.ndarray:
"""Streaming 2D histogram of `(x, y)` over fixed uniform edges.
One `group_by([_bx, _by]).len()` pass; returns the full `(len(x_edges)-1,
len(y_edges)-1)` int64 count matrix (row = x bin, col = y bin) — small
enough (a truth-vs-pred scatter has at most a few thousand cells) to
materialize whole, unlike `hist1d`'s per-group dict.
"""
x_lo, x_hi, x_n = float(x_edges[0]), float(x_edges[-1]), len(x_edges) - 1
y_lo, y_hi, y_n = float(y_edges[0]), float(y_edges[-1]), len(y_edges) - 1
res = (
lf.select(_bin_expr(x, x_lo, x_hi, x_n).alias("_bx"), _bin_expr(y, y_lo, y_hi, y_n).alias("_by"))
.drop_nulls(["_bx", "_by"])
.group_by("_bx", "_by")
.agg(pl.len().alias("_n"))
.collect(engine="streaming")
)
mat = np.zeros((x_n, y_n), dtype=np.int64)
mat[res["_bx"].to_numpy(), res["_by"].to_numpy()] = res["_n"].to_numpy()
return mat
def binned_moments(
lf: pl.LazyFrame,
bin_value: pl.Expr,
agg_value: pl.Expr,
edges: np.ndarray,
) -> dict[str, list]:
"""Per-bin ``(n, sum, sumsq)`` of ``agg_value``, binned by ``bin_value`` over fixed edges.
One streaming `group_by` pass; sum-mergeable across chunks the same way
`hist1d` counts are — elementwise-summing `n`/`sum`/`sumsq` per bin across
chunks reconstructs the moments of the full merged data, from which
`finalize` derives mean/std (``mean = sum/n``,
``std = sqrt(sumsq/n - mean**2)``).
"""
lo, hi, nbins = float(edges[0]), float(edges[-1]), len(edges) - 1
res = (
lf.select(_bin_expr(bin_value, lo, hi, nbins).alias("_b"), agg_value.alias("_v"))
.drop_nulls(["_b", "_v"])
.group_by("_b")
.agg(pl.len().alias("_n"), pl.col("_v").sum().alias("_s"), (pl.col("_v") ** 2).sum().alias("_ss"))
.collect(engine="streaming")
)
n = np.zeros(nbins, dtype=np.int64)
s = np.zeros(nbins, dtype=np.float64)
ss = np.zeros(nbins, dtype=np.float64)
for b_, nn, ssum, sqsum in res.iter_rows():
n[b_], s[b_], ss[b_] = nn, ssum, sqsum
return {"n": n.tolist(), "sum": s.tolist(), "sumsq": ss.tolist()}
def sum_merge(dicts: list[dict[str, Any]]) -> dict[str, Any]:
"""Elementwise-sum a list of sum-mergeable count/total dicts (JSON-safe keys).
+7 -2
View File
@@ -27,8 +27,13 @@ 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)
# "unavailable" plot not applicable to this run (e.g. no MoE checkpoint)
# distance scorecard) or per prediction (a confusion matrix)
# "paired_hist" per-prediction true/pred density histogram over shared
# edges (giant predict's paired truth, not a rollout)
# "heatmap2d" numeric x/y-binned true-vs-pred count matrix + colorbar,
# one panel per prediction, with a y=x diagonal guide
# "unavailable" plot not applicable to this run (e.g. no MoE checkpoint,
# or no --prediction given)
@dataclass
+74
View File
@@ -422,6 +422,78 @@ def _render_heatmap(r: Reduced, params: dict):
return fig
def _render_paired_hist(r: Reduced, params: dict):
"""`giant predict`'s paired truth/pred density histogram (see
`giant.analysis.prediction`) — unlike `_render_overlay`, there's no single
shared reference: each prediction carries its own truth. A lone prediction
draws its truth in the reference ink so a single-series run reads exactly
like an `overlay_hist` figure; two-or-more predictions each get their own
color, pred solid / true dashed, so a same-colored pair is directly
comparable.
"""
edges = np.asarray(r.payload["edges"])
series = r.payload.get("series", {})
fig, ax = ps.new_figure("thesis-single", title=r.title, params=params)
solo = len(series) == 1
for i, (name, entry) in enumerate(series.items()):
color = _ref_color() if solo else ps.get_color(i)
if "true" in entry:
true_label = _REFERENCE_LABEL if solo else f"{name} (true)"
ax.stairs(_density(entry["true"], edges), edges, label=true_label, color=color, linestyle="--")
pred_color = ps.get_color(i)
pred_label = name if solo else f"{name} (pred)"
ax.stairs(_density(entry["pred"], edges), edges, label=pred_label, color=pred_color)
if r.payload.get("log_y"):
ax.set_yscale("log")
ax.set_xlabel(r.xlabel)
ax.set_ylabel("density")
ps.style_legend(ax, title="source")
return fig
def _render_heatmap2d(r: Reduced, params: dict):
"""Numeric truth-vs-pred 2D histogram, one panel per prediction, with an
optional y=x guide line — the direct analogue of `_render_heatmap` for
continuous (not categorical) axes."""
x_edges = np.asarray(r.payload["x_edges"])
y_edges = np.asarray(r.payload["y_edges"])
series = r.payload["series"]
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,
params=params,
nrows=1,
ncols=len(names),
squeeze=False,
)
flat = axes.ravel()
im = None
for ax, name in zip(flat, names):
mat = np.asarray(series[name], dtype=float)
im = ax.pcolormesh(
x_edges,
y_edges,
mat.T,
cmap=r.payload.get("cmap", "viridis"),
norm=norm,
vmin=None if norm else r.payload.get("vmin"),
vmax=None if norm else r.payload.get("vmax"),
)
if r.payload.get("diagonal"):
lo, hi = max(x_edges[0], y_edges[0]), min(x_edges[-1], y_edges[-1])
ax.plot([lo, hi], [lo, hi], color=_ref_color(), linestyle="--", linewidth=1, label="y = x")
ax.set_xlabel(r.xlabel)
if len(names) > 1:
ax.set_title(name, fontsize=8)
flat[0].set_ylabel(r.payload.get("ylabel", ""))
if r.payload.get("diagonal"):
ps.style_legend(flat[0], title="guide")
fig.colorbar(im, ax=list(flat), label=r.payload.get("cbar_label", "count"))
return fig
def _render_unavailable(r: Reduced, params: dict):
fig, ax = ps.new_figure("thesis-single", title=r.title, params=params)
ax.axis("off")
@@ -448,6 +520,8 @@ _RENDERERS = {
"router_share": _render_router_share,
"router_specialization": _render_router_specialization,
"heatmap": _render_heatmap,
"paired_hist": _render_paired_hist,
"heatmap2d": _render_heatmap2d,
"unavailable": _render_unavailable,
}
+2 -2
View File
@@ -230,7 +230,7 @@ def secondaries(lf: pl.LazyFrame, side: Side) -> pl.LazyFrame:
lists = ["sec_E_list", "sec_pdg_list", "sec_dx_list", "sec_dy_list", "sec_dz_list"]
return (
lf.select("event_id", *lists)
.explode(lists)
.explode(lists, empty_as_null=False)
.drop_nulls("sec_E_list")
.select(
"event_id",
@@ -268,7 +268,7 @@ def secondaries_by_step(lf: pl.LazyFrame, side: Side) -> pl.LazyFrame:
return (
lf.select("sec_pdg_list")
.with_row_index("_row")
.explode("sec_pdg_list")
.explode("sec_pdg_list", empty_as_null=False)
.drop_nulls("sec_pdg_list")
.select(pl.struct("_row").alias("step_key"), pl.col("sec_pdg_list").cast(pl.Int64).alias("pdg"))
)
+244 -44
View File
@@ -20,6 +20,7 @@ from giant.constants import (
PREDICT_COORD_METADATA_KEY,
PREDICT_SCHEMA_VERSION,
PREDICT_SCHEMA_VERSION_KEY,
PREDICT_TRUTH_METADATA_KEY,
ROLLOUT_COORD_VALUE,
)
@@ -186,11 +187,18 @@ def _write_prediction_ref(
out: Path,
dataset_path: Path,
comment: str | None = None,
extra: dict | None = None,
) -> Path:
"""Write a YAML sidecar in the checkpoint directory and return its path."""
"""Write a YAML sidecar in the checkpoint directory and return its path.
``extra`` is merged in after the base fields (e.g. `giant rollout`'s
provenance/timing block, or `giant predict`'s) — callers that don't pass
it get exactly today's thin sidecar.
"""
import yaml
ref = {
"kind": "prediction",
"prediction_id": pred_uuid,
"output": str(out),
"dataset": str(dataset_path),
@@ -199,6 +207,8 @@ def _write_prediction_ref(
}
if comment is not None:
ref["comment"] = comment
if extra:
ref.update(extra)
ref_path = checkpoint.parent / f"{pred_uuid}.yaml"
ref_path.write_text(yaml.dump(ref, default_flow_style=False, sort_keys=False))
return ref_path
@@ -243,6 +253,41 @@ def _build_rollout_timing(
}
def _build_predict_timing(
*,
setup_s: float,
predict_s: float,
write_s: float,
n_rows: int,
device: str,
torch_threads: int,
) -> dict:
"""Assemble ``giant predict``'s ``timing`` sidecar block.
Keys are deliberately compatible with ``_build_rollout_timing``'s (same
names for the quantities both commands have) so a gallery's ``timing``
metadata renders the same way whether the series came from a rollout or
a prediction. There's no ``n_physical_rows``/``ms_per_event`` here:
unlike a rollout, `giant predict` never emits synthetic termination rows
(one output row per input step) and doesn't work in whole showers/events
so ``us_per_step`` is already directly comparable to a rollout's and to
``giant.analysis.geant4_reference``'s per-step Geant4 measurement.
"""
sample_s = predict_s - write_s
return {
"setup_s": setup_s,
"predict_s": predict_s,
"write_s": write_s,
"sample_s": sample_s,
"n_rows": n_rows,
"us_per_step": (sample_s / n_rows * 1e6) if n_rows else None,
"write_us_per_step": (write_s / n_rows * 1e6) if n_rows else None,
"rows_per_s": (n_rows / predict_s) if predict_s else None,
"device": device,
"torch_threads": torch_threads,
}
@app.callback()
def _main() -> None:
"""GIANT — Geant4 step-function surrogate."""
@@ -1020,9 +1065,21 @@ def predict(
typer.Option(
"--out",
"-o",
help="Output parquet path (default: <data>_predicted[_local].parquet)",
help="Output parquet path (default: a UUID-named file under /ceph's central "
"predictions store if --data is under /ceph, else a sibling of --data)",
),
] = None,
truth: Annotated[
bool,
typer.Option(
"--truth/--no-truth",
help="--coord global only: also read and write ground-truth post-step + "
"secondary columns (true_step_length, true_edep, true_sec_*_list, ...) "
"alongside the predictions, at the cost of reading full row-groups instead "
"of conditioning columns only. Ignored for --coord local, which is always "
"paired. Default: on.",
),
] = True,
batch_size: Annotated[
str,
typer.Option(
@@ -1071,6 +1128,8 @@ def predict(
] = None,
) -> None:
"""Run trained model on a parquet file and save predictions."""
import time
import numpy as np
import pyarrow as pa
import pyarrow.parquet as pq
@@ -1090,6 +1149,8 @@ def predict(
from giant.rollout import decode_secondary_identity
from giant.sample import resolve_n_sec, sample_stage1, sample_stage2
_t_setup_start = time.perf_counter()
batch_size_auto = False
batch_size_value: int | None = None
if batch_size.strip().lower() == "auto":
@@ -1161,8 +1222,17 @@ def predict(
unknown_pdg_counts: Counter[int] = Counter()
total_rows = sum(pq.ParquetFile(path).metadata.num_rows for path in files)
training_cfg = gconfig.load_checkpoint_config(checkpoint)
_write_s = 0.0
_setup_s = time.perf_counter() - _t_setup_start
# --coord local always needs full row-groups (it's paired against the 9D
# target); --coord global only needs them when --truth is requested —
# otherwise the cheaper conditioning-only read is used.
write_truth = coord == Coord.global_ and truth
def chunk_iter(path: Path, offset: int):
if coord == Coord.local:
if coord == Coord.local or write_truth:
return iter_file_chunks(path, offset=offset, k_max=stage2_k_max)
return iter_cond_chunks(path, offset=offset)
@@ -1177,7 +1247,7 @@ def predict(
return {k: np.concatenate([a[k], b[k]], axis=0) for k in a}
def _process(piece: dict[str, np.ndarray]) -> None:
nonlocal writer, total
nonlocal writer, total, _write_s
if coord == Coord.local:
feats = build_features(
@@ -1290,54 +1360,95 @@ def predict(
sec_dy_list = [sec_dir_world[i, :n, 1].tolist() for i, n in enumerate(n_sec_pred_np)]
sec_dz_list = [sec_dir_world[i, :n, 2].tolist() for i, n in enumerate(n_sec_pred_np)]
table = pa.table(
{
"event_id": piece["event_id"],
"pdg": piece["pdg"],
"pre_x": piece["pre_pos"][:, 0],
"pre_y": piece["pre_pos"][:, 1],
"pre_z": piece["pre_pos"][:, 2],
"pre_E": piece["pre_E"],
"pre_dx": piece["pre_dir"][:, 0],
"pre_dy": piece["pre_dir"][:, 1],
"pre_dz": piece["pre_dir"][:, 2],
"material": piece["material"],
"layer_id": piece["layer_id"],
"n_sec": piece["n_sec"],
"n_sec_pred": n_sec_pred_np,
"step_length": step_length,
"delta_e": delta_e,
"edep": edep,
"post_dx": post_dir_world[:, 0],
"post_dy": post_dir_world[:, 1],
"post_dz": post_dir_world[:, 2],
"post_x": post_pos_world[:, 0],
"post_y": post_pos_world[:, 1],
"post_z": post_pos_world[:, 2],
"sec_pdg_list": sec_pdg_list,
"sec_E_list": sec_E_list,
"sec_dx_list": sec_dx_list,
"sec_dy_list": sec_dy_list,
"sec_dz_list": sec_dz_list,
}
)
columns = {
"event_id": piece["event_id"],
"pdg": piece["pdg"],
"pre_x": piece["pre_pos"][:, 0],
"pre_y": piece["pre_pos"][:, 1],
"pre_z": piece["pre_pos"][:, 2],
"pre_E": piece["pre_E"],
"pre_dx": piece["pre_dir"][:, 0],
"pre_dy": piece["pre_dir"][:, 1],
"pre_dz": piece["pre_dir"][:, 2],
"material": piece["material"],
"layer_id": piece["layer_id"],
"n_sec": piece["n_sec"],
"n_sec_pred": n_sec_pred_np,
"step_length": step_length,
"delta_e": delta_e,
"edep": edep,
"post_dx": post_dir_world[:, 0],
"post_dy": post_dir_world[:, 1],
"post_dz": post_dir_world[:, 2],
"post_x": post_pos_world[:, 0],
"post_y": post_pos_world[:, 1],
"post_z": post_pos_world[:, 2],
"sec_pdg_list": sec_pdg_list,
"sec_E_list": sec_E_list,
"sec_dx_list": sec_dx_list,
"sec_dy_list": sec_dy_list,
"sec_dz_list": sec_dz_list,
}
if write_truth:
n_sec_true = piece["n_sec"]
columns.update(
{
"true_step_length": piece["step_length"],
"true_delta_e": piece["delta_e"],
"true_edep": piece["edep"],
"true_post_E": piece["post_E"],
"true_post_dx": piece["post_dir"][:, 0],
"true_post_dy": piece["post_dir"][:, 1],
"true_post_dz": piece["post_dir"][:, 2],
"true_post_x": piece["post_pos"][:, 0],
"true_post_y": piece["post_pos"][:, 1],
"true_post_z": piece["post_pos"][:, 2],
"true_e_sec": piece["e_sec"],
"process": piece["process"],
}
)
if "sec_E_list" in piece:
columns.update(
{
"true_sec_pdg_list": [
piece["sec_pdg_list"][i, :n].tolist() for i, n in enumerate(n_sec_true)
],
"true_sec_E_list": [piece["sec_E_list"][i, :n].tolist() for i, n in enumerate(n_sec_true)],
"true_sec_dx_list": [
piece["sec_dir_list"][i, :n, 0].tolist() for i, n in enumerate(n_sec_true)
],
"true_sec_dy_list": [
piece["sec_dir_list"][i, :n, 1].tolist() for i, n in enumerate(n_sec_true)
],
"true_sec_dz_list": [
piece["sec_dir_list"][i, :n, 2].tolist() for i, n in enumerate(n_sec_true)
],
}
)
table = pa.table(columns)
table = table.replace_schema_metadata(
{
PREDICT_COORD_METADATA_KEY: coord.value,
PREDICT_SCHEMA_VERSION_KEY: PREDICT_SCHEMA_VERSION,
PREDICT_TRUTH_METADATA_KEY: "1" if (coord == Coord.local or write_truth) else "0",
}
)
_t0 = time.perf_counter()
if writer is None:
writer = pq.ParquetWriter(out, table.schema)
writer.write_table(table)
_write_s += time.perf_counter() - _t0
total += len(piece["event_id"])
# Buffer rows across row-group boundaries so the inference batch size
# isn't capped by however the source file happens to be chunked.
buffer: dict[str, np.ndarray] | None = None
_t_predict_start = time.perf_counter()
bar = tqdm(total=total_rows, desc="predict", unit="row", dynamic_ncols=True)
for i, path in enumerate(files):
for chunk in chunk_iter(path, offset=event_id_offset(i)):
@@ -1365,9 +1476,56 @@ def predict(
bar.close()
if writer is not None:
writer.close()
_predict_s = time.perf_counter() - _t_predict_start
ref_path = _write_prediction_ref(checkpoint, pred_uuid, out, dataset_path, comment)
timing = _build_predict_timing(
setup_s=_setup_s,
predict_s=_predict_s,
write_s=_write_s,
n_rows=total,
device=str(_device),
torch_threads=torch.get_num_threads(),
)
ref_path = _write_prediction_ref(
checkpoint,
pred_uuid,
out,
dataset_path,
comment,
extra={
"coord": coord.value,
"has_truth": coord == Coord.local or write_truth,
"schema_version": PREDICT_SCHEMA_VERSION,
"steps": steps,
"weights": weights.value,
"device": str(_device),
"batch_size": bs,
"batch_size_auto": batch_size_auto,
"n_input_rows": total_rows,
"n_files": len(files),
"n_rows": total,
"n_skipped_rows": skipped,
"unknown_pdg_counts": {str(pdg): count for pdg, count in unknown_pdg_counts.items()},
"timing": timing,
# Full architecture spec baked into the checkpoint — see the
# matching comment in `rollout`.
"model_config": dict(ctx.model_config),
"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)
# — empty dicts if the checkpoint has no config.toml next to it.
"training_config": dict(training_cfg.get("train", {})),
"training_meta": dict(training_cfg.get("meta", {})),
},
)
typer.echo(f"reference: {ref_path}")
if timing["us_per_step"] is not None:
typer.echo(
f"timing: {_predict_s:.1f}s total ({timing['sample_s']:.1f}s sample + {_write_s:.1f}s write), "
f"{timing['us_per_step']:.1f} us/step over {total:,} step(s)"
)
if skipped:
codes = ", ".join(f"{pdg} ({count})" for pdg, count in sorted(unknown_pdg_counts.items()))
@@ -1499,7 +1657,6 @@ def rollout(
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
@@ -1628,10 +1785,12 @@ def rollout(
l1_summary = l1_dist_collector.summary()
ref_path = _write_prediction_ref(checkpoint, pred_uuid, out, dataset_path)
ref = yaml.safe_load(ref_path.read_text())
ref.update(
{
ref_path = _write_prediction_ref(
checkpoint,
pred_uuid,
out,
dataset_path,
extra={
"kind": "rollout",
"geometry_oracle": str(geometry.resolve()),
"energy_cutoff": energy_cutoff,
@@ -1668,9 +1827,8 @@ def rollout(
# — empty dicts if the checkpoint has no config.toml next to it.
"training_config": dict(training_cfg.get("train", {})),
"training_meta": dict(training_cfg.get("meta", {})),
}
},
)
ref_path.write_text(yaml.dump(ref, default_flow_style=False, sort_keys=False))
typer.echo(f"wrote {summary['n_rows']:,} step rows → {out}")
typer.echo(f"terminations: {summary['termination_reason_counts']}")
@@ -1708,6 +1866,25 @@ def analyze_prep(
"single YAML).",
),
] = None,
prediction: Annotated[
list[Path] | None,
typer.Option(
"--prediction",
help="giant predict YAML sidecar(s) (paired truth/pred comparison, the "
"`prediction` plot family) — optional add-on to the rollout comparison. "
"Every one must be seeded from the same `dataset` as the rollout(s) and "
"share one predict --coord.",
),
] = None,
prediction_label: Annotated[
list[str] | None,
typer.Option(
"--prediction-label",
help="Series name for a --prediction YAML, positionally matched to it — give "
'none, or exactly one per YAML. Defaults to the YAML stem (or "prediction" '
"for a single YAML).",
),
] = None,
run_dir: Annotated[
Path | None,
typer.Option(
@@ -1724,7 +1901,7 @@ def analyze_prep(
typer.Option("--chunks", help="Split each plot's data into this many event_id chunks"),
] = 1,
) -> None:
"""Read the rollout YAML(s) → shared.json + run_meta.json in the run directory."""
"""Read the rollout (+ optional prediction) YAML(s) → shared.json + run_meta.json."""
from giant.analysis import prep
path = prep(
@@ -1733,6 +1910,8 @@ def analyze_prep(
n_chunks=chunks,
default_base=Path.cwd() / "analysis_runs",
labels=label,
prediction_yamls=prediction or (),
prediction_labels=prediction_label,
n_energy_bins=n_energy_bins,
n_marginal_bins=n_marginal_bins,
top_k_pdg=top_k_pdg,
@@ -1831,6 +2010,25 @@ def analyze_submit(
"single YAML).",
),
] = None,
prediction: Annotated[
list[Path] | None,
typer.Option(
"--prediction",
help="giant predict YAML sidecar(s) (paired truth/pred comparison, the "
"`prediction` plot family) — optional add-on to the rollout comparison. "
"Every one must be seeded from the same `dataset` as the rollout(s) and "
"share one predict --coord.",
),
] = None,
prediction_label: Annotated[
list[str] | None,
typer.Option(
"--prediction-label",
help="Series name for a --prediction YAML, positionally matched to it — give "
'none, or exactly one per YAML. Defaults to the YAML stem (or "prediction" '
"for a single YAML).",
),
] = None,
run_dir: Annotated[
Path | None,
typer.Option(
@@ -1868,6 +2066,8 @@ def analyze_submit(
n_chunks=chunks,
default_base=Path.cwd() / "analysis_runs",
labels=label,
prediction_yamls=prediction or (),
prediction_labels=prediction_label,
n_energy_bins=n_energy_bins,
n_marginal_bins=n_marginal_bins,
top_k_pdg=top_k_pdg,
+6 -1
View File
@@ -64,7 +64,12 @@ LOCAL_TARGET_NAMES = [
# guessing from its column names.
PREDICT_COORD_METADATA_KEY = "giant.predict.coord"
PREDICT_SCHEMA_VERSION_KEY = "giant.predict.schema_version"
PREDICT_SCHEMA_VERSION = "2"
PREDICT_SCHEMA_VERSION = "3"
# Whether a --coord global predict parquet also carries true_* / true_sec_*
# columns (v3+; "1"/"0"). Lets analysis code tell a paired prediction file
# apart from a --no-truth one without sniffing for column presence.
PREDICT_TRUTH_METADATA_KEY = "giant.predict.has_truth"
# Coord-metadata value tagging a `giant rollout` steps parquet (world frame,
# autoregressive shower output). Distinct from predict's "global"/"local".
+5 -1
View File
@@ -76,7 +76,11 @@ def _pooled_pdg_lazy(path: Path, has_sec_pdg_list: bool) -> pl.LazyFrame:
lf = pl.scan_parquet(path, row_index_name="__row")
parts = [lf.select(pl.col("pdg").alias("__val"), "__row")]
if has_sec_pdg_list:
parts.append(lf.select(pl.col("sec_pdg_list").alias("__val"), "__row").explode("__val").drop_nulls("__val"))
parts.append(
lf.select(pl.col("sec_pdg_list").alias("__val"), "__row")
.explode("__val", empty_as_null=False)
.drop_nulls("__val")
)
combined = pl.concat(parts)
return combined.group_by("__val").agg(pl.len().alias("__count"), pl.col("__row").min().alias("__first_row"))
+1 -1
View File
@@ -58,7 +58,7 @@ def _add_secondary_attributes(df: pl.DataFrame) -> tuple[pl.DataFrame, int]:
exploded = (
df.select(["event_id", "child_track_ids"])
.with_row_index("_step_row")
.explode("child_track_ids")
.explode("child_track_ids", empty_as_null=False)
.rename({"child_track_ids": "child_track_id"})
.drop_nulls("child_track_id")
)
+4 -4
View File
@@ -1,12 +1,12 @@
[project]
name = "giant"
version = "0.3.21"
version = "0.3.23"
description = "Geant4 step-function surrogate via conditional flow matching"
readme = "README.md"
requires-python = ">=3.12"
dependencies = [
"numpy>=1.26,<3",
"polars>=1.0,<2",
"polars>=1.43,<2",
"pyarrow>=16,<26",
"tqdm>=4.60,<5",
"typer>=0.12,<1",
@@ -43,11 +43,11 @@ wandb = [
convert = [
"uproot>=5.3,<6",
"awkward>=2.6,<3",
"polars>=1.0,<2",
"polars>=1.43,<2",
]
analysis = [
"matplotlib>=3.8,<4",
"polars>=1.0,<2",
"polars>=1.43,<2",
"ipykernel>=7.3.0",
# KIT matplotlib theme, published from git.larsbogner.de. Only the local
# `giant analyze render` step imports it; compute workers never do.
+278
View File
@@ -0,0 +1,278 @@
"""Tests for giant.analysis.prediction (paired truth/pred frames for `giant predict`
output) and the `prediction` family of catalog specs."""
from __future__ import annotations
import numpy as np
import polars as pl
import pytest
from giant.analysis.catalog import Bundle, get_spec
from giant.analysis.context import Context, build_context
from giant.analysis.prediction import (
PAIRED_SCALARS,
PredictionSpec,
open_prediction,
paired_frame,
paired_secondaries,
prediction_secondaries,
)
from giant.analysis.reduce import hist2d
from giant.analysis.sources import RolloutSpec
from tests.test_analysis_reduce import _reference_frame, _rollout_frame
def _global_prediction_frame() -> pl.LazyFrame:
"""A `--coord global --truth` predict parquet, as a LazyFrame (schema per
`giant.cli.predict`'s global-coord table, `giant/cli.py:1310-1379`)."""
return pl.DataFrame(
{
"event_id": [1, 1, 2],
"pdg": [11, 11, 22],
"pre_x": [0.0, 0.0, 0.0],
"pre_y": [0.0, 0.0, 0.0],
"pre_z": [0.0, 1.0, 0.0],
"pre_E": [100.0, 60.0, 50.0],
"pre_dx": [0.0, 0.0, 0.0],
"pre_dy": [0.0, 0.0, 0.0],
"pre_dz": [1.0, 1.0, 1.0],
"material": ["G4_PbWO4", "G4_PbWO4", "G4_Pb"],
"layer_id": [0, 1, 0],
"n_sec": [1, 0, 2],
"n_sec_pred": [1, 0, 1],
# predicted (unprefixed) values
"step_length": [1.2, 0.9, 1.1],
"delta_e": [42.0, 29.0, 31.0],
"edep": [35.0, 29.0, 25.0],
"post_dx": [0.0, 0.0, 0.0],
"post_dy": [0.0, 0.0, 0.0],
"post_dz": [1.0, 1.0, 1.0],
"post_x": [0.0, 0.0, 0.0],
"post_y": [0.0, 0.0, 0.0],
"post_z": [1.2, 1.9, 1.1],
"sec_pdg_list": [[22], [], [22]],
"sec_E_list": [[5.0], [], [4.0]],
"sec_dx_list": [[0.0], [], [0.0]],
"sec_dy_list": [[0.0], [], [0.0]],
"sec_dz_list": [[1.0], [], [1.0]],
# truth
"true_step_length": [1.0, 1.0, 1.0],
"true_delta_e": [40.0, 30.0, 30.0],
"true_edep": [40.0, 30.0, 20.0],
"true_post_E": [60.0, 30.0, 20.0],
"true_post_dx": [0.0, 0.0, 0.0],
"true_post_dy": [0.0, 0.0, 0.0],
"true_post_dz": [1.0, 1.0, 1.0],
"true_post_x": [0.0, 0.0, 0.0],
"true_post_y": [0.0, 0.0, 0.0],
"true_post_z": [1.0, 2.0, 1.0],
"true_e_sec": [0.0, 0.0, 10.0],
"process": ["compt", "phot", "compt"],
"true_sec_pdg_list": [[22], [], [22, 11]],
"true_sec_E_list": [[6.0], [], [7.0, 3.0]],
"true_sec_dx_list": [[0.0], [], [0.0, 1.0]],
"true_sec_dy_list": [[0.0], [], [0.0, 0.0]],
"true_sec_dz_list": [[1.0], [], [1.0, 0.0]],
}
).lazy()
def _local_prediction_frame() -> pl.LazyFrame:
"""A `--coord local` predict parquet — always paired, never has secondaries."""
return pl.DataFrame(
{
"event_id": [1, 2],
"pdg": [11, 22],
"pre_x": [0.0, 0.0],
"pre_y": [0.0, 0.0],
"pre_z": [0.0, 0.0],
"pre_E": [100.0, 50.0],
"pre_dx": [0.0, 0.0],
"pre_dy": [0.0, 0.0],
"pre_dz": [1.0, 1.0],
"material": ["G4_PbWO4", "G4_Pb"],
"layer_id": [0, 0],
"n_sec": [1, 0],
# ALR logits: [edep_logit, sec_logit] -> softmax([z1,z2,0]) * pre_E
"pred_log_step_length": [np.log(1.2 + 1e-6), np.log(0.9 + 1e-6)],
"pred_edep_logit": [1.0, 0.5],
"pred_sec_logit": [0.0, -1.0],
"pred_post_dx": [0.0, 0.0],
"pred_post_dy": [0.0, 0.0],
"pred_post_dz": [1.0, 1.0],
"pred_travel_dx": [0.0, 0.0],
"pred_travel_dy": [0.0, 0.0],
"pred_travel_dz": [1.0, 1.0],
"true_log_step_length": [np.log(1.0 + 1e-6), np.log(1.0 + 1e-6)],
"true_edep_logit": [0.8, 0.6],
"true_sec_logit": [0.2, -2.0],
"true_post_dx": [0.0, 0.0],
"true_post_dy": [0.0, 0.0],
"true_post_dz": [1.0, 1.0],
"true_travel_dx": [0.0, 0.0],
"true_travel_dy": [0.0, 0.0],
"true_travel_dz": [1.0, 1.0],
}
).lazy()
def test_open_prediction_detects_coord_and_truth():
g = open_prediction(_global_prediction_frame())
assert g.coord == "global" and g.has_truth
loc = open_prediction(_local_prediction_frame())
assert loc.coord == "local" and loc.has_truth
def test_paired_frame_global_matches_source_columns():
lf = _global_prediction_frame()
p = paired_frame(lf, "global", has_truth=True).collect()
assert p["pred_step_length"].to_list() == [1.2, 0.9, 1.1]
assert p["true_step_length"].to_list() == [1.0, 1.0, 1.0]
assert p["pred_edep"].to_list() == [35.0, 29.0, 25.0]
assert p["true_edep"].to_list() == [40.0, 30.0, 20.0]
# post_E isn't written directly for the prediction (energy conservation:
# pre_E - delta_e); truth carries it verbatim.
assert p["pred_post_E"].to_list() == pytest.approx([100.0 - 42.0, 60.0 - 29.0, 50.0 - 31.0])
assert p["true_post_E"].to_list() == [60.0, 30.0, 20.0]
# cos_scatter: pre_dir . post_dir, both (0,0,1) here -> 1.0
assert p["pred_cos_scatter"].to_list() == pytest.approx([1.0, 1.0, 1.0])
assert p["true_cos_scatter"].to_list() == pytest.approx([1.0, 1.0, 1.0])
def test_paired_frame_local_decodes_energy_simplex():
lf = _local_prediction_frame()
p = paired_frame(lf, "local", has_truth=True).collect()
# softmax([1.0, 0.0, 0.0]) * 100 for row 0's pred edep
z = np.exp([1.0, 0.0, 0.0])
expected_edep_0 = (z[0] / z.sum()) * 100.0
assert p["pred_edep"][0] == pytest.approx(expected_edep_0)
assert p["pred_step_length"][0] == pytest.approx(1.2, abs=1e-4)
# local coord never has a meaningful cos_travel (no reconstructed post_pos)
assert "cos_travel" not in [c.rsplit("_", 1)[-1] for c in ["pred_cos_travel"] if c in p.columns] or True
assert "pred_cos_travel" not in p.columns
def test_prediction_secondaries_and_pairing():
lf = _global_prediction_frame()
true_sec = prediction_secondaries(lf, "true").collect()
pred_sec = prediction_secondaries(lf, "pred").collect()
assert true_sec["pdg"].to_list() == [22, 22, 11]
assert pred_sec["pdg"].to_list() == [22, 22]
pairs = paired_secondaries(lf).collect()
# event 1: 1 true, 1 pred -> paired (22, 22); event 2: 2 true, 1 pred -> paired rank0 only (22, 22)
assert pairs["true_pdg"].to_list() == [22, 22]
assert pairs["pred_pdg"].to_list() == [22, 22]
def test_hist2d_basic():
lf = pl.DataFrame({"x": [0.1, 0.5, 0.9, 0.5], "y": [0.1, 0.9, 0.9, 0.1]}).lazy()
edges = np.linspace(0.0, 1.0, 3) # 2 bins: [0,0.5), [0.5,1]
mat = hist2d(lf, pl.col("x"), pl.col("y"), edges, edges)
assert mat.sum() == 4
assert mat.shape == (2, 2)
def _ctx_with_predictions(n_marginal_bins: int = 10) -> Context:
return build_context(
[RolloutSpec("rollout", _rollout_frame())],
_reference_frame(),
predictions=[PredictionSpec("pred", _global_prediction_frame())],
n_energy_bins=2,
n_marginal_bins=n_marginal_bins,
top_k_pdg=3,
sample_rows=1000,
)
def test_build_context_resolves_prediction_ranges():
ctx = _ctx_with_predictions()
assert "edep" in ctx.pred_var_ranges
assert "edep" in ctx.pred_residual_ranges
assert ctx.pred_top_sec_pdgs # secondaries present in the fixture
def test_prediction_specs_compute_valid_reduced():
ctx = _ctx_with_predictions()
bundle = Bundle.open(
[RolloutSpec("rollout", _rollout_frame())],
_reference_frame(),
ctx,
predictions=[PredictionSpec("pred", _global_prediction_frame())],
)
for spec_id in (
"pred_marginal_edep",
"pred_scatter_edep",
"pred_residual_edep",
"pred_relative_residual_edep",
"pred_residual_profile_edep",
"pred_ks_summary",
"pred_bias_summary",
"pred_rmse_summary",
"pred_n_sec_confusion",
"pred_sec_species_confusion",
"pred_dir_alignment_post",
"pred_dir_alignment_travel",
"pred_constraint_violations",
"pred_correlation_delta",
):
spec = get_spec(spec_id)
r = spec.finalize([spec.compute_partial(bundle)], ctx)
assert r.id == spec_id
assert r.kind != "unavailable", f"{spec_id} unexpectedly unavailable"
assert "pred" in r.payload["series"]
def test_prediction_specs_unavailable_without_predictions():
ctx = _ctx_with_predictions()
bundle = Bundle.open([RolloutSpec("rollout", _rollout_frame())], _reference_frame(), ctx)
for spec_id in ("pred_marginal_edep", "pred_scatter_edep", "pred_n_sec_confusion", "pred_ks_summary"):
spec = get_spec(spec_id)
r = spec.finalize([spec.compute_partial(bundle)], ctx)
assert r.kind == "unavailable"
assert r.payload["note"]
@pytest.mark.parametrize(
"spec_id",
["pred_marginal_edep", "pred_scatter_edep", "pred_n_sec_confusion", "pred_ks_summary", "pred_correlation_delta"],
)
def test_prediction_chunked_matches_unchunked(spec_id: str):
ctx = _ctx_with_predictions()
specs = [RolloutSpec("rollout", _rollout_frame())]
preds = [PredictionSpec("pred", _global_prediction_frame())]
spec = get_spec(spec_id)
unchunked_bundle = Bundle.open(specs, _reference_frame(), ctx, predictions=preds)
unchunked = spec.finalize([spec.compute_partial(unchunked_bundle)], ctx)
n_chunks = 2
parts = [
spec.compute_partial(Bundle.open(specs, _reference_frame(), ctx, chunk=(k, n_chunks), predictions=preds))
for k in range(n_chunks)
]
chunked = spec.finalize(parts, ctx)
assert chunked.kind == unchunked.kind
_assert_close(unchunked.payload, chunked.payload)
def _assert_close(a, b) -> None:
"""Recursively compare two JSON-shaped payloads (float-tolerant)."""
if isinstance(a, dict):
assert set(a) == set(b)
for k in a:
_assert_close(a[k], b[k])
elif isinstance(a, list):
assert len(a) == len(b)
for x, y in zip(a, b):
_assert_close(x, y)
elif isinstance(a, float):
assert np.isclose(a, b, atol=1e-9) or (np.isnan(a) and np.isnan(b))
else:
assert a == b
def test_paired_scalars_are_subset_of_all_vars():
assert set(PAIRED_SCALARS) <= {"step_length", "edep", "delta_e", "post_E"}
+111
View File
@@ -6,6 +6,7 @@ from typer.testing import CliRunner
from giant.cli import (
_CEPH_PREDICTIONS,
_build_predict_timing,
_resolve_prediction_output,
_write_prediction_ref,
app,
@@ -103,6 +104,7 @@ def test_ref_yaml_contains_expected_fields(tmp_path):
ref_path = _write_prediction_ref(checkpoint, pred_uuid, out, dataset)
data = yaml.safe_load(ref_path.read_text())
assert data["kind"] == "prediction"
assert data["prediction_id"] == pred_uuid
assert data["output"] == str(out)
assert data["dataset"] == str(dataset)
@@ -144,6 +146,46 @@ def test_ref_timestamp_is_iso_format(tmp_path):
assert ts.tzinfo is not None
def test_ref_yaml_merges_extra_after_base_fields(tmp_path):
ckpt_dir = tmp_path / "checkpoints"
ckpt_dir.mkdir()
checkpoint = ckpt_dir / "best.pt"
checkpoint.touch()
out = tmp_path / "pred.parquet"
dataset = tmp_path / "full.manifest"
pred_uuid = str(uuid.uuid4())
ref_path = _write_prediction_ref(
checkpoint,
pred_uuid,
out,
dataset,
extra={"coord": "global", "n_rows": 42, "timing": {"setup_s": 1.0}},
)
data = yaml.safe_load(ref_path.read_text())
# Base fields untouched, extras layered on top.
assert data["kind"] == "prediction"
assert data["prediction_id"] == pred_uuid
assert data["coord"] == "global"
assert data["n_rows"] == 42
assert data["timing"] == {"setup_s": 1.0}
def test_ref_yaml_without_extra_matches_today(tmp_path):
ckpt_dir = tmp_path / "checkpoints"
ckpt_dir.mkdir()
checkpoint = ckpt_dir / "best.pt"
checkpoint.touch()
pred_uuid = str(uuid.uuid4())
ref_path = _write_prediction_ref(checkpoint, pred_uuid, tmp_path / "p.parquet", tmp_path / "d")
data = yaml.safe_load(ref_path.read_text())
assert set(data) == {"kind", "prediction_id", "output", "dataset", "checkpoint", "timestamp"}
def test_ref_checkpoint_path_is_absolute(tmp_path):
ckpt_dir = tmp_path / "checkpoints"
ckpt_dir.mkdir()
@@ -212,3 +254,72 @@ def test_predict_set_flag_disallowed_path_surfaces_compat_error(tmp_path):
assert result.exit_code == 1
assert "not an inference-safe override" in result.output
# ---------------------------------------------------------------------------
# schema v3 constants (truth-column tagging)
# ---------------------------------------------------------------------------
def test_predict_schema_version_is_v3():
from giant.constants import PREDICT_SCHEMA_VERSION
assert PREDICT_SCHEMA_VERSION == "3"
def test_predict_truth_metadata_key_exists():
from giant.constants import PREDICT_TRUTH_METADATA_KEY
assert PREDICT_TRUTH_METADATA_KEY == "giant.predict.has_truth"
# ---------------------------------------------------------------------------
# _build_predict_timing
# ---------------------------------------------------------------------------
def test_build_predict_timing_computes_per_step_cost():
timing = _build_predict_timing(
setup_s=1.0,
predict_s=10.0,
write_s=2.0,
n_rows=100,
device="cpu",
torch_threads=4,
)
assert timing["n_rows"] == 100
assert timing["sample_s"] == 8.0 # predict_s - write_s
assert timing["us_per_step"] == 8.0 / 100 * 1e6
assert timing["write_us_per_step"] == 2.0 / 100 * 1e6
assert timing["rows_per_s"] == 10.0
assert timing["device"] == "cpu" and timing["torch_threads"] == 4
def test_build_predict_timing_handles_zero_rows():
timing = _build_predict_timing(
setup_s=1.0,
predict_s=0.0,
write_s=0.0,
n_rows=0,
device="cpu",
torch_threads=1,
)
assert timing["us_per_step"] is None
assert timing["write_us_per_step"] is None
assert timing["rows_per_s"] is None
def test_predict_has_truth_flag_default_on():
# Inspecting rendered --help text is brittle across terminal
# widths/color settings (wraps or re-colors mid-flag); go straight to
# the underlying click command's registered option instead.
from typing import cast
import typer
from click import Group
predict_cmd = cast(Group, typer.main.get_command(app)).commands["predict"]
truth_param = next(p for p in predict_cmd.params if p.name == "truth")
assert truth_param.opts == ["--truth"]
assert truth_param.secondary_opts == ["--no-truth"]
assert truth_param.default is True
+137 -2
View File
@@ -16,6 +16,8 @@ from giant.analysis import (
compute_one,
compute_reduced,
derive_run_dir,
load_prediction_yaml,
load_prediction_yamls,
load_rollout_yaml,
load_rollout_yamls,
merge_one,
@@ -25,7 +27,8 @@ from giant.analysis import (
from giant.analysis.catalog import get_spec
from giant.analysis.condor import Context
from giant.analysis.reduced import Partial, Reduced
from giant.constants import PREDICT_COORD_METADATA_KEY, ROLLOUT_COORD_VALUE
from giant.constants import PREDICT_COORD_METADATA_KEY, PREDICT_TRUTH_METADATA_KEY, ROLLOUT_COORD_VALUE
from tests.test_analysis_prediction import _global_prediction_frame
from tests.test_analysis_reduce import _reference_frame, _rollout_frame
@@ -86,6 +89,30 @@ def _write_two_inputs(tmp_path: Path) -> tuple[Path, Path]:
return paths[0], paths[1]
def _write_prediction(path: Path, coord: str = "global") -> None:
tbl = _global_prediction_frame().collect().to_arrow()
tbl = tbl.replace_schema_metadata({PREDICT_COORD_METADATA_KEY: coord, PREDICT_TRUTH_METADATA_KEY: "1"})
pq.write_table(tbl, path)
def _write_prediction_yaml(tmp_path: Path, reference: Path, tag: str = "p", coord: str = "global") -> Path:
pred = tmp_path / f"pred_{tag}.parquet"
_write_prediction(pred, coord=coord)
yaml_path = tmp_path / f"pred_{tag}.yaml"
yaml_path.write_text(
yaml.safe_dump(
{
"kind": "prediction",
"prediction_id": f"{tag}pred1234",
"output": str(pred),
"dataset": str(reference),
"checkpoint": f"/ckpt/{tag}.pt",
}
)
)
return yaml_path
def _fake_venv(repo_dir: Path) -> None:
"""Stand in for a `uv sync`'d venv: write_submit checks `.venv/bin/giant` exists."""
giant = repo_dir / ".venv" / "bin" / "giant"
@@ -94,13 +121,14 @@ def _fake_venv(repo_dir: Path) -> None:
giant.chmod(0o755)
def _prep(rollout_yamls, run_dir: str | Path | None = None, chunks: int = 1, labels=None) -> Path:
def _prep(rollout_yamls, run_dir: str | Path | None = None, chunks: int = 1, labels=None, prediction_yamls=()) -> Path:
"""``prep`` with small test-sized context bins/sampling."""
return prep(
rollout_yamls,
run_dir,
n_chunks=chunks,
labels=labels,
prediction_yamls=prediction_yamls,
n_energy_bins=2,
n_marginal_bins=8,
top_k_pdg=3,
@@ -161,6 +189,113 @@ def test_load_rollout_yamls_rejects_mismatched_reference(tmp_path: Path):
load_rollout_yamls([a, c])
def test_load_prediction_yaml_requires_paths(tmp_path: Path):
bad = tmp_path / "bad.yaml"
bad.write_text(yaml.safe_dump({"output": "x.parquet"})) # no dataset
with pytest.raises(ValueError):
load_prediction_yaml(bad)
def test_load_prediction_yaml_rejects_rollout_kind(tmp_path: Path):
y = tmp_path / "r.yaml"
y.write_text(yaml.safe_dump({"output": "x.parquet", "dataset": "d.parquet", "kind": "rollout"}))
with pytest.raises(ValueError, match="kind"):
load_prediction_yaml(y)
def test_load_prediction_yamls_single_defaults_to_prediction_name(tmp_path: Path):
reference = tmp_path / "reference.parquet"
_reference_frame().collect().write_parquet(reference)
y = _write_prediction_yaml(tmp_path, reference)
loaded = load_prediction_yamls([y], str(reference))
assert [lp.name for lp in loaded] == ["prediction"]
assert loaded[0].coord == "global"
def test_load_prediction_yamls_multi_defaults_to_stem_and_labels(tmp_path: Path):
reference = tmp_path / "reference.parquet"
_reference_frame().collect().write_parquet(reference)
a = _write_prediction_yaml(tmp_path, reference, tag="a")
b = _write_prediction_yaml(tmp_path, reference, tag="b")
loaded = load_prediction_yamls([a, b], str(reference))
assert [lp.name for lp in loaded] == ["pred_a", "pred_b"]
loaded = load_prediction_yamls([a, b], str(reference), labels=["ep20", "ep50"])
assert [lp.name for lp in loaded] == ["ep20", "ep50"]
def test_load_prediction_yamls_rejects_mismatched_reference(tmp_path: Path):
reference = tmp_path / "reference.parquet"
_reference_frame().collect().write_parquet(reference)
other_ref = tmp_path / "other_reference.parquet"
_reference_frame().collect().write_parquet(other_ref)
y = _write_prediction_yaml(tmp_path, other_ref)
with pytest.raises(ValueError, match="same reference"):
load_prediction_yamls([y], str(reference))
def test_load_prediction_yamls_rejects_mixed_coord(tmp_path: Path):
reference = tmp_path / "reference.parquet"
_reference_frame().collect().write_parquet(reference)
a = _write_prediction_yaml(tmp_path, reference, tag="a", coord="global")
b = _write_prediction_yaml(tmp_path, reference, tag="b", coord="local")
with pytest.raises(ValueError, match="coord"):
load_prediction_yamls([a, b], str(reference))
def test_prep_with_prediction_writes_run_meta(tmp_path: Path):
rollout_yaml = _write_inputs(tmp_path)
reference = load_rollout_yaml(rollout_yaml)["dataset"]
pred_yaml = _write_prediction_yaml(tmp_path, Path(reference))
run_dir = _prep([rollout_yaml], prediction_yamls=[pred_yaml])
meta = RunMeta.load(run_dir / "run_meta.json")
assert [p["name"] for p in meta.predictions] == ["prediction"]
assert meta.predictions[0]["plot_meta"]["checkpoint"] == "/ckpt/p.pt"
computed = compute_one("pred_marginal_edep", run_dir, chunk_index=0)
partial = Partial.load(computed)
assert partial.data["available"]
def test_prep_forwards_predict_only_metadata_keys(tmp_path: Path):
"""A rich `giant predict` sidecar's provenance/timing keys reach
run_meta.json's plot_meta, same as a rollout's do a thin legacy
sidecar (no such keys) still loads fine (see _write_prediction_yaml)."""
rollout_yaml = _write_inputs(tmp_path)
reference = load_rollout_yaml(rollout_yaml)["dataset"]
pred = tmp_path / "pred_rich.parquet"
_write_prediction(pred, coord="global")
yaml_path = tmp_path / "pred_rich.yaml"
yaml_path.write_text(
yaml.safe_dump(
{
"kind": "prediction",
"prediction_id": "richpred12",
"output": str(pred),
"dataset": str(reference),
"checkpoint": "/ckpt/rich.pt",
"coord": "global",
"has_truth": True,
"schema_version": "3",
"n_input_rows": 1000,
"n_files": 1,
"n_skipped_rows": 3,
"unknown_pdg_counts": {"999999": 3},
"batch_size_auto": False,
"timing": {"us_per_step": 12.5},
}
)
)
run_dir = _prep([rollout_yaml], prediction_yamls=[yaml_path])
meta = RunMeta.load(run_dir / "run_meta.json")
plot_meta = meta.predictions[0]["plot_meta"]
assert plot_meta["coord"] == "global"
assert plot_meta["has_truth"] is True
assert plot_meta["n_input_rows"] == 1000
assert plot_meta["n_skipped_rows"] == 3
assert plot_meta["unknown_pdg_counts"] == {"999999": 3}
assert plot_meta["timing"] == {"us_per_step": 12.5}
def test_derive_run_dir_next_to_rollout():
y = {"output": "/data/roll.parquet", "prediction_id": "abcd1234ef", "dataset": "d"}
assert derive_run_dir([y]) == Path("/data/analysis_abcd1234")
+147
View File
@@ -0,0 +1,147 @@
"""Tests for the scheduled dependency-bump automation (weekly uv.lock
refresh, monthly pyproject upper-bound raise). The workflow YAMLs themselves
can only be exercised by a real scheduled/dispatched run (same reasoning as
tests/test_release_tooling.py for the release workflow), so this checks the
script logic they drive plus the one piece of cross-file coupling that would
silently misbehave if it drifted: each workflow's script must target the
same standing branch that deps-pr.sh (invoked by that script) manages.
"""
import importlib.util
import re
import sys
import tomllib
from pathlib import Path
import pytest
import yaml
from packaging.requirements import Requirement
from packaging.version import Version
_ROOT = Path(__file__).resolve().parents[1]
_SCRIPTS = _ROOT / ".gitea" / "scripts"
_WORKFLOWS = _ROOT / ".gitea" / "workflows"
def _load_check_dep_bounds():
spec = importlib.util.spec_from_file_location("check_dep_bounds", _SCRIPTS / "check_dep_bounds.py")
assert spec is not None and spec.loader is not None
module = importlib.util.module_from_spec(spec)
sys.modules[spec.name] = module
spec.loader.exec_module(module)
return module
@pytest.fixture(scope="module")
def cdb():
return _load_check_dep_bounds()
def test_canonical_requirements_finds_real_dependencies_and_skips_excluded(cdb):
pyproject = tomllib.loads((_ROOT / "pyproject.toml").read_text())
requirements = cdb.canonical_requirements(pyproject)
assert "numpy" in requirements
assert "polars" in requirements
# polars appears in [project.dependencies], [convert], and [analysis]
# with the same specifier — canonicalization must dedupe by name.
assert isinstance(requirements["polars"], Requirement)
for skipped in ("torch", "plotstyle", "giant"):
assert skipped in cdb.SKIP_REASONS, f"{skipped} must stay on the never-auto-raise skip list"
def test_next_ceiling_pre_and_post_1_0():
from importlib import import_module
cdb = import_module("check_dep_bounds")
# Post-1.0: next ceiling is the major above latest.
assert cdb.next_ceiling(Version("2.5.3")) == "<3"
assert cdb.next_ceiling(Version("25.0.1")) == "<26"
# Pre-1.0: next ceiling is the minor above latest (matches this repo's
# own pins, e.g. ruff>=0.15,<1 and ty>=0.0.50,<0.1).
assert cdb.next_ceiling(Version("0.16.2")) == "<0.17"
def test_find_findings_flags_out_of_range_and_skips_excluded(cdb, monkeypatch):
requirements = {
"numpy": Requirement("numpy>=1.26,<3"),
"widget": Requirement("widget>=1,<2"),
"torch": Requirement("torch>=2.3,<2.4"),
}
fake_latest = {"numpy": "2.5.3", "widget": "3.1.0", "torch": "2.9.0"}
monkeypatch.setattr(cdb, "fetch_latest_version", lambda name: fake_latest[name])
findings = cdb.find_findings(requirements, {"torch": "pinned deliberately"})
assert len(findings) == 1
finding = findings[0]
assert finding.name == "widget"
assert set(finding.old_specifier.split(",")) == {">=1", "<2"}
# SpecifierSet doesn't guarantee clause order, so compare as a set.
assert set(finding.new_specifier.split(",")) == {">=1", "<4"}
assert finding.latest == "3.1.0"
def test_apply_findings_rewrites_every_occurrence_and_nothing_else(cdb):
text = (
"[project]\n"
"dependencies = [\n"
' "polars>=1.43,<2",\n'
"]\n\n"
"[project.optional-dependencies]\n"
"convert = [\n"
' "polars>=1.43,<2",\n'
"]\n"
"other = [\n"
' "numpy>=1.26,<3",\n'
"]\n"
)
finding = cdb.Finding(name="polars", old_specifier=">=1.43,<2", new_specifier=">=1.43,<3", latest="2.0.0")
new_text = cdb.apply_findings(text, [finding])
assert new_text.count('"polars>=1.43,<3"') == 2
assert '"polars>=1.43,<2"' not in new_text
# Untouched dependency (numpy) survives byte-for-byte.
assert '"numpy>=1.26,<3"' in new_text
def test_apply_findings_on_real_pyproject_is_a_noop_when_no_findings(cdb):
text = (_ROOT / "pyproject.toml").read_text()
assert cdb.apply_findings(text, []) == text
@pytest.mark.parametrize(
("workflow_file", "script_file"),
[
("deps-lock.yml", "deps-lock-pr.sh"),
("deps-bounds.yml", "deps-bounds-pr.sh"),
],
)
def test_workflow_yaml_parses_and_calls_matching_script(workflow_file, script_file):
workflow = yaml.safe_load((_WORKFLOWS / workflow_file).read_text())
assert "schedule" in workflow["on"]
assert "workflow_dispatch" in workflow["on"]
workflow_text = (_WORKFLOWS / workflow_file).read_text()
assert script_file in workflow_text, f"{workflow_file} must invoke .gitea/scripts/{script_file}"
@pytest.mark.parametrize(
("script_file", "branch"),
[
("deps-lock-pr.sh", "chore/uv-lock-upgrade"),
("deps-bounds-pr.sh", "chore/dep-bounds"),
],
)
def test_driver_script_branch_matches_deps_pr_invocations(script_file, branch):
"""Each driver script must define BRANCH as the expected literal, then
route every git/deps-pr.sh call through that one $BRANCH variable a
hardcoded mismatch would silently leave a stray branch/PR or open a
second one each run."""
text = (_SCRIPTS / script_file).read_text()
assert f'BRANCH="{branch}"' in text, f"{script_file} should set BRANCH={branch!r}"
assert 'git checkout -B "$BRANCH"' in text
assert re.search(r'deps-pr\.sh open "\$BRANCH"', text)
assert re.search(r'deps-pr\.sh close "\$BRANCH"', text)
+52
View File
@@ -329,6 +329,58 @@ def test_render_one_of_each_kind(tmp_path: Path):
"log_color": True,
},
),
Reduced(
"ph1",
"prediction",
"paired_hist",
"Paired hist (single prediction)",
"x",
{"edges": [0, 1, 2, 3], "series": {"pred": {"pred": [1, 2, 3], "true": [2, 2, 2]}}, "log_y": False},
),
Reduced(
"ph2",
"prediction",
"paired_hist",
"Paired hist (two predictions)",
"x",
{
"edges": [0, 1, 2, 3],
"series": {"a": {"pred": [1, 2, 3], "true": [2, 2, 2]}, "b": {"pred": [3, 2, 1]}},
"log_y": False,
},
),
Reduced(
"hm2d",
"prediction",
"heatmap2d",
"Scatter (truth vs pred)",
"true x",
{
"x_edges": [0, 1, 2],
"y_edges": [0, 1, 2],
"series": {"pred": [[2, 0], [1, 3]]},
"ylabel": "predicted x",
"cbar_label": "count",
"log_color": True,
"diagonal": True,
},
),
Reduced(
"profile_noref",
"prediction",
"profile",
"Residual profile (no reference)",
"true x",
{"edges": [0, 1, 2], "series": {"pred": {"mean": [0.1, -0.1], "std": [0.2, 0.2]}}},
),
Reduced(
"bar_noref",
"prediction",
"bar",
"Constraint violations (no reference)",
"check",
{"labels": ["a", "b"], "series": {"pred": [0.01, 0.0]}, "ylabel": "rate"},
),
]
try:
pdfs = _try_render(reduced, tmp_path)
Generated
+4 -4
View File
@@ -825,7 +825,7 @@ wheels = [
[[package]]
name = "giant"
version = "0.3.21"
version = "0.3.23"
source = { editable = "." }
dependencies = [
{ name = "numpy" },
@@ -892,9 +892,9 @@ requires-dist = [
{ name = "pandas", marker = "extra == 'dev'", specifier = ">=2.2,<4" },
{ name = "particle", specifier = ">=1.0,<2" },
{ name = "plotstyle", marker = "extra == 'analysis'", specifier = ">=1.0.0", index = "https://git.larsbogner.de/api/packages/lars/pypi/simple/" },
{ name = "polars", specifier = ">=1.0,<2" },
{ name = "polars", marker = "extra == 'analysis'", specifier = ">=1.0,<2" },
{ name = "polars", marker = "extra == 'convert'", specifier = ">=1.0,<2" },
{ name = "polars", specifier = ">=1.43,<2" },
{ name = "polars", marker = "extra == 'analysis'", specifier = ">=1.43,<2" },
{ name = "polars", marker = "extra == 'convert'", specifier = ">=1.43,<2" },
{ name = "pyarrow", specifier = ">=16,<26" },
{ name = "pytest", marker = "extra == 'dev'", specifier = ">=8,<10" },
{ name = "pytest-cov", marker = "extra == 'dev'", specifier = ">=5,<8" },