Files
MLXIDDemo/scripts/03_eval_mc_truth.py
xie_x1andClaude Sonnet 5 2683fd3576 Initial commit: MLXIDDemo, a 12 keV end-to-end walkthrough of the MLXID pipeline
MC generation -> pile-up assembly -> CNN vs. classical eta-interpolation
(scored against MC ground truth) -> optional training -> qualitative
inference on real measurement data, self-contained and runnable from a
single conda env (see README.md).

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
2026-08-21 16:50:06 +02:00

127 lines
5.6 KiB
Python

"""
Score the three pretrained 12 keV models (1/2/3-photon) directly against MC
ground truth. For pile-up, the network predicts an unordered set of points,
so each event is scored against whichever ground-truth permutation
minimizes squared error (same idea as Train_2Photon.py / Train_3Photon.py's
set loss, generalized to N=1..3 here).
For single photon only, also scores classical eta interpolation (Rosenblatt
LUT, see EtaInterpolation/EtaInterpolationFunctions.py) against the same
truth, both *before* interpolation (raw eta value used directly as a
position) and *after* it -- showing what that step buys you -- alongside
the CNN. Eta interpolation has no multi-photon extension in this codebase,
so this comparison is 1-photon only. The LUT is fit on this same dataset
(no train/test split): in real usage there's no separate "training set" for
eta interpolation, a measurement interpolates itself from its own
charge-sharing statistics.
Needs: torch, numpy. Double/triple-photon models require CUDA (hard-coded
in src/models.py). Use the `mlxid_demo` conda env.
"""
import sys
import itertools
from pathlib import Path
import numpy as np
import torch
HERE = Path(__file__).resolve().parent.parent
sys.path.append(str(HERE / 'src'))
from datasets import singlePhotonDataset, doublePhotonDataset, triplePhotonDataset
from eta_interpolation_functions import build_xy_lut_Rosenblatt, bilinear_xy_lookup
from model_zoo import load_pretrained
DEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'
N_ETA_BINS = 101 ### smaller than production's 201 -- these demo samples are much smaller
POSITION_WEIGHTS = np.array([-1, 0, 1]) + 0.5 ### same convention as EtaInterpolation/etaInterpolationFromClusters.py
def matched_squared_error(pred_xy, gt_xy):
"""pred_xy, gt_xy: (B, N, 2) tensors of unordered points -> (B,) best-match summed sq. error."""
n = pred_xy.shape[1]
best = None
for perm in itertools.permutations(range(n)):
d = ((pred_xy - gt_xy[:, perm, :]) ** 2).sum(dim=(1, 2))
best = d if best is None else torch.minimum(best, d)
return best
def compute_eta(clusters_3x3):
sums = clusters_3x3.sum(axis=(1, 2))
etaX = np.clip((clusters_3x3.sum(axis=1) * POSITION_WEIGHTS).sum(axis=-1) / sums, 0, 1)
etaY = np.clip((clusters_3x3.sum(axis=2) * POSITION_WEIGHTS).sum(axis=-1) / sums, 0, 1)
return etaX, etaY
def eta_interpolate(etaX, etaY):
"""Fit the Rosenblatt LUT on these same events and map eta -> interpolated sub-pixel (x, y)."""
hist2D, _, _ = np.histogram2d(etaX, etaY, bins=N_ETA_BINS, range=[[0, 1], [0, 1]])
U_tab, V_tab = build_xy_lut_Rosenblatt(hist2D)
return bilinear_xy_lookup(etaX, etaY, U_tab, V_tab)
def eval_single_photon():
files = sorted(str(f) for f in (HERE / 'data/mc_samples').glob('12keV_Moench040_150V_*.npz'))
### classical eta interpolation, before vs. after the Rosenblatt LUT
all_samples, all_labels = [], []
for f in files:
d = np.load(f)
all_samples.append(d['samples'])
all_labels.append(d['labels'])
samples5 = np.concatenate(all_samples, axis=0) # (N, 5, 5)
labels = np.concatenate(all_labels, axis=0).copy()
samples3 = samples5[:, 1:4, 1:4] # (N, 3, 3), matching the CNN's 3x3 input
gt_xy = labels[:, :2] - 1.0 # shift to the 3x3 crop's coordinate frame
etaX, etaY = compute_eta(samples3)
raw_xy = np.stack([etaX + 1, etaY + 1], axis=1) ### before interpolation: eta used directly as position
raw_rms = np.sqrt(np.mean(np.sum((raw_xy - gt_xy) ** 2, axis=1)))
x_frac, y_frac = eta_interpolate(etaX, etaY)
interpolated_xy = np.stack([x_frac + 1, y_frac + 1], axis=1) ### after interpolation
interpolated_rms = np.sqrt(np.mean(np.sum((interpolated_xy - gt_xy) ** 2, axis=1)))
### CNN
dataset = singlePhotonDataset(files, sampleRatio=1.0, datasetName='MC-1ph-truth')
loader = torch.utils.data.DataLoader(dataset, batch_size=4096, shuffle=False)
model = load_pretrained(1, DEVICE)
sq_err, n = 0.0, 0
with torch.no_grad():
for sample, label in loader:
sample, label = sample.to(DEVICE), label.to(DEVICE)
pred_xy = model(sample)[:, :2]
sq_err += ((pred_xy - label[:, :2]) ** 2).sum().item()
n += sample.size(0)
cnn_rms = (sq_err / n) ** 0.5
# print(f'[1ph] Eta, before interpolation: {raw_rms:.4f} pixels RMS vs truth')
print(f'[1ph] Eta interpolation: {interpolated_rms:.4f} pixels RMS vs truth(n={n} events)')
print(f'[1ph] Deep learning: {cnn_rms:.4f} pixels RMS vs truth (n={n} events)')
def eval_pileup(n_photons):
sample_size = {2: 6, 3: 9}[n_photons]
dataset_cls = {2: doublePhotonDataset, 3: triplePhotonDataset}[n_photons]
f = HERE / f'data/pileup_samples/{n_photons}photon_pileup_12keV_demo.npz'
dataset = dataset_cls([str(f)], sampleRatio=1.0, datasetName=f'MC-{n_photons}ph-truth')
loader = torch.utils.data.DataLoader(dataset, batch_size=1024, shuffle=False)
model = load_pretrained(n_photons, DEVICE)
sq_err = 0.0
n = 0
with torch.no_grad():
for sample, label in loader:
sample, label = sample.to(DEVICE), label.to(DEVICE)
pred = model(sample).view(-1, n_photons, 2)
gt_xy = label[:, :, :2] if label.dim() == 3 else label.view(-1, n_photons, 4)[:, :, :2]
sq_err += matched_squared_error(pred, gt_xy).sum().item()
n += sample.size(0)
rms = (sq_err / (n * n_photons)) ** 0.5
print(f'[{n_photons}ph] Deep learning: {rms:.4f} pixels (n={n} events)')
if __name__ == '__main__':
print(f'[03] Using device: {DEVICE}')
eval_single_photon()
eval_pileup(2)
eval_pileup(3)