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

Function main

pytorch-model/classification_train.py:178–246  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

176
177
178def main():
179
180 # --- Check for GPU ---
181 device = "cuda:0" if torch.cuda.is_available() else "cpu"
182
183 # --- Args ---
184 args = opts.get_train_args()
185 print("\n" + "-" * 30 + " Args " + "-" * 30)
186 for k, v in vars(args).items():
187 print(f"{k}: {v}")
188 print()
189
190 # --- Model and viz save dir ---
191 model_save_dir = constants.get_model_dir(args.dataset, args.model_type, args.model_name)
192 viz_save_dir = constants.get_viz_dir(args.dataset, args.model_type, args.model_name)
193 if os.path.isdir(model_save_dir):
194 raise RuntimeError(f"Error: {model_save_dir} already exists! Exiting...")
195 elif os.path.isdir(viz_save_dir):
196 raise RuntimeError(f"Error: {viz_save_dir} already exists! Exiting...")
197 else:
198 print(f"--> Creating directory {model_save_dir}...")
199 os.makedirs(model_save_dir)
200 print(f"--> Creating directory {viz_save_dir}...")
201 os.makedirs(viz_save_dir)
202 print("Done!\n")
203
204 # --- Setup dataset ---
205 print("--> Setting up dataset...")
206 train_dataset = datasets.DATASETS[args.dataset](mode="train")
207 val_dataset = datasets.DATASETS[args.dataset](mode="val")
208 print("Done!\n")
209
210 # --- Dataloaders ---
211 print("--> Setting up dataloaders...")
212 train_dataloader = DataLoader(train_dataset, batch_size=args.batch_size, shuffle=True, num_workers=1)
213 val_dataloader = DataLoader(val_dataset, batch_size=args.batch_size, shuffle=False, num_workers=1)
214 print("Done!\n")
215
216 # --- Setup model ---
217 print("--> Setting up model...")
218 model = models.MODEL_TYPES[args.model_type](train_dataset)
219 # torch.cuda.set_device(device)
220 # model = model.cuda(device)
221 print("Done!\n")
222
223 # --- Optimizer ---
224 print("--> Setting up optimizer/criterion...")
225 opt = torch.optim.Adam(model.parameters(), lr=args.lr)
226
227 # --- Loss fn ---
228 criterion = nn.CrossEntropyLoss(weight=train_dataset.get_weights())#.cuda(constants.GPU)
229 # criterion = nn.CrossEntropyLoss()#.cuda(constants.GPU)
230 print("Done!\n")
231
232 # --- Train ---
233 train_losses, train_accs, val_losses, val_accs =\
234 train(args, model, train_dataloader, val_dataloader, criterion, opt)
235

Callers 1

Calls 3

trainFunction · 0.70
save_train_statsFunction · 0.70
get_weightsMethod · 0.45

Tested by

no test coverage detected