MCPcopy Create free account
hub / github.com/NVIDIA/semantic-segmentation / main

Function main

train.py:324–462  ·  view source on GitHub ↗

Main Function

()

Source from the content-addressed store, hash-verified

322
323
324def 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)

Callers 1

train.pyFile · 0.70

Calls 15

assert_and_infer_cfgFunction · 0.90
prep_experimentFunction · 0.90
get_lossFunction · 0.90
get_optimizerFunction · 0.90
restore_optFunction · 0.90
restore_netFunction · 0.90
validate_topnFunction · 0.90
update_epochFunction · 0.90
validateFunction · 0.85
trainFunction · 0.85
check_terminationFunction · 0.85
stepMethod · 0.80

Tested by

no test coverage detected