Merge energy-conservation-poc into phase2-secondary-prediction

Brings the energy-conservation PoC work (dwarf CLI unification, dwarf
status improvements, predict --comment, ODE-step comparison scripts,
predict-parquet-only analysis refactor) onto the Phase 2 branch.

Conflict resolution:
- giant/analysis.py: took the energy-conservation-poc version wholesale.
  That branch deliberately removed the live checkpoint+sampler diagnostics
  path (ModelBundle/load_model_bundle/make_val_loader/collect_samples) in
  favor of reading `giant predict --coord local` parquet output. Phase 2's
  only edits to this file adapted the removed path to the new dataset API,
  so nothing Phase-2-specific is lost; no external code called those funcs.

Fixes for pre-existing breakage surfaced by the merge (both predate it):
- giant/cli.py: predict's `_process` unpacked build_features into 5 values,
  but Phase 2 made it return 8 (added n_sec/sec_cont/sec_pdg_idx). Expanded
  the unpack; `giant predict --coord local` would have crashed otherwise.
- tests/test_steps_to_parquet.py: Phase 2 renamed _add_secondary_energy ->
  _add_secondary_attributes without updating this test. Renamed the calls
  and extended the fixture with the pdg/pre_d{x,y,z} columns the expanded
  function reads; e_sec assertions unchanged.
- analysis/compare_ode_steps_energy_conservation.py: E731 lambda assignment
  (added in the un-linted final PoC commit) rewritten as a def.

