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