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

Function train

pytorch-model/classification_train.py:112–156  ·  view source on GitHub ↗

Trains and evaluates model.

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

Source from the content-addressed store, hash-verified

110
111
112def 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
159def save_train_stats(train_losses, train_accs, val_losses, val_accs, 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