import sys sys.path.append('./src') 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}_2ph_{conf.data.energy}keV_v{conf.model.version}_{exp_index:02d}' if not Path(f'Results/{exp_name}').exists(): break exp_index += 1 Path(f'Results/{exp_name}').mkdir(parents=True, exist_ok=True) Path(f'Results/{exp_name}/Models').mkdir(parents=True, exist_ok=True) Path(f'Results/{exp_name}/Plots').mkdir(parents=True, exist_ok=True) OmegaConf.save(conf, f'Results/{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/{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): 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/{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 = 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 = -1, test_rms_y = -1, exp_name=exp_name, conf=conf)