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:
+193
-77
File diff suppressed because one or more lines are too long
+170
-3
@@ -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
@@ -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",
|
||||
]
|
||||
|
||||
@@ -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"
|
||||
|
||||
Reference in New Issue
Block a user