()
| 176 | |
| 177 | |
| 178 | def 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 |
no test coverage detected