MCPcopy Create free account
hub / github.com/OpenGVLab/UniFormerV2 / train

Function train

tools/train_net.py:387–538  ·  view source on GitHub ↗

Train a video model for many epochs on train set and evaluate it on val set. Args: cfg (CfgNode): configs. Details can be found in slowfast/config/defaults.py

(cfg)

Source from the content-addressed store, hash-verified

385
386
387def train(cfg):
388 """
389 Train a video model for many epochs on train set and evaluate it on val set.
390 Args:
391 cfg (CfgNode): configs. Details can be found in
392 slowfast/config/defaults.py
393 """
394 # Set up environment.
395 du.init_distributed_training(cfg)
396 # Set random seed from configs.
397 np.random.seed(cfg.RNG_SEED)
398 torch.manual_seed(cfg.RNG_SEED)
399
400 # Setup logging format.
401 logging.setup_logging(cfg.OUTPUT_DIR)
402
403 # Init multigrid.
404 multigrid = None
405 if cfg.MULTIGRID.LONG_CYCLE or cfg.MULTIGRID.SHORT_CYCLE:
406 multigrid = MultigridSchedule()
407 cfg = multigrid.init_multigrid(cfg)
408 if cfg.MULTIGRID.LONG_CYCLE:
409 cfg, _ = multigrid.update_long_cycle(cfg, cur_epoch=0)
410 # Print config.
411 logger.info("Train with config:")
412 logger.info(pprint.pformat(cfg))
413
414 # Loss scaler
415 loss_scaler = NativeScaler()
416
417 # Build the video model and print model statistics.
418 model = build_model(cfg)
419 if du.is_master_proc() and cfg.LOG_MODEL_INFO:
420 misc.log_model_info(model, cfg, use_train_input=True)
421
422 # Construct the optimizer.
423 optimizer = optim.construct_optimizer(model, cfg)
424
425 # Load a checkpoint to resume training if applicable.
426 start_epoch = cu.load_train_checkpoint(cfg, model, optimizer, loss_scaler)
427
428 # Create the video train and val loaders.
429 train_loader = loader.construct_loader(cfg, "train")
430 val_loader = loader.construct_loader(cfg, "val")
431 precise_bn_loader = (
432 loader.construct_loader(cfg, "train", is_precise_bn=True)
433 if cfg.BN.USE_PRECISE_STATS
434 else None
435 )
436
437 # Create meters.
438 if cfg.DETECTION.ENABLE:
439 train_meter = AVAMeter(len(train_loader), cfg, mode="train")
440 val_meter = AVAMeter(len(val_loader), cfg, mode="val")
441 else:
442 train_meter = TrainMeter(len(train_loader), cfg)
443 val_meter = ValMeter(len(val_loader), cfg)
444

Callers

nothing calls this directly

Calls 15

init_multigridMethod · 0.95
update_long_cycleMethod · 0.95
epoch_ticMethod · 0.95
epoch_tocMethod · 0.95
last_epoch_timeMethod · 0.95
avg_epoch_timeMethod · 0.95
median_epoch_timeMethod · 0.95
closeMethod · 0.95
MultigridScheduleClass · 0.90
build_modelFunction · 0.90
AVAMeterClass · 0.90
TrainMeterClass · 0.90

Tested by

no test coverage detected