diff --git a/zamba/models/model_manager.py b/zamba/models/model_manager.py index a80c8014..fb32ac63 100644 --- a/zamba/models/model_manager.py +++ b/zamba/models/model_manager.py @@ -280,19 +280,17 @@ def train_model( accelerator, devices = configure_accelerator_and_devices_from_gpus(train_config.gpus) + + multiprocessing_strategy = getattr(train_config,"multiprocessing_strategy",None) + trainer = pl.Trainer( accelerator=accelerator, devices=devices, max_epochs=train_config.max_epochs, logger=tensorboard_logger, - callbacks=callbacks, - fast_dev_run=train_config.dry_run, - strategy=( - DDPStrategy(find_unused_parameters=False) - if (data_module.multiprocessing_context is not None) and (train_config.gpus > 1) - else "auto" - ), + strategy = multiprocessing_strategy, ) + #Set the strategy within trainer to reflect changes if video_loader_config.cache_dir is None: logger.info("No cache dir is specified. Videos will not be cached.") diff --git a/zamba/pytorch_lightning/utils.py b/zamba/pytorch_lightning/utils.py index 041f63e9..b7bff343 100644 --- a/zamba/pytorch_lightning/utils.py +++ b/zamba/pytorch_lightning/utils.py @@ -65,11 +65,8 @@ def __init__( transform=transform, video_loader_config=video_loader_config, ) - self.multiprocessing_context: BaseContext = ( - None - if (multiprocessing_context is None) or (num_workers == 0) - else multiprocessing_context - ) + self.multiprocessing_context: BaseContext = multiprocessing_context #Modified the multiprocessing context to not factor in num_workers decoupliing it + super().__init__(*args, **kwargs)