Files
giant/analysis/compare_ode_steps_energy_conservation.py
T
lars 1a22ae022c Add energy-conservation PoC ODE-step comparison scripts
Add analysis scripts for the 10-vs-20 flow-matching ODE-step ablation on the
energy-conservation PoC predict outputs:

- compare_ode_steps_energy_conservation.py: per-event energy-budget table +
  20-step plots and the 10-vs-20 overlay.
- compare_ode_steps_kl.py: per-step marginal KL(real||gen) per target dim over
  fixed shared bins, so the two runs are directly comparable dim-by-dim.

Also commit export_energy_conservation_poc.py (the baseline event-level budget
export) and repoint validation.ipynb at the PoC predict file at sample_frac=1.0.

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
2026-07-06 11:58:25 +02:00

148 lines
5.4 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}")
disp_ratio = lambda r: (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")