Add mean/median deposited energy and step length plots per event

Move ipykernel into the analysis extra instead of a separate
dependency group, since it's needed wherever analysis plotting runs.

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
2026-06-24 12:46:17 +02:00
parent c802a59033
commit b4ce04e772
4 changed files with 366 additions and 93 deletions
+193 -77
View File
File diff suppressed because one or more lines are too long
+170 -3
View File
@@ -40,6 +40,8 @@ corrupted by partial events::
obs = compute_event_observables_pl("path/to/steps_predicted_local.parquet")
plot_total_energy(obs)
plot_total_length(obs)
plot_mean_energy_per_step(obs)
plot_mean_length_per_step(obs)
plot_longitudinal_profile(obs)
plot_transverse_profile(obs)
plot_shower_max_depth(obs)
@@ -610,7 +612,9 @@ def plot_marginals(
table = marginal_table(
collection, group_by=group_by, n_energy_bins=n_energy_bins, bins=bins
)
by_group = table.groupby("group").agg(kl_max=("kl_real_gen", "max"), n=("n", "first"))
by_group = table.groupby("group").agg(
kl_max=("kl_real_gen", "max"), n=("n", "first")
)
worst_first = (by_group["kl_max"] * by_group["n"]).sort_values(ascending=False)
order = {label: rank for rank, label in enumerate(worst_first.index)}
groups = [g for g in groups if g[0] in order]
@@ -629,7 +633,9 @@ def plot_marginals(
ax = axes[row][col]
edges = _hist_edges(real[:, j], gen[:, j], bins=bins)
ax.hist(real[:, j], bins=edges, density=True, histtype="step", label="real")
ax.hist(gen[:, j], bins=edges, density=True, histtype="step", label="generated")
ax.hist(
gen[:, j], bins=edges, density=True, histtype="step", label="generated"
)
ax.set_yscale("log")
if row == 0:
ax.set_title(dims[col], fontsize=9)
@@ -648,7 +654,9 @@ def _limit_groups_by_kl_n(table: pd.DataFrame, max_groups: int) -> pd.DataFrame:
alone, so a rare pdg/material with a noisy, high-variance KL estimate from
a handful of samples doesn't crowd out groups that actually matter.
"""
by_group = table.groupby("group").agg(kl_max=("kl_real_gen", "max"), n=("n", "first"))
by_group = table.groupby("group").agg(
kl_max=("kl_real_gen", "max"), n=("n", "first")
)
worst_first = (by_group["kl_max"] * by_group["n"]).sort_values(ascending=False)
keep = set(worst_first.index[:max_groups])
return table[table["group"].isin(keep)]
@@ -1202,6 +1210,50 @@ def compute_event_observables_pl(
real_max_depth = depth_centers[np.argmax(real_depth_bin_edep, axis=1)]
gen_max_depth = depth_centers[np.argmax(gen_depth_bin_edep, axis=1)]
safe_n_steps = np.where(n_steps > 0, n_steps, 1)
real_mean_edep = real_total_edep / safe_n_steps
gen_mean_edep = gen_total_edep / safe_n_steps
real_mean_length = real_total_length / safe_n_steps
gen_mean_length = gen_total_length / safe_n_steps
# The inverse transform must be applied *before* the median, not after:
# median commutes with inv_log_transform only for odd-length groups. For an
# even number of steps polars' .median() averages the two central order
# statistics, and that averaging does not commute with the nonlinear exp
# (it would yield their geometric mean rather than the true median). So take
# the median on the raw (exp-transformed) values directly. No post_pos
# reconstruction is needed — only the magnitudes matter.
medians = (
lf.select(
"event_id",
"true_log_edep",
"pred_log_edep",
"true_log_step_length",
"pred_log_step_length",
)
.group_by("event_id")
.agg(
(pl.col("true_log_edep").exp() - _LOG_EPS)
.median()
.alias("real_median_edep"),
(pl.col("pred_log_edep").exp() - _LOG_EPS)
.median()
.alias("gen_median_edep"),
(pl.col("true_log_step_length").exp() - _LOG_EPS)
.median()
.alias("real_median_length"),
(pl.col("pred_log_step_length").exp() - _LOG_EPS)
.median()
.alias("gen_median_length"),
)
.sort("event_id")
.collect()
)
real_median_edep = medians["real_median_edep"].to_numpy()
gen_median_edep = medians["gen_median_edep"].to_numpy()
real_median_length = medians["real_median_length"].to_numpy()
gen_median_length = medians["gen_median_length"].to_numpy()
event_table = pl.DataFrame(
{
"event_id": event_ids,
@@ -1210,6 +1262,14 @@ def compute_event_observables_pl(
"gen_total_edep": gen_total_edep,
"real_total_length": real_total_length,
"gen_total_length": gen_total_length,
"real_mean_edep": real_mean_edep,
"gen_mean_edep": gen_mean_edep,
"real_mean_length": real_mean_length,
"gen_mean_length": gen_mean_length,
"real_median_edep": real_median_edep,
"gen_median_edep": gen_median_edep,
"real_median_length": real_median_length,
"gen_median_length": gen_median_length,
"real_centroid_depth": real_centroid_depth,
"gen_centroid_depth": gen_centroid_depth,
"real_transverse_rms": real_transverse_rms,
@@ -1292,6 +1352,113 @@ def plot_total_length(observables: EventObservables, bins: int = 50):
return fig
def _plot_mean_median_per_step(
real_mean: np.ndarray,
gen_mean: np.ndarray,
real_median: np.ndarray,
gen_median: np.ndarray,
mean_xlabel: str,
median_xlabel: str,
bins: int,
median_bins: int,
):
"""Side-by-side real-vs-generated histograms: per-event mean (left), median (right).
The mean is pulled down by a compressed/under-sampled right tail (rare
large values), while the median is robust to that tail. Splitting them into
separate panels (rather than overlaying on one axis) keeps each comparison
legible; if a tail-compression bias is the explanation for a low generated
mean, the median panel should overlap far more closely than the mean panel.
Each panel gets its own bin edges (`bins` for the mean, `median_bins` for
the median): the median is far less spread than the mean, so reusing the
mean's range would crush it into a few bins — `median_bins` is larger to
resolve that tighter distribution.
"""
mean_edges = _hist_edges(real_mean, gen_mean, bins=bins)
median_edges = _hist_edges(real_median, gen_median, bins=median_bins)
prop_colors = plt.rcParams["axes.prop_cycle"].by_key()["color"]
real_color, gen_color = prop_colors[0], prop_colors[1]
fig, (ax_mean, ax_median) = plt.subplots(1, 2, figsize=(12, 4))
for ax, real, gen, edges, title, xlabel in [
(ax_mean, real_mean, gen_mean, mean_edges, "mean", mean_xlabel),
(ax_median, real_median, gen_median, median_edges, "median", median_xlabel),
]:
ax.hist(
real,
bins=edges,
density=True,
histtype="step",
color=real_color,
label=f"real (σ/μ={real.std() / real.mean():.3f})",
)
ax.hist(
gen,
bins=edges,
density=True,
histtype="step",
color=gen_color,
label=f"generated (σ/μ={gen.std() / gen.mean():.3f})",
)
ax.set_yscale("log")
ax.set_title(title)
ax.set_xlabel(xlabel)
ax.legend(fontsize=8)
fig.tight_layout()
return fig
def plot_mean_energy_per_step(
observables: EventObservables, bins: int = 50, median_bins: int = 150
):
"""Real-vs-generated histograms of mean/median deposited energy per step, per event.
Per event: `total_edep / n_steps` (mean) and the per-step median edep —
distinct from `plot_total_energy`, which histograms the per-event *total*;
these instead ask whether the typical step's energy deposit is right,
independent of how many steps the event happened to have. Mean and median
are drawn in separate panels, each with its own binning (`bins` /
`median_bins`); see `_plot_mean_median_per_step`.
"""
table = observables.event_table
return _plot_mean_median_per_step(
table["real_mean_edep"].to_numpy(),
table["gen_mean_edep"].to_numpy(),
table["real_median_edep"].to_numpy(),
table["gen_median_edep"].to_numpy(),
mean_xlabel="mean deposited energy per step, per event [MeV]",
median_xlabel="median deposited energy per step, per event [MeV]",
bins=bins,
median_bins=median_bins,
)
def plot_mean_length_per_step(
observables: EventObservables, bins: int = 50, median_bins: int = 150
):
"""Real-vs-generated histograms of mean/median step length per step, per event.
Per event: `total_length / n_steps` (mean) and the per-step median step
length — distinct from `plot_total_length`, which histograms the per-event
*total*; these instead ask whether the typical step length is right,
independent of how many steps the event happened to have. Mean and median
are drawn in separate panels, each with its own binning (`bins` /
`median_bins`); see `_plot_mean_median_per_step`.
"""
table = observables.event_table
return _plot_mean_median_per_step(
table["real_mean_length"].to_numpy(),
table["gen_mean_length"].to_numpy(),
table["real_median_length"].to_numpy(),
table["gen_median_length"].to_numpy(),
mean_xlabel="mean step length per step, per event [mm]",
median_xlabel="median step length per step, per event [mm]",
bins=bins,
median_bins=median_bins,
)
def _plot_profile(
centers: np.ndarray,
real_mean: np.ndarray,
+1 -5
View File
@@ -32,6 +32,7 @@ convert = [
analysis = [
"matplotlib>=3.8,<4",
"polars>=1.0,<2",
"ipykernel>=7.3.0",
]
[project.scripts]
@@ -67,8 +68,3 @@ explicit = true
name = "pytorch-cu118"
url = "https://download.pytorch.org/whl/cu118"
explicit = true
[dependency-groups]
analysis = [
"ipykernel>=7.3.0",
]
Generated
+2 -8
View File
@@ -446,6 +446,7 @@ dependencies = [
[package.optional-dependencies]
analysis = [
{ name = "ipykernel" },
{ name = "matplotlib" },
{ name = "polars" },
]
@@ -467,14 +468,10 @@ dev = [
{ name = "ty" },
]
[package.dev-dependencies]
analysis = [
{ name = "ipykernel" },
]
[package.metadata]
requires-dist = [
{ name = "awkward", marker = "extra == 'convert'", specifier = ">=2.6,<3" },
{ name = "ipykernel", marker = "extra == 'analysis'", specifier = ">=7.3.0" },
{ name = "matplotlib", marker = "extra == 'analysis'", specifier = ">=3.8,<4" },
{ name = "numpy", specifier = ">=1.26,<3" },
{ name = "pandas", specifier = ">=2.2,<4" },
@@ -492,9 +489,6 @@ requires-dist = [
]
provides-extras = ["cpu", "cuda", "dev", "convert", "analysis"]
[package.metadata.requires-dev]
analysis = [{ name = "ipykernel", specifier = ">=7.3.0" }]
[[package]]
name = "iniconfig"
version = "2.3.0"