Correct the gt matching, rms and loss calculation; use existing pileup training samples; add rms plots

Co-authored-by: Copilot <copilot@github.com>
This commit is contained in:
2026-06-10 11:22:34 +02:00
co-authored by Copilot
parent 346845f6ea
commit 9fc2a32686
+119 -27
View File
@@ -2,6 +2,7 @@ import sys
sys.path.append('./src') sys.path.append('./src')
from omegaconf import OmegaConf ### for yaml config parsing from omegaconf import OmegaConf ### for yaml config parsing
import torch import torch
from torch import nn
import numpy as np import numpy as np
import torch.optim as optim import torch.optim as optim
from tqdm import tqdm from tqdm import tqdm
@@ -39,31 +40,86 @@ def prepare_output_folder(conf):
def get_loss_function(conf): def get_loss_function(conf):
if conf.loss.type == "two_point_set_loss_l2": if conf.loss.type == "two_point_set_loss_l2":
def two_point_set_loss_l2(pred_xy, gt_xy): def two_point_set_loss_l2(pred_xy, gt_xy):
def pair_cost_l2sq(p, q): # p,q: (...,2)
return ((p - q)**2).sum(dim=-1) # squared L2
p1, p2 = pred_xy[:,0], pred_xy[:,1] p1, p2 = pred_xy[:,0], pred_xy[:,1]
g1, g2 = gt_xy[:,0], gt_xy[:,1] g1, g2 = gt_xy[:,0], gt_xy[:,1]
c_a = pair_cost_l2sq(p1,g1) + pair_cost_l2sq(p2,g2)
c_b = pair_cost_l2sq(p1,g2) + pair_cost_l2sq(p2,g1) # 1. matching phase: use L2 distance to determine the matching, regardless of the final loss type
return torch.minimum(c_a, c_b).mean() with torch.no_grad():
c_a = ((p1 - g1)**2).sum(dim=-1) + ((p2 - g2)**2).sum(dim=-1)
c_b = ((p1 - g2)**2).sum(dim=-1) + ((p2 - g1)**2).sum(dim=-1)
swap_mask = (c_b < c_a).unsqueeze(-1)
# 2. Reorganize Ground Truth
g1_matched = torch.where(swap_mask, g2, g1)
g2_matched = torch.where(swap_mask, g1, g2)
# 3. Punishment phase: calculate L2 Loss
loss_p1 = ((p1 - g1_matched)**2).sum(dim=-1)
loss_p2 = ((p2 - g2_matched)**2).sum(dim=-1)
# Divide by 2.0 to make the loss scale represent the "average MSE per photon", fully aligned with RMS logic
return (loss_p1 + loss_p2).mean() / 2.0
return two_point_set_loss_l2 return two_point_set_loss_l2
elif conf.loss.type == "two_point_set_loss_smooth_l1": elif conf.loss.type == "two_point_set_loss_smooth_l1":
def two_point_set_loss_smooth_l1(pred_xy, gt_xy): def two_point_set_loss_smooth_l1(pred_xy, gt_xy):
loss_fn = torch.nn.SmoothL1Loss(reduction='none') loss_fn = nn.SmoothL1Loss(reduction='none', beta=conf.loss.huber_beta)
p1, p2 = pred_xy[:,0], pred_xy[:,1] p1, p2 = pred_xy[:,0], pred_xy[:,1]
g1, g2 = gt_xy[:,0], gt_xy[:,1] g1, g2 = gt_xy[:,0], gt_xy[:,1]
c_a = loss_fn(p1, g1).sum(dim=-1) + loss_fn(p2, g2).sum(dim=-1) # 1. matching phase: use L2 distance to determine the matching, regardless of the final loss type
c_b = loss_fn(p1, g2).sum(dim=-1) + loss_fn(p2, g1).sum(dim=-1) with torch.no_grad():
c_a_l2 = ((p1 - g1)**2).sum(dim=-1) + ((p2 - g2)**2).sum(dim=-1)
c_b_l2 = ((p1 - g2)**2).sum(dim=-1) + ((p2 - g1)**2).sum(dim=-1)
swap_mask = (c_b_l2 < c_a_l2).unsqueeze(-1)
# 2. Reorganize Ground Truth
g1_matched = torch.where(swap_mask, g2, g1)
g2_matched = torch.where(swap_mask, g1, g2)
return torch.minimum(c_a, c_b).mean() # 3. Punishment phase: apply Smooth L1 loss to the correctly matched coordinate pairs
loss_p1 = loss_fn(p1, g1_matched).sum(dim=-1)
loss_p2 = loss_fn(p2, g2_matched).sum(dim=-1)
return (loss_p1 + loss_p2).mean() / 2.0
return two_point_set_loss_smooth_l1 return two_point_set_loss_smooth_l1
def get_mse(pred_xy, gt_xy):
p1, p2 = pred_xy[:,0], pred_xy[:,1] # [B, 2]
g1, g2 = gt_xy[:,0], gt_xy[:,1] # [B, 2]
# 1. matching phase: use L2 distance to determine the matching, regardless of the final loss type
c_a = ((p1 - g1)**2).sum(dim=-1) + ((p2 - g2)**2).sum(dim=-1)
c_b = ((p1 - g2)**2).sum(dim=-1) + ((p2 - g1)**2).sum(dim=-1)
swap_mask = (c_b < c_a).unsqueeze(-1)
g1_matched = torch.where(swap_mask, g2, g1)
g2_matched = torch.where(swap_mask, g1, g2)
# get squared errors for each axis
# p1[:, 0] is x coordinate, p1[:, 1] is y coordinate
err_sq_p1_x = (p1[:, 0] - g1_matched[:, 0])**2
err_sq_p1_y = (p1[:, 1] - g1_matched[:, 1])**2
err_sq_p2_x = (p2[:, 0] - g2_matched[:, 0])**2
err_sq_p2_y = (p2[:, 1] - g2_matched[:, 1])**2
# 2D Euclidean distance squared mean
mse_2d = (err_sq_p1_x + err_sq_p1_y + err_sq_p2_x + err_sq_p2_y).mean() / 2.0
# 1D single-axis squared error mean
mse_x = (err_sq_p1_x + err_sq_p2_x).mean() / 2.0
mse_y = (err_sq_p1_y + err_sq_p2_y).mean() / 2.0
return mse_2d, mse_x, mse_y
def train(model, trainLoader, optimizer, loss_fn): def train(model, trainLoader, optimizer, loss_fn):
model.train() model.train()
batchLoss = 0 batchLoss = 0
sum_squared_error_2d = 0; event_count = 0
sum_squared_error_x = 0; sum_squared_error_y = 0
for batch_idx, (sample, label) in enumerate(trainLoader): for batch_idx, (sample, label) in enumerate(trainLoader):
sample, label = sample.cuda(), label.cuda() sample, label = sample.cuda(), label.cuda()
x1, y1, z1, e1 = label[:,0], label[:,1], label[:,2], label[:,3] x1, y1, z1, e1 = label[:,0], label[:,1], label[:,2], label[:,3]
@@ -76,13 +132,20 @@ def train(model, trainLoader, optimizer, loss_fn):
loss.backward() loss.backward()
optimizer.step() optimizer.step()
batchLoss += loss.item() * sample.shape[0] batchLoss += loss.item() * sample.shape[0]
avgLoss = batchLoss / len(trainLoader.dataset) / 4 ### divide by 4 to get the average loss per photon per axis mse_2d, mse_x, mse_y = get_mse(pred_xy, gt_xy)
print(f"[Train]\t Average Loss: {avgLoss:.6f} (RMS = {np.sqrt(avgLoss):.6f})") sum_squared_error_2d += mse_2d.item() * sample.shape[0]
return avgLoss sum_squared_error_x += mse_x.item() * sample.shape[0]
sum_squared_error_y += mse_y.item() * sample.shape[0]
event_count += sample.shape[0]
avgLoss = batchLoss / len(trainLoader.dataset)
rms_2d = np.sqrt(sum_squared_error_2d / event_count)
print(f"[Train]\t Average Loss: {avgLoss:.6f} (RMS_2d = {rms_2d:.6f}, RMS_x = {np.sqrt(sum_squared_error_x / event_count):.6f}, RMS_y = {np.sqrt(sum_squared_error_y / event_count):.6f})")
return avgLoss, rms_2d.item()
def evaluate(model, valLoader, loss_fn): def evaluate(model, valLoader, loss_fn):
model.eval() model.eval()
batchLoss = 0 batchLoss = 0
sum_squared_error_2d = 0; sum_squared_error_x = 0; sum_squared_error_y = 0; event_count = 0
with torch.no_grad(): with torch.no_grad():
for batch_idx, (sample, label) in enumerate(valLoader): for batch_idx, (sample, label) in enumerate(valLoader):
sample, label = sample.cuda(), label.cuda() sample, label = sample.cuda(), label.cuda()
@@ -93,9 +156,17 @@ def evaluate(model, valLoader, loss_fn):
pred_xy = torch.stack((output[:,0:2], output[:,2:4]), axis=1) pred_xy = torch.stack((output[:,0:2], output[:,2:4]), axis=1)
loss = loss_fn(pred_xy, gt_xy) loss = loss_fn(pred_xy, gt_xy)
batchLoss += loss.item() * sample.shape[0] batchLoss += loss.item() * sample.shape[0]
avgLoss = batchLoss / len(valLoader.dataset) / 4 ### divide by 4 to get the average loss per photon per axis mse_2d, mse_x, mse_y = get_mse(pred_xy, gt_xy)
print(f"[Val]\t Average Loss: {avgLoss:.6f} (RMS = {np.sqrt(avgLoss):.6f})") sum_squared_error_2d += mse_2d.item() * sample.shape[0]
return avgLoss sum_squared_error_x += mse_x.item() * sample.shape[0]
sum_squared_error_y += mse_y.item() * sample.shape[0]
event_count += sample.shape[0]
avgLoss = batchLoss / len(valLoader.dataset)
rms_2d = np.sqrt(sum_squared_error_2d / event_count)
rms_x = np.sqrt(sum_squared_error_x / event_count)
rms_y = np.sqrt(sum_squared_error_y / event_count)
print(f"[Val]\t Average Loss: {avgLoss:.6f} (RMS_2d = {rms_2d:.6f}, RMS_x = {rms_x:.6f}, RMS_y = {rms_y:.6f})")
return avgLoss, rms_2d.item()
def get_dataloaders(conf): def get_dataloaders(conf):
"""construct all dataloaders""" """construct all dataloaders"""
@@ -107,16 +178,12 @@ def get_dataloaders(conf):
file_range_keys = ['train_file_range', 'val_file_range', 'test_file_range'] file_range_keys = ['train_file_range', 'val_file_range', 'test_file_range']
for split, key, batch_key, file_range_key in zip(splits, keys, batch_keys, file_range_keys): for split, key, batch_key, file_range_key in zip(splits, keys, batch_keys, file_range_keys):
files = [f"{conf.data.sample_folder}/{conf.data.energy}keV_Moench040_150V_{i}.npz" for i in range(conf.data[file_range_key][0], conf.data[file_range_key][1] + 1)] files = [f"{conf.data.sample_folder}/pileupOf2phs_sample_{i}.npz" for i in range(conf.data[file_range_key][0], conf.data[file_range_key][1] + 1)]
datasets[split] = doublePhotonDataset( datasets[split] = doublePhotonDataset(
files, files,
sampleRatio = 1.0, sampleRatio = conf.data.sample_ratio,
datasetName = split.capitalize(), datasetName = split.capitalize(),
noiseKeV = conf.data.noise_keV,
nSize = conf.data.n_size,
noiseThreshold = conf.data.noise_threshold * conf.data.noise_keV,
normalize = conf.data.normalize
) )
loaders[split] = torch.utils.data.DataLoader( loaders[split] = torch.utils.data.DataLoader(
@@ -136,17 +203,33 @@ def plot_loss_curves(train_losses, val_losses, test_loss, exp_name, conf):
if test_loss > 0: if test_loss > 0:
plt.axhline(y=test_loss, color='green', linestyle='--', label='Test Loss') plt.axhline(y=test_loss, color='green', linestyle='--', label='Test Loss')
plt.xlabel('Epoch') plt.xlabel('Epoch')
plt.ylabel('MSE Loss') plt.ylabel(f'Loss ({conf.loss.type})')
plt.yscale('log') plt.yscale('log')
plt.legend() plt.legend()
plt.grid() plt.grid()
plotName = f'loss_curve_doublePhoton_{conf.model.version}.png' plotName = f'loss_curve_doublePhoton_{conf.model.version}.png'
plt.savefig(f'Results/{exp_name}/Plots/{plotName}') plt.savefig(f'Results/{exp_name}/Plots/{plotName}')
plt.close()
def plot_rms_curve(train_rms_x, train_rms_y, val_rms_x, val_rms_y, test_rms_x, test_rms_y, exp_name, conf):
import matplotlib.pyplot as plt
plt.figure(figsize=(8,6))
plt.plot(train_rms_x, label='Train RMS X', color='blue')
plt.plot(train_rms_y, label='Train RMS Y', color='cyan')
plt.plot(val_rms_x, label='Val RMS X', color='orange')
plt.plot(val_rms_y, label='Val RMS Y', color='magenta')
if test_rms_x > 0 and test_rms_y > 0:
plt.axhline(y=test_rms_x, color='green', linestyle='--', label='Test RMS X')
plt.axhline(y=test_rms_y, color='lime', linestyle='--', label='Test RMS Y')
plt.xlabel('Epoch')
plt.ylabel('RMS Error [pixels]')
plt.legend()
plt.grid()
plotName = f'rms_curve_doublePhoton_{conf.model.version}'
plt.savefig(f'Results/{exp_name}/Plots/{plotName}.png')
def get_model_name(conf): def get_model_name(conf):
modelName = f'doublePhoton{conf.model.version}_{conf.data.energy}keV_Noise{conf.data.noise_keV}keV' modelName = f'doublePhoton{conf.model.version}_{conf.data.energy}keV_Noise{conf.data.noise_keV}keV'
if conf.data.normalize:
modelName += '_normalized'
return modelName return modelName
if __name__ == "__main__": if __name__ == "__main__":
@@ -161,11 +244,18 @@ if __name__ == "__main__":
trainLoader, valLoader, testLoader = get_dataloaders(conf) trainLoader, valLoader, testLoader = get_dataloaders(conf)
TrainLosses, ValLosses = [], [] TrainLosses, ValLosses = [], []
train_rms_xs, train_rms_ys = [], []
val_rms_xs, val_rms_ys = [], []
for epoch in tqdm(range(1, conf.training.epochs + 1)): for epoch in tqdm(range(1, conf.training.epochs + 1)):
train_loss = train(model, trainLoader, optimizer, loss_fn) train_loss, train_rms = train(model, trainLoader, optimizer, loss_fn)
val_loss = evaluate(model, valLoader, loss_fn) val_loss, val_rms = evaluate(model, valLoader, loss_fn)
TrainLosses.append(train_loss) TrainLosses.append(train_loss)
ValLosses.append(val_loss) ValLosses.append(val_loss)
train_rms_xs.append(train_rms)
train_rms_ys.append(train_rms)
val_rms_xs.append(val_rms)
val_rms_ys.append(val_rms)
scheduler.step(val_loss) scheduler.step(val_loss)
print(f"Learning Rate: {optimizer.param_groups[0]['lr']:.2e}") print(f"Learning Rate: {optimizer.param_groups[0]['lr']:.2e}")
if epoch in conf.training.checkpoint_epochs or epoch == conf.training.epochs: if epoch in conf.training.checkpoint_epochs or epoch == conf.training.epochs:
@@ -173,5 +263,7 @@ if __name__ == "__main__":
torch.save(model.state_dict(), f'Results/{exp_name}/Models/{modelName}_E{epoch}.pth') torch.save(model.state_dict(), f'Results/{exp_name}/Models/{modelName}_E{epoch}.pth')
print(f"Saved model checkpoint: {modelName}_E{epoch}.pth") print(f"Saved model checkpoint: {modelName}_E{epoch}.pth")
plot_loss_curves(TrainLosses, ValLosses, test_loss=-1, exp_name=exp_name, conf=conf) plot_loss_curves(TrainLosses, ValLosses, test_loss=-1, exp_name=exp_name, conf=conf)
plot_rms_curve(train_rms_xs, train_rms_ys, val_rms_xs, val_rms_ys, test_rms_x = -1, test_rms_y = -1, exp_name=exp_name, conf=conf)
test_loss = evaluate(model, testLoader, loss_fn) test_loss = evaluate(model, testLoader, loss_fn)
plot_loss_curves(TrainLosses, ValLosses, test_loss=test_loss, exp_name=exp_name, conf=conf) plot_loss_curves(TrainLosses, ValLosses, test_loss=test_loss, exp_name=exp_name, conf=conf)
plot_rms_curve(train_rms_xs, train_rms_ys, val_rms_xs, val_rms_ys, test_rms_x = -1, test_rms_y = -1, exp_name=exp_name, conf=conf)