ruff, ty, and pytest (179 passed) all green.

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
This commit is contained in:
2026-07-06 12:18:30 +02:00
33 changed files with 2316 additions and 964 deletions
@@ -0,0 +1,148 @@
"""Compare energy-conservation PoC at 10 vs 20 ODE steps.
Runs the same event-level energy-budget analysis as
`analysis/export_energy_conservation_poc.py` on both predict outputs and prints a
side-by-side table. Also regenerates the two incident-energy comparison plots for
the 20-step run (prefix `giant-energy-conservation-poc-ode20-`).
"""
from pathlib import Path
import matplotlib.pyplot as plt
import numpy as np
import polars as pl
from giant.analysis import _edep_pl, _hist_edges
FILES = {
"10-step (baseline)": "/home/lars/Programming/giant/9879e806-5e88-4b06-b1fa-0e61de9cda6f.parquet",
"20-step": "/home/lars/Programming/giant/0e236919-65ae-4b9b-9957-d31a8211aca4.parquet",
}
OUT = Path("/home/lars/knowledge-base/meta/attachments")
PREFIX20 = "giant-energy-conservation-poc-ode20"
def per_event(file: str) -> dict:
pe = (
pl.scan_parquet(file)
.group_by("event_id")
.agg(
pl.col("pre_E").max().alias("primary_E"),
_edep_pl("true").sum().alias("real_total_edep"),
_edep_pl("pred").sum().alias("gen_total_edep"),
)
.collect(engine="streaming")
)
primary_E = pe["primary_E"].to_numpy()
assert np.unique(primary_E).size == 1, "expected a single fixed incident energy"
E0 = float(primary_E[0])
real = pe["real_total_edep"].to_numpy()
gen = pe["gen_total_edep"].to_numpy()
n_steps = (
pl.scan_parquet(file).select(pl.len()).collect(engine="streaming").item()
)
return {
"E0": E0,
"n_events": pe.height,
"n_rows": n_steps,
"real": real,
"gen": gen,
}
results = {name: per_event(f) for name, f in FILES.items()}
def fmt_row(label, fn):
cells = " ".join(f"{fn(r):>14}" for r in results.values())
print(f"{label:<32}{cells}")
print("=" * 80)
header = " ".join(f"{name:>14}" for name in results)
print(f"{'metric':<32}{header}")
print("-" * 80)
fmt_row("n_events", lambda r: r["n_events"])
fmt_row("n_rows (steps)", lambda r: r["n_rows"])
fmt_row("E0 [MeV]", lambda r: f"{r['E0']:.1f}")
print("-- REAL --")
fmt_row("real mean [MeV]", lambda r: f"{r['real'].mean():.3f}")
fmt_row("real std [MeV]", lambda r: f"{r['real'].std():.3f}")
fmt_row("real sigma/mu", lambda r: f"{r['real'].std() / r['real'].mean():.4f}")
print("-- GENERATED --")
fmt_row("gen mean [MeV]", lambda r: f"{r['gen'].mean():.3f}")
fmt_row("gen std [MeV]", lambda r: f"{r['gen'].std():.3f}")
fmt_row("gen sigma/mu", lambda r: f"{r['gen'].std() / r['gen'].mean():.4f}")
fmt_row("gen max [MeV]", lambda r: f"{r['gen'].max():.3f}")
fmt_row("gen mean/E0", lambda r: f"{r['gen'].mean() / r['E0']:.4f}")
fmt_row("gen max/E0", lambda r: f"{r['gen'].max() / r['E0']:.4f}")
fmt_row("gen p99/E0", lambda r: f"{np.quantile(r['gen'], 0.99) / r['E0']:.4f}")
fmt_row("frac events gen>E0", lambda r: f"{np.mean(r['gen'] > r['E0']):.4f}")
def disp_ratio(r):
return (r["gen"].std() / r["gen"].mean()) / (r["real"].std() / r["real"].mean())
fmt_row("dispersion ratio gen/real", lambda r: f"{disp_ratio(r):.2f}x")
print("=" * 80)
# --- plots for the 20-step run (mirror the baseline export) ---
r = results["20-step"]
E0, real_tot, gen_tot = r["E0"], r["real"], r["gen"]
fig, ax = plt.subplots(figsize=(6, 4))
edges = _hist_edges(real_tot, gen_tot, bins=50).tolist()
ax.hist(
real_tot, bins=edges, density=True, histtype="step",
label=f"real (σ/μ={real_tot.std() / real_tot.mean():.3f})",
)
ax.hist(
gen_tot, bins=edges, density=True, histtype="step",
label=f"generated (σ/μ={gen_tot.std() / gen_tot.mean():.3f}, "
f"{np.mean(gen_tot > E0):.1%} > E0)",
)
ax.axvline(E0, color="k", linestyle="--", linewidth=1, label=f"incident energy E0={E0:.0f} MeV")
ax.set_yscale("log")
ax.set_xlabel("total deposited energy per event [MeV]")
ax.set_title("20 ODE steps")
ax.legend(fontsize=8)
fig.tight_layout()
fig.savefig(OUT / f"{PREFIX20}-event-total-energy-vs-E0.png", dpi=150, bbox_inches="tight")
fig, ax = plt.subplots(figsize=(6, 4))
ratio_real = real_tot / E0
ratio_gen = gen_tot / E0
edges = _hist_edges(ratio_real, ratio_gen, bins=60).tolist()
ax.hist(ratio_real, bins=edges, density=True, histtype="step", label="real")
ax.hist(ratio_gen, bins=edges, density=True, histtype="step", label="generated")
ax.axvline(1.0, color="k", linestyle="--", linewidth=1, label="conservation limit (=1)")
ax.set_yscale("log")
ax.set_xlabel("total deposited energy / incident energy, per event")
ax.set_title("20 ODE steps")
ax.legend(fontsize=8)
fig.tight_layout()
fig.savefig(OUT / f"{PREFIX20}-event-energy-ratio.png", dpi=150, bbox_inches="tight")
# --- overlay: generated total-edep, 10 vs 20 steps, against real+E0 ---
fig, ax = plt.subplots(figsize=(6, 4))
all_arrays = [results["10-step (baseline)"]["real"]] + [
r2["gen"] for r2 in results.values()
]
edges = _hist_edges(*all_arrays, bins=60).tolist()
ax.hist(
results["20-step"]["real"], bins=edges, density=True, histtype="step",
color="k", label="real",
)
for name, r2 in results.items():
ax.hist(
r2["gen"], bins=edges, density=True, histtype="step",
label=f"gen {name} ({np.mean(r2['gen'] > r2['E0']):.1%} > E0)",
)
ax.axvline(E0, color="gray", linestyle="--", linewidth=1, label=f"E0={E0:.0f} MeV")
ax.set_yscale("log")
ax.set_xlabel("total deposited energy per event [MeV]")
ax.set_title("Generated event energy: 10 vs 20 ODE steps")
ax.legend(fontsize=8)
fig.tight_layout()
fig.savefig(OUT / f"{PREFIX20}-compare-event-total-energy.png", dpi=150, bbox_inches="tight")
print("DONE")
+71
View File
@@ -0,0 +1,71 @@
"""Per-step marginal KL(real||gen) per target dim, 10 vs 20 ODE steps.
Streaming histogram over shared bin edges (computed from the true distribution),
so the two runs are directly comparable dim-by-dim.
"""
import numpy as np
import polars as pl
from giant.analysis import RAW_TARGET_NAMES, _kl_from_counts, _raw_dim_expr
FILES = {
"10-step": "/home/lars/Programming/giant/9879e806-5e88-4b06-b1fa-0e61de9cda6f.parquet",
"20-step": "/home/lars/Programming/giant/0e236919-65ae-4b9b-9957-d31a8211aca4.parquet",
}
BINS = 100
# Fixed shared edges from the true distribution (identical across both files), using
# robust quantiles to avoid a few outliers dominating the range.
base = FILES["10-step"]
edges = {}
for j, name in enumerate(RAW_TARGET_NAMES):
lo, hi = (
pl.scan_parquet(base)
.select(
_raw_dim_expr("true", j).quantile(0.001).alias("lo"),
_raw_dim_expr("true", j).quantile(0.999).alias("hi"),
)
.collect(engine="streaming")
.row(0)
)
if not (hi - lo > 1e-9):
lo, hi = lo - 0.5, hi + 0.5
edges[name] = np.linspace(lo, hi, BINS + 1)
def counts(file, prefix, j, e):
vals = (
pl.scan_parquet(file)
.select(_raw_dim_expr(prefix, j).alias("v"))
.collect(engine="streaming")["v"]
.to_numpy()
)
c, _ = np.histogram(vals, bins=e)
return c
kls = {name: {} for name in FILES}
for j, name in enumerate(RAW_TARGET_NAMES):
e = edges[name]
real_c = counts(base, "true", j, e) # identical true dist across files
for run, f in FILES.items():
gen_c = counts(f, "pred", j, e)
kls[run][name] = _kl_from_counts(real_c, gen_c)
print(f"{'dim':<14}{'KL 10-step':>14}{'KL 20-step':>14}{'ratio 20/10':>14}")
print("-" * 56)
tot = {"10-step": 0.0, "20-step": 0.0}
for name in RAW_TARGET_NAMES:
a, b = kls["10-step"][name], kls["20-step"][name]
tot["10-step"] += a
tot["20-step"] += b
print(f"{name:<14}{a:>14.5f}{b:>14.5f}{b / a if a else float('nan'):>14.2f}")
print("-" * 56)
print(
f"{'SUM':<14}{tot['10-step']:>14.5f}{tot['20-step']:>14.5f}"
f"{tot['20-step'] / tot['10-step']:>14.2f}"
)
print(
f"{'MEAN':<14}{tot['10-step'] / 9:>14.5f}{tot['20-step'] / 9:>14.5f}"
)
+101
View File
@@ -0,0 +1,101 @@
"""One-off export of the energy-budget-violation plot for the
energy-conservation PoC checkpoint (ALR simplex output space, retrained on
the patched-Geant4 regenerated dataset). Not part of the package; run
manually.
Unlike `plot_total_energy` (per-step conservation only, no reference to the
fixed primary/incident energy), this adds an explicit comparison against
`primary_E` (== max(pre_E) per event, since this PoC dataset uses a single
fixed incident energy) to show event-level conservation violations that the
per-step ALR simplex constraint does not prevent.
"""
from pathlib import Path
import matplotlib.pyplot as plt
import numpy as np
import polars as pl
from giant.analysis import _edep_pl, _hist_edges
FILE = "/home/lars/Programming/giant/9879e806-5e88-4b06-b1fa-0e61de9cda6f.parquet"
OUT = Path("/home/lars/knowledge-base/meta/attachments")
PREFIX = "giant-energy-conservation-poc"
print("=== per-event primary energy + total edep ===")
per_event = (
pl.scan_parquet(FILE)
.group_by("event_id")
.agg(
pl.col("pre_E").max().alias("primary_E"),
_edep_pl("true").sum().alias("real_total_edep"),
_edep_pl("pred").sum().alias("gen_total_edep"),
)
.collect(engine="streaming")
)
print("n events:", per_event.height)
primary_E = per_event["primary_E"].to_numpy()
real_tot = per_event["real_total_edep"].to_numpy()
gen_tot = per_event["gen_total_edep"].to_numpy()
assert np.unique(primary_E).size == 1, "expected a single fixed incident energy"
E0 = float(primary_E[0])
print(f"fixed incident energy E0 = {E0} MeV")
print()
print(
f"real: mean={real_tot.mean():.3f} std={real_tot.std():.3f} "
f"sigma/mu={real_tot.std() / real_tot.mean():.4f} max={real_tot.max():.3f} "
f"frac>E0={np.mean(real_tot > E0):.4f}"
)
print(
f"gen: mean={gen_tot.mean():.3f} std={gen_tot.std():.3f} "
f"sigma/mu={gen_tot.std() / gen_tot.mean():.4f} max={gen_tot.max():.3f} "
f"frac>E0={np.mean(gen_tot > E0):.4f}"
)
print(
f"gen/E0 ratio: mean={np.mean(gen_tot / E0):.4f} max={np.max(gen_tot / E0):.4f} "
f"p99={np.quantile(gen_tot / E0, 0.99):.4f}"
)
print("=== plot: total edep per event, marked against incident energy ===")
fig, ax = plt.subplots(figsize=(6, 4))
edges = _hist_edges(real_tot, gen_tot, bins=50).tolist()
ax.hist(
real_tot,
bins=edges,
density=True,
histtype="step",
label=f"real (σ/μ={real_tot.std() / real_tot.mean():.3f})",
)
ax.hist(
gen_tot,
bins=edges,
density=True,
histtype="step",
label=f"generated (σ/μ={gen_tot.std() / gen_tot.mean():.3f}, "
f"{np.mean(gen_tot > E0):.1%} > E0)",
)
ax.axvline(E0, color="k", linestyle="--", linewidth=1, label=f"incident energy E0={E0:.0f} MeV")
ax.set_yscale("log")
ax.set_xlabel("total deposited energy per event [MeV]")
ax.legend(fontsize=8)
fig.tight_layout()
fig.savefig(OUT / f"{PREFIX}-event-total-energy-vs-E0.png", dpi=150, bbox_inches="tight")
print("=== plot: total edep / incident energy ratio ===")
fig, ax = plt.subplots(figsize=(6, 4))
ratio_real = real_tot / E0
ratio_gen = gen_tot / E0
edges = _hist_edges(ratio_real, ratio_gen, bins=60).tolist()
ax.hist(ratio_real, bins=edges, density=True, histtype="step", label="real")
ax.hist(ratio_gen, bins=edges, density=True, histtype="step", label="generated")
ax.axvline(1.0, color="k", linestyle="--", linewidth=1, label="conservation limit (=1)")
ax.set_yscale("log")
ax.set_xlabel("total deposited energy / incident energy, per event")
ax.legend(fontsize=8)
fig.tight_layout()
fig.savefig(OUT / f"{PREFIX}-event-energy-ratio.png", dpi=150, bbox_inches="tight")
print("DONE")
+12 -6
View File
@@ -23,15 +23,23 @@ fig.savefig(OUT / f"{PREFIX}-event-total-length.png", dpi=150, bbox_inches="tigh
print("=== mean/median energy & length per step ===")
fig = a.plot_mean_energy_per_step(obs)
fig.savefig(OUT / f"{PREFIX}-event-mean-energy-per-step.png", dpi=150, bbox_inches="tight")
fig.savefig(
OUT / f"{PREFIX}-event-mean-energy-per-step.png", dpi=150, bbox_inches="tight"
)
fig = a.plot_mean_length_per_step(obs)
fig.savefig(OUT / f"{PREFIX}-event-mean-length-per-step.png", dpi=150, bbox_inches="tight")
fig.savefig(
OUT / f"{PREFIX}-event-mean-length-per-step.png", dpi=150, bbox_inches="tight"
)
print("=== longitudinal / transverse profiles ===")
fig = a.plot_longitudinal_profile(obs)
fig.savefig(OUT / f"{PREFIX}-event-longitudinal-profile.png", dpi=150, bbox_inches="tight")
fig.savefig(
OUT / f"{PREFIX}-event-longitudinal-profile.png", dpi=150, bbox_inches="tight"
)
fig = a.plot_transverse_profile(obs)
fig.savefig(OUT / f"{PREFIX}-event-transverse-profile.png", dpi=150, bbox_inches="tight")
fig.savefig(
OUT / f"{PREFIX}-event-transverse-profile.png", dpi=150, bbox_inches="tight"
)
print("=== shower-max depth ===")
fig = a.plot_shower_max_depth(obs)
@@ -45,8 +53,6 @@ fig = a.plot_pdg_length_share(pdg_table)
fig.savefig(OUT / f"{PREFIX}-pdg-length-share.png", dpi=150, bbox_inches="tight")
print("=== summary stats ===")
import numpy as np # noqa: E402
for label, real_col, gen_col in [
("total_edep", "real_total_edep", "gen_total_edep"),
("total_length", "real_total_length", "gen_total_length"),
+41 -5
View File
@@ -17,8 +17,26 @@ edep_idx = a.RAW_TARGET_NAMES.index("edep")
real = samples.real_raw[mask, edep_idx]
gen = samples.gen_raw[mask, edep_idx]
print("n photon rows:", mask.sum())
print("real: mean", real.mean(), "std", real.std(), "max", real.max(), "frac==0", (real == 0).mean())
print("gen: mean", gen.mean(), "std", gen.std(), "max", gen.max(), "frac==0", (gen == 0).mean())
print(
"real: mean",
real.mean(),
"std",
real.std(),
"max",
real.max(),
"frac==0",
(real == 0).mean(),
)
print(
"gen: mean",
gen.mean(),
"std",
gen.std(),
"max",
gen.max(),
"frac==0",
(gen == 0).mean(),
)
for q in [0.5, 0.9, 0.99, 0.999]:
print(f"q={q}: real={np.quantile(real, q):.4f} gen={np.quantile(gen, q):.4f}")
@@ -30,11 +48,29 @@ axes[0].set_yscale("log")
axes[0].set_xlabel("edep (photons, pdg=22)")
axes[0].legend()
axes[1].hist(real, bins=bins, alpha=0.6, label="real", density=True, cumulative=True, histtype="step")
axes[1].hist(gen, bins=bins, alpha=0.6, label="gen", density=True, cumulative=True, histtype="step")
axes[1].hist(
real,
bins=bins,
alpha=0.6,
label="real",
density=True,
cumulative=True,
histtype="step",
)
axes[1].hist(
gen,
bins=bins,
alpha=0.6,
label="gen",
density=True,
cumulative=True,
histtype="step",
)
axes[1].set_xlabel("edep (photons, pdg=22) - CDF")
axes[1].legend()
fig.tight_layout()
fig.savefig(OUT / "giant-h1024n8d0.1lr3e-4-photon-edep-zoom.png", dpi=150, bbox_inches="tight")
fig.savefig(
OUT / "giant-h1024n8d0.1lr3e-4-photon-edep-zoom.png", dpi=150, bbox_inches="tight"
)
print("saved photon-edep-zoom")
+9 -2
View File
@@ -31,9 +31,16 @@ print("n photon rows:", df.height)
# Process order by abundance, so the legend is stable and the busiest on top.
processes = (
df.group_by("process").len().sort("len", descending=True).get_column("process").to_list()
df.group_by("process")
.len()
.sort("len", descending=True)
.get_column("process")
.to_list()
)
series = {p: df.filter(pl.col("process") == p).get_column("edep_ev").to_numpy() for p in processes}
series = {
p: df.filter(pl.col("process") == p).get_column("edep_ev").to_numpy()
for p in processes
}
all_edep = df.get_column("edep_ev").to_numpy()
lin_bins = np.linspace(0, np.quantile(all_edep, 0.999), 80)
+8 -2
View File
@@ -36,7 +36,9 @@ ax.set_ylabel("density")
ax.legend()
fig.tight_layout()
fig.savefig(OUT / "giant-h1024n8d0.1lr3e-4-photon-edep-ev.png", dpi=150, bbox_inches="tight")
fig.savefig(
OUT / "giant-h1024n8d0.1lr3e-4-photon-edep-ev.png", dpi=150, bbox_inches="tight"
)
print("saved photon-edep-ev")
# Second version: log-spaced energy axis to expose the low-deposition structure.
@@ -53,5 +55,9 @@ ax.set_ylabel("density")
ax.legend()
fig.tight_layout()
fig.savefig(OUT / "giant-h1024n8d0.1lr3e-4-photon-edep-ev-logx.png", dpi=150, bbox_inches="tight")
fig.savefig(
OUT / "giant-h1024n8d0.1lr3e-4-photon-edep-ev-logx.png",
dpi=150,
bbox_inches="tight",
)
print("saved photon-edep-ev-logx")
+14 -2
View File
@@ -91,10 +91,22 @@ axes[0].set_yscale("log")
axes[0].set_xlabel("edep (photons, pdg=22)")
axes[0].legend()
axes[1].hist(
real, bins=bins, alpha=0.6, label="real", density=True, cumulative=True, histtype="step"
real,
bins=bins,
alpha=0.6,
label="real",
density=True,
cumulative=True,
histtype="step",
)
axes[1].hist(
gen, bins=bins, alpha=0.6, label="gen", density=True, cumulative=True, histtype="step"
gen,
bins=bins,
alpha=0.6,
label="gen",
density=True,
cumulative=True,
histtype="step",
)
axes[1].set_xlabel("edep (photons, pdg=22) - CDF")
axes[1].legend()
File diff suppressed because one or more lines are too long