Files
giant/analysis/compare_ode_steps_energy_conservation.py
T
lars a14a4f973a Apply ruff format across the codebase
Whitespace-only reflow (line wrapping, blank lines between defs); no
logic changes.

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
2026-07-08 14:44:53 +02:00

168 lines
5.5 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""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")