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)
| 385 | |
| 386 | |
| 387 | def 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 |
nothing calls this directly
no test coverage detected