Trains and evaluates model.
(args, model, train_dataloader, val_dataloader, criterion, opt)
| 108 | |
| 109 | |
| 110 | def train(args, model, train_dataloader, val_dataloader, criterion, opt): |
| 111 | """ |
| 112 | Trains and evaluates model. |
| 113 | """ |
| 114 | print("\n--- Begin training! ---\n") |
| 115 | train_losses = dict() |
| 116 | train_residuals = dict() |
| 117 | val_losses = dict() |
| 118 | val_residuals = dict() |
| 119 | model_save_dir = constants.get_model_dir(args.dataset, args.model_type, args.model_name) |
| 120 | viz_path = constants.get_viz_dir(args.dataset, args.model_type, args.model_name) |
| 121 | |
| 122 | for epoch in range(args.num_epochs): |
| 123 | |
| 124 | if epoch % args.eval_every == 0: |
| 125 | # --- Time to evaluate! --- |
| 126 | val_avg_loss, val_avg_residual, _ = eval_model(model, val_dataloader, criterion, args) |
| 127 | val_losses[epoch] = val_avg_loss |
| 128 | val_residuals[epoch] = val_avg_residual |
| 129 | |
| 130 | # --- Report and plot losses/ious --- |
| 131 | print(f" Val avg loss: {val_avg_loss} | Val avg residual: {val_avg_residual}\n") |
| 132 | # viz_utils.plot_losses_ious(val_losses, val_ious, viz_path, prefix="val") |
| 133 | |
| 134 | # --- Train --- |
| 135 | train_avg_loss, train_avg_residual, _ = train_one_epoch(model, train_dataloader, criterion, opt, args) |
| 136 | train_losses[epoch] = train_avg_residual |
| 137 | train_residuals[epoch] = train_avg_loss |
| 138 | |
| 139 | # --- Print results --- |
| 140 | if epoch % args.print_every == 0: |
| 141 | print(f"Epoch: {epoch} | Train avg loss: {train_avg_loss} | Train avg residual: {train_avg_residual}") |
| 142 | |
| 143 | # --- Plot loss/metrics so far (TODO: do this every epoch?) --- |
| 144 | # --- Then save all train stats --- |
| 145 | # viz_utils.plot_losses_ious(train_losses, train_residuals, viz_path, prefix="train") |
| 146 | save_train_stats(train_losses, train_residuals, val_losses, val_residuals, args) |
| 147 | |
| 148 | # --- Only save actual model files every so often --- |
| 149 | if epoch % args.save_every == 0: |
| 150 | model_save_path = os.path.join(model_save_dir, f"model_epoch_{epoch}.pth") |
| 151 | print(f"Saving model to {model_save_path}...") |
| 152 | torch.save(model.state_dict(), model_save_path) |
| 153 | |
| 154 | return train_losses, train_residuals, val_losses, val_residuals |
| 155 | |
| 156 | |
| 157 | def save_train_stats(train_losses, train_residuals, val_losses, val_residuals, args): |
no test coverage detected