Skip to content
Merged
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
68 changes: 56 additions & 12 deletions Python/DeepSSMUtilsPackage/DeepSSMUtils/trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -56,6 +56,19 @@ def log_print(logger, values):
logger.write(log_string + '\n')


# The TL-net stages report different terms, so every stage gets its own columns and leaves the
# columns it does not compute blank. Keep this in sync with log_tl_net().
TL_NET_LOG_HEADER = ["Training_Stage", "Epoch", "LR",
"Train_Err_AE", "Train_Rel_Err_AE", "Val_Err_AE", "Val_Rel_Err_AE",
"Train_Err_TF", "Train_Rel_Err_TF", "Val_Err_TF", "Val_Rel_Err_TF",
"Train_Err_Joint", "Val_Err_Joint", "Sec"]


def log_tl_net(logger, stage, epoch, lr, seconds, ae=("", "", "", ""), tf=("", "", "", ""), joint=("", "")):
# ae and tf are (train_err, train_rel_err, val_err, val_rel_err), joint is (train_err, val_err)
log_print(logger, [stage, epoch, lr, *ae, *tf, *joint, seconds])


def set_scheduler(opt, sched_params):
if sched_params["type"] == "Step":
step_size = sched_params['parameters']['step_size']
Expand All @@ -70,6 +83,16 @@ def set_scheduler(opt, sched_params):
return scheduler


def restart_scheduler(opt, sched_params, learning_rate):
# The TL-net stages train in sequence on one optimizer and each anneals from the base rate, so
# the schedule restarts at every stage boundary. StepLR.get_lr() is relative to the current
# rate, so the group rates have to be restored explicitly; a new scheduler is not enough.
for group in opt.param_groups:
group['lr'] = learning_rate
group['initial_lr'] = learning_rate
return set_scheduler(opt, sched_params)


def train(project, config_file):
net_utils.set_seed(42)
sw.utils.initialize_project_mesh_warper(project)
Expand Down Expand Up @@ -410,6 +433,7 @@ def supervised_train_tl(config_file):
learning_rate = parameters['trainer']['learning_rate']
eval_freq = parameters['trainer']['val_freq']
decay_lr = parameters['trainer']['decay_lr']['enabled']
decay_lr_params = parameters['trainer']['decay_lr']
a_ae = parameters["tl_net"]["a_ae"]
c_ae = parameters["tl_net"]["c_ae"]
a_lat = parameters["tl_net"]["a_lat"]
Expand All @@ -432,8 +456,6 @@ def supervised_train_tl(config_file):
train_params = net.parameters()
opt = torch.optim.Adam(train_params, learning_rate)
opt.zero_grad()
if decay_lr:
scheduler = set_scheduler(opt, parameters['trainer']['decay_lr'])
print("Done.")
# train
print("Beginning training on device = " + device + '\n')
Expand Down Expand Up @@ -461,10 +483,12 @@ def supervised_train_tl(config_file):

# train the AE first

log_print(logger, ["Training_Stage", "Epoch", "LR", "Train_Err", "Train_Rel_Err", "Val_Err", "Val_Rel_Err", "Sec"])
log_print(logger, TL_NET_LOG_HEADER)

mean_mdl = torch.mean(train_loader.dataset.mdl_target, 0).to(device)

if decay_lr:
scheduler = restart_scheduler(opt, decay_lr_params, learning_rate)
for e in range(1, ae_epochs + 1):
if sw_check_abort():
sw_message("Aborted")
Expand Down Expand Up @@ -533,9 +557,8 @@ def supervised_train_tl(config_file):
last_learning_rate = learning_rate
if decay_lr:
last_learning_rate = scheduler.get_last_lr()[0]
log_print(logger,
["AE", e, last_learning_rate, train_ae_err, train_rel_ae_err, val_ae_err, val_rel_ae_err,
time.time() - t0])
log_tl_net(logger, "AE", e, last_learning_rate, time.time() - t0,
ae=(train_ae_err, train_rel_ae_err, val_ae_err, val_rel_ae_err))
# plot
epochs.append(e)
plot_train_losses.append(train_ae_err)
Expand All @@ -546,6 +569,8 @@ def supervised_train_tl(config_file):
train_plot.canvas.draw()
train_plot.savefig(model_dir + C.TRAINING_PLOT_AE_FILE)
t0 = time.time()
if decay_lr:
scheduler.step()
# save
torch.save(net.state_dict(), os.path.join(model_dir, C.FINAL_MODEL_AE_FILE))
# fix the autoencoder and train the TL-net
Expand Down Expand Up @@ -573,6 +598,8 @@ def supervised_train_tl(config_file):
epochs = []
plot_train_losses = []
plot_val_losses = []
if decay_lr:
scheduler = restart_scheduler(opt, decay_lr_params, learning_rate)
for e in range(1, tf_epochs + 1):
if sw_check_abort():
sw_message("Aborted")
Expand Down Expand Up @@ -643,9 +670,8 @@ def supervised_train_tl(config_file):
last_learning_rate = learning_rate
if decay_lr:
last_learning_rate = scheduler.get_last_lr()[0]
log_print(logger,
["T-Flank", e, last_learning_rate, train_tf_err, train_rel_tf_err, val_tf_err, val_rel_tf_err,
time.time() - t0])
log_tl_net(logger, "T-Flank", e, last_learning_rate, time.time() - t0,
tf=(train_tf_err, train_rel_tf_err, val_tf_err, val_rel_tf_err))
# plot
epochs.append(e)
plot_train_losses.append(train_tf_err)
Expand All @@ -656,6 +682,8 @@ def supervised_train_tl(config_file):
train_plot.canvas.draw()
train_plot.savefig(model_dir + C.TRAINING_PLOT_TF_FILE)
t0 = time.time()
if decay_lr:
scheduler.step()
# save
torch.save(net.state_dict(), os.path.join(model_dir, C.FINAL_MODEL_TF_FILE))
# jointly train the model
Expand Down Expand Up @@ -683,6 +711,8 @@ def supervised_train_tl(config_file):
plot_train_losses = []
plot_val_losses = []
t0 = time.time()
if decay_lr:
scheduler = restart_scheduler(opt, decay_lr_params, learning_rate)
for e in range(1, joint_epochs + 1):
if sw_check_abort():
sw_message("Aborted")
Expand All @@ -692,7 +722,9 @@ def supervised_train_tl(config_file):

# train
net.train()
ae_train_losses = []
ae_train_rel_losses = []
tf_train_losses = []
tf_train_rel_losses = []
joint_train_losses = []
pred_particles = []
Expand Down Expand Up @@ -720,7 +752,9 @@ def supervised_train_tl(config_file):
loss.backward()
opt.step()
joint_train_losses.append(loss.item())
ae_train_losses.append(loss_ae.item())
ae_train_rel_losses.append(ae_train_rel_loss.item())
tf_train_losses.append(loss_tf.item())
tf_train_rel_losses.append(tf_train_rel_loss.item())
pred_particles.extend(pred_pt.detach().cpu().numpy())
true_particles.extend(mdl.detach().cpu().numpy())
Expand All @@ -732,7 +766,9 @@ def supervised_train_tl(config_file):
val_names = []
if ((e % eval_freq) == 0 or e == 1):
net.eval()
ae_val_losses = []
ae_val_rel_losses = []
tf_val_losses = []
tf_val_rel_losses = []
joint_val_losses = []
for img, pca, mdl, names in val_loader:
Expand All @@ -748,25 +784,33 @@ def supervised_train_tl(config_file):
tf_val_rel_loss = losses.MSE(lat, lat_img) / losses.MSE(lat * 0, lat_img)

joint_val_losses.append(loss_ae.item() + alpha * loss_tf.item())
ae_val_losses.append(loss_ae.item())
ae_val_rel_losses.append(ae_val_rel_loss.item())
tf_val_losses.append(loss_tf.item())
tf_val_rel_losses.append(tf_val_rel_loss.item())
pred_particles.extend(pred_pt.detach().cpu().numpy())
true_particles.extend(mdl.detach().cpu().numpy())
train_viz.write_examples(np.array(pred_particles), np.array(true_particles), val_names,
model_dir + "examples/validation_")

# log
train_ae_err = np.mean(ae_train_losses)
train_rel_ae_err = np.mean(ae_train_rel_losses)
train_rel_tf_err = np.mean(tf_train_rel_losses)
val_ae_err = np.mean(ae_val_losses)
val_rel_ae_err = np.mean(ae_val_rel_losses)
train_tf_err = np.mean(tf_train_losses)
train_rel_tf_err = np.mean(tf_train_rel_losses)
val_tf_err = np.mean(tf_val_losses)
val_rel_tf_err = np.mean(tf_val_rel_losses)
joint_train_losses = np.mean(joint_train_losses)
joint_val_losses = np.mean(joint_val_losses)
last_learning_rate = learning_rate
if decay_lr:
last_learning_rate = scheduler.get_last_lr()[0]
log_print(logger, ["Joint", e, last_learning_rate, train_rel_ae_err, val_rel_ae_err, train_rel_tf_err,
val_rel_tf_err, time.time() - t0])
log_tl_net(logger, "Joint", e, last_learning_rate, time.time() - t0,
ae=(train_ae_err, train_rel_ae_err, val_ae_err, val_rel_ae_err),
tf=(train_tf_err, train_rel_tf_err, val_tf_err, val_rel_tf_err),
joint=(joint_train_losses, joint_val_losses))
# plot
epochs.append(e)
plot_train_losses.append(joint_train_losses)
Expand Down
Loading