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:
+118
-26
@@ -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)
|
||||||
|
|
||||||
return torch.minimum(c_a, c_b).mean()
|
# 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
|
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)
|
||||||
Reference in New Issue
Block a user