(cfg_dict: DictConfig)
| 26 | config_name="pretrain", |
| 27 | ) |
| 28 | def pretrain(cfg_dict: DictConfig) -> None: |
| 29 | cfg = get_typed_root_config(cfg_dict, PretrainCfg) |
| 30 | callbacks, logger, checkpoint_path, _ = run_common_training_setup(cfg, cfg_dict) |
| 31 | |
| 32 | # Configure the datasets to load the model's desired image shape. |
| 33 | multiplier = cfg.cropping.flow_scale_multiplier |
| 34 | flow_shape = tuple(x * multiplier for x in cfg.cropping.image_shape) |
| 35 | for dataset_cfg in cfg.dataset: |
| 36 | dataset_cfg.image_shape = flow_shape |
| 37 | |
| 38 | # Set up the model. |
| 39 | model = Model(cfg.model) |
| 40 | losses = get_losses(cfg.loss) |
| 41 | visualizers = get_visualizers(cfg.visualizer) |
| 42 | model_wrapper = ModelWrapperPretrain( |
| 43 | cfg.model_wrapper, |
| 44 | cfg.cropping, |
| 45 | cfg.flow, |
| 46 | model, |
| 47 | losses, |
| 48 | visualizers, |
| 49 | ) |
| 50 | trainer = Trainer( |
| 51 | max_epochs=-1, |
| 52 | accelerator="gpu", |
| 53 | logger=logger, |
| 54 | devices="auto", |
| 55 | strategy=( |
| 56 | "ddp_find_unused_parameters_true" |
| 57 | if torch.cuda.device_count() > 1 |
| 58 | else "auto" |
| 59 | ), |
| 60 | callbacks=callbacks, |
| 61 | val_check_interval=cfg.trainer.val_check_interval, |
| 62 | max_steps=cfg.trainer.max_steps, |
| 63 | plugins=[SLURMEnvironment(auto_requeue=False)], |
| 64 | log_every_n_steps=1, |
| 65 | ) |
| 66 | trainer.fit( |
| 67 | model_wrapper, |
| 68 | datamodule=DataModulePretrain( |
| 69 | cfg.dataset, |
| 70 | cfg.data_module, |
| 71 | cfg.frame_sampler, |
| 72 | trainer.global_rank, |
| 73 | ), |
| 74 | ckpt_path=checkpoint_path, |
| 75 | ) |
| 76 | |
| 77 | |
| 78 | if __name__ == "__main__": |
no test coverage detected