MCPcopy Create free account
hub / github.com/dcharatan/flowmap / pretrain

Function pretrain

flowmap/pretrain.py:28–75  ·  view source on GitHub ↗
(cfg_dict: DictConfig)

Source from the content-addressed store, hash-verified

26 config_name="pretrain",
27)
28def 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
78if __name__ == "__main__":

Callers 1

pretrain.pyFile · 0.85

Calls 7

get_typed_root_configFunction · 0.85
ModelClass · 0.85
get_lossesFunction · 0.85
get_visualizersFunction · 0.85
DataModulePretrainClass · 0.85

Tested by

no test coverage detected