Validation loss curve PyTorch - how to store all the losses while training not only last?

Viewed 229

I'm very new in this field. The original code is from this GitHub: https://github.com/YuliaRubanova/latent_ode (run_models.py) file. I want to draw a learning curve using the validation loss. The code part for training:

# Training

log_path = "logs/" + file_name + "_" + str(experimentID) + ".log"
if not os.path.exists("logs/"):
    utils.makedirs("logs/")
logger = utils.get_logger(logpath=log_path, filepath=os.path.abspath(__file__))
logger.info(input_command)

optimizer = optim.Adamax(model.parameters(), lr=args.lr)

num_batches = data_obj["n_train_batches"]

for itr in range(1, num_batches * (args.niters + 1)):
    optimizer.zero_grad()
    utils.update_learning_rate(optimizer, decay_rate = 0.999, lowest = args.lr / 10)

    wait_until_kl_inc = 10
    if itr // num_batches < wait_until_kl_inc:
        kl_coef = 0.
    else:
        kl_coef = (1-0.99** (itr // num_batches - wait_until_kl_inc))

    batch_dict = utils.get_next_batch(data_obj["train_dataloader"])
    train_res = model.compute_all_losses(batch_dict, n_traj_samples = 3, kl_coef = kl_coef)
    train_res["loss"].backward()
    optimizer.step()

    n_iters_to_viz = 1
    if itr % (n_iters_to_viz * num_batches) == 0:
        with torch.no_grad():

            test_res = compute_loss_all_batches(model, 
                data_obj["test_dataloader"], args,
                n_batches = data_obj["n_test_batches"],
                experimentID = experimentID,
                device = device,
                n_traj_samples = 3, kl_coef = kl_coef)

            message = 'Epoch {:04d} [Test seq (cond on sampled tp)] | Loss {:.6f} | Likelihood {:.6f} | KL fp {:.4f} | FP STD {:.4f}|'.format(
                itr//num_batches, 
                test_res["loss"].detach(), test_res["likelihood"].detach(), 
                test_res["kl_first_p"], test_res["std_first_p"])
        
            logger.info("Experiment " + str(experimentID))
            logger.info(message)
            logger.info("KL coef: {}".format(kl_coef))
            logger.info("Train loss (one batch): {}".format(train_res["loss"].detach()))
            logger.info("Train CE loss (one batch): {}".format(train_res["ce_loss"].detach()))
            
            if "auc" in test_res:
                logger.info("Classification AUC (TEST): {:.4f}".format(test_res["auc"]))

            if "mse" in test_res:
                logger.info("Test MSE: {:.4f}".format(test_res["mse"]))

            if "accuracy" in train_res:
                logger.info("Classification accuracy (TRAIN): {:.4f}".format(train_res["accuracy"]))

            if "accuracy" in test_res:
                logger.info("Classification accuracy (TEST): {:.4f}".format(test_res["accuracy"]))

            if "pois_likelihood" in test_res:
                logger.info("Poisson likelihood: {}".format(test_res["pois_likelihood"]))

            if "ce_loss" in test_res:
                logger.info("CE loss: {}".format(test_res["ce_loss"]))

        torch.save({
            'args': args,
            'state_dict': model.state_dict(),
        }, ckpt_path)

Also I added this part for plotting the loss curve:

import matplotlib.pyplot as plt

import seaborn as sns

# Use plot styling from seaborn.
sns.set(style='darkgrid')

      # Increase the plot size and font size.
sns.set(font_scale=1.5)
plt.rcParams["figure.figsize"] = (12,6)

      # Plot the learning curve.
plt.plot(test_res["mse"], 'r-o')
      #plt.plot(tsloss_val2, 'c-o')

plt.legend(["RNN"])
      # Label the plot.
plt.title("Validation loss")
plt.xlabel("Epoch")
plt.ylabel("Loss") 

plt.savefig('/content/latent_ode/foo.png')

The problem I'm getting a curve only for the final loss, but not for all iterations. My curve: enter image description here

What should be added to the code to get standard learning curve with losses during all epochs, not only from the last one? Like this (from other project): enter image description here

1 Answers

save losses in a list (define a list before for scope on epochs , and append each loss.item to list). after that, draw plot

Related