Files
DeepLearning/Train_2Photon.py
T
2026-06-11 16:22:56 +02:00

273 lines
12 KiB
Python

import sys
sys.path.append('./src')
import matplotlib
matplotlib.use('Agg')
from omegaconf import OmegaConf ### for yaml config parsing
import torch
from torch import nn
import numpy as np
import torch.optim as optim
from tqdm import tqdm
from torchinfo import summary
from pathlib import Path
from models import get_double_photon_model_class
from datasets import doublePhotonDataset
### random seed for reproducibility
torch.manual_seed(0)
torch.cuda.manual_seed(0)
np.random.seed(0)
torch.backends.cudnn.deterministic = True
torch.backends.cudnn.benchmark = False
conf = OmegaConf.load("Configs/train_2photon.yaml")
def prepare_output_folder(conf):
from datetime import datetime
date = datetime.now().strftime("%y%m%d") ## YYMMDD format
# find the next index for experiment name
exp_index = 0
while True:
exp_name = f'{date}_{conf.data.energy}keV_v{conf.model.version}_{exp_index:02d}'
if not Path(f'Results/2ph/{exp_name}').exists():
break
exp_index += 1
Path(f'Results/2ph/{exp_name}').mkdir(parents=True, exist_ok=True)
Path(f'Results/2ph/{exp_name}/Models').mkdir(parents=True, exist_ok=True)
Path(f'Results/2ph/{exp_name}/Plots').mkdir(parents=True, exist_ok=True)
OmegaConf.save(conf, f'Results/2ph/{exp_name}/config.yaml')
return exp_name
def get_loss_function(conf):
if conf.loss.type == "two_point_set_loss_l2":
def two_point_set_loss_l2(pred_xy, gt_xy):
p1, p2 = pred_xy[:,0], pred_xy[:,1]
g1, g2 = gt_xy[:,0], gt_xy[:,1]
# 1. matching phase: use L2 distance to determine the matching, regardless of the final loss type
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
elif conf.loss.type == "two_point_set_loss_smooth_l1":
def two_point_set_loss_smooth_l1(pred_xy, gt_xy):
loss_fn = nn.SmoothL1Loss(reduction='none', beta=conf.loss.huber_beta)
p1, p2 = pred_xy[:,0], pred_xy[:,1]
g1, g2 = gt_xy[:,0], gt_xy[:,1]
# 1. matching phase: use L2 distance to determine the matching, regardless of the final loss type
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)
# 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
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):
model.train()
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):
sample, label = sample.cuda(), label.cuda()
x1, y1, z1, e1 = label[:,0], label[:,1], label[:,2], label[:,3]
x2, y2, z2, e2 = label[:,4], label[:,5], label[:,6], label[:,7]
gt_xy = torch.stack((torch.stack((x1, y1), axis=1), torch.stack((x2, y2), axis=1)), axis=1)
optimizer.zero_grad()
output = model(sample)
pred_xy = torch.stack((output[:,0:2], output[:,2:4]), axis=1)
loss = loss_fn(pred_xy, gt_xy)
loss.backward()
optimizer.step()
batchLoss += loss.item() * sample.shape[0]
mse_2d, mse_x, mse_y = get_mse(pred_xy, gt_xy)
sum_squared_error_2d += mse_2d.item() * sample.shape[0]
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):
model.eval()
batchLoss = 0
sum_squared_error_2d = 0; sum_squared_error_x = 0; sum_squared_error_y = 0; event_count = 0
with torch.no_grad():
for batch_idx, (sample, label) in enumerate(valLoader):
sample, label = sample.cuda(), label.cuda()
x1, y1, z1, e1 = label[:,0], label[:,1], label[:,2], label[:,3]
x2, y2, z2, e2 = label[:,4], label[:,5], label[:,6], label[:,7]
gt_xy = torch.stack((torch.stack((x1, y1), axis=1), torch.stack((x2, y2), axis=1)), axis=1)
output = model(sample)
pred_xy = torch.stack((output[:,0:2], output[:,2:4]), axis=1)
loss = loss_fn(pred_xy, gt_xy)
batchLoss += loss.item() * sample.shape[0]
mse_2d, mse_x, mse_y = get_mse(pred_xy, gt_xy)
sum_squared_error_2d += mse_2d.item() * sample.shape[0]
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):
"""construct all dataloaders"""
datasets = {}
loaders = {}
splits = ['Train', 'Val', 'Test']
keys = ['train_files', 'val_files', 'test_files']
batch_keys = ['batch_size_train', 'batch_size_val', 'batch_size_test']
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):
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(
files,
sampleRatio = conf.data.sample_ratio,
datasetName = split.capitalize(),
)
loaders[split] = torch.utils.data.DataLoader(
datasets[split],
batch_size=conf.data[batch_key],
shuffle=(split=='Train'),
num_workers=conf.data.num_workers,
pin_memory=True
)
return loaders['Train'], loaders['Val'], loaders['Test']
def plot_loss_curves(train_losses, val_losses, test_loss, exp_name, conf):
import matplotlib.pyplot as plt
plt.figure(figsize=(8,6))
plt.plot(train_losses, label='Train Loss')
plt.plot(val_losses, label='Val Loss')
if test_loss > 0:
plt.axhline(y=test_loss, color='green', linestyle='--', label='Test Loss')
plt.xlabel('Epoch')
plt.ylabel(f'Loss ({conf.loss.type})')
plt.yscale('log')
plt.legend()
plt.grid()
plotName = f'loss_curve_doublePhoton_{conf.model.version}.png'
plt.savefig(f'Results/2ph/{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/2ph/{exp_name}/Plots/{plotName}.png')
plt.close()
def get_model_name(conf):
modelName = f'doublePhoton{conf.model.version}_{conf.data.energy}keV_Noise{conf.data.noise_keV}keV'
return modelName
if __name__ == "__main__":
exp_name = prepare_output_folder(conf)
model = get_double_photon_model_class(conf.model.version)().cuda()
# summary(model, input_size=(128, 3, conf.data.n_size, conf.data.n_size))
loss_fn = get_loss_function(conf)
optimizer = torch.optim.Adam(model.parameters(), lr=conf.training.learning_rate, weight_decay=conf.training.weight_decay)
scheduler = optim.lr_scheduler.ReduceLROnPlateau(optimizer, 'min', factor=conf.training.scheduler_factor, patience=conf.training.scheduler_patience)
trainLoader, valLoader, testLoader = get_dataloaders(conf)
TrainLosses, ValLosses = [], []
train_rms_xs, train_rms_ys = [], []
val_rms_xs, val_rms_ys = [], []
for epoch in tqdm(range(1, conf.training.epochs + 1)):
train_loss, train_rms = train(model, trainLoader, optimizer, loss_fn)
val_loss, val_rms = evaluate(model, valLoader, loss_fn)
TrainLosses.append(train_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)
print(f"Learning Rate: {optimizer.param_groups[0]['lr']:.2e}")
if epoch in conf.training.checkpoint_epochs or epoch == conf.training.epochs:
modelName = get_model_name(conf)
torch.save(model.state_dict(), f'Results/2ph/{exp_name}/Models/{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_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, rms_2d_test = evaluate(model, testLoader, loss_fn)
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 = rms_2d_test[0], test_rms_y = rms_2d_test[1], exp_name=exp_name, conf=conf)