MCPcopy Create free account
hub / github.com/Modulus-Labs/RockyBot / train

Function train

pytorch-model/regression_train.py:110–154  ·  view source on GitHub ↗

Trains and evaluates model.

(args, model, train_dataloader, val_dataloader, criterion, opt)

Source from the content-addressed store, hash-verified

108
109
110def 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
157def save_train_stats(train_losses, train_residuals, val_losses, val_residuals, args):

Callers 1

mainFunction · 0.70

Calls 3

eval_modelFunction · 0.70
train_one_epochFunction · 0.70
save_train_statsFunction · 0.70

Tested by

no test coverage detected