Main Function
()
| 322 | |
| 323 | |
| 324 | def main(): |
| 325 | """ |
| 326 | Main Function |
| 327 | """ |
| 328 | if AutoResume: |
| 329 | AutoResume.init() |
| 330 | |
| 331 | assert args.result_dir is not None, 'need to define result_dir arg' |
| 332 | logx.initialize(logdir=args.result_dir, |
| 333 | tensorboard=True, hparams=vars(args), |
| 334 | global_rank=args.global_rank) |
| 335 | |
| 336 | # Set up the Arguments, Tensorboard Writer, Dataloader, Loss Fn, Optimizer |
| 337 | assert_and_infer_cfg(args) |
| 338 | prep_experiment(args) |
| 339 | train_loader, val_loader, train_obj = \ |
| 340 | datasets.setup_loaders(args) |
| 341 | criterion, criterion_val = get_loss(args) |
| 342 | |
| 343 | auto_resume_details = None |
| 344 | if AutoResume: |
| 345 | auto_resume_details = AutoResume.get_resume_details() |
| 346 | |
| 347 | if auto_resume_details: |
| 348 | checkpoint_fn = auto_resume_details.get("RESUME_FILE", None) |
| 349 | checkpoint = torch.load(checkpoint_fn, |
| 350 | map_location=torch.device('cpu')) |
| 351 | args.result_dir = auto_resume_details.get("TENSORBOARD_DIR", None) |
| 352 | args.start_epoch = int(auto_resume_details.get("EPOCH", None)) + 1 |
| 353 | args.restore_net = True |
| 354 | args.restore_optimizer = True |
| 355 | msg = ("Found details of a requested auto-resume: checkpoint={}" |
| 356 | " tensorboard={} at epoch {}") |
| 357 | logx.msg(msg.format(checkpoint_fn, args.result_dir, |
| 358 | args.start_epoch)) |
| 359 | elif args.resume: |
| 360 | checkpoint = torch.load(args.resume, |
| 361 | map_location=torch.device('cpu')) |
| 362 | args.arch = checkpoint['arch'] |
| 363 | args.start_epoch = int(checkpoint['epoch']) + 1 |
| 364 | args.restore_net = True |
| 365 | args.restore_optimizer = True |
| 366 | msg = "Resuming from: checkpoint={}, epoch {}, arch {}" |
| 367 | logx.msg(msg.format(args.resume, args.start_epoch, args.arch)) |
| 368 | elif args.snapshot: |
| 369 | if 'ASSETS_PATH' in args.snapshot: |
| 370 | args.snapshot = args.snapshot.replace('ASSETS_PATH', cfg.ASSETS_PATH) |
| 371 | checkpoint = torch.load(args.snapshot, |
| 372 | map_location=torch.device('cpu')) |
| 373 | args.restore_net = True |
| 374 | msg = "Loading weights from: checkpoint={}".format(args.snapshot) |
| 375 | logx.msg(msg) |
| 376 | |
| 377 | net = network.get_net(args, criterion) |
| 378 | optim, scheduler = get_optimizer(args, net) |
| 379 | |
| 380 | if args.fp16: |
| 381 | net, optim = amp.initialize(net, optim, opt_level=args.amp_opt_level) |
no test coverage detected