Build training model and its associated tools, including optimizer, dataloaders and meters. Args: cfg (CfgNode): configs. Details can be found in slowfast/config/defaults.py Returns: model (nn.Module): training model. optimizer (Optimizer): optimi
(cfg)
| 335 | |
| 336 | |
| 337 | def build_trainer(cfg): |
| 338 | """ |
| 339 | Build training model and its associated tools, including optimizer, |
| 340 | dataloaders and meters. |
| 341 | Args: |
| 342 | cfg (CfgNode): configs. Details can be found in |
| 343 | slowfast/config/defaults.py |
| 344 | Returns: |
| 345 | model (nn.Module): training model. |
| 346 | optimizer (Optimizer): optimizer. |
| 347 | train_loader (DataLoader): training data loader. |
| 348 | val_loader (DataLoader): validatoin data loader. |
| 349 | precise_bn_loader (DataLoader): training data loader for computing |
| 350 | precise BN. |
| 351 | train_meter (TrainMeter): tool for measuring training stats. |
| 352 | val_meter (ValMeter): tool for measuring validation stats. |
| 353 | """ |
| 354 | # Build the video model and print model statistics. |
| 355 | model = build_model(cfg) |
| 356 | if du.is_master_proc() and cfg.LOG_MODEL_INFO: |
| 357 | misc.log_model_info(model, cfg, use_train_input=True) |
| 358 | |
| 359 | # Construct the optimizer. |
| 360 | optimizer = optim.construct_optimizer(model, cfg) |
| 361 | |
| 362 | # Loss scaler |
| 363 | loss_scaler = NativeScaler() |
| 364 | |
| 365 | # Create the video train and val loaders. |
| 366 | train_loader = loader.construct_loader(cfg, "train") |
| 367 | val_loader = loader.construct_loader(cfg, "val") |
| 368 | precise_bn_loader = loader.construct_loader( |
| 369 | cfg, "train", is_precise_bn=True |
| 370 | ) |
| 371 | # Create meters. |
| 372 | train_meter = TrainMeter(len(train_loader), cfg) |
| 373 | val_meter = ValMeter(len(val_loader), cfg) |
| 374 | |
| 375 | return ( |
| 376 | model, |
| 377 | optimizer, |
| 378 | loss_scaler, |
| 379 | train_loader, |
| 380 | val_loader, |
| 381 | precise_bn_loader, |
| 382 | train_meter, |
| 383 | val_meter, |
| 384 | ) |
| 385 | |
| 386 | |
| 387 | def train(cfg): |
no test coverage detected