diff --git a/examples/new_project_templates/trainer_gpu_cluster_template.py b/examples/new_project_templates/trainer_gpu_cluster_template.py index f852a74d..6a75622d 100644 --- a/examples/new_project_templates/trainer_gpu_cluster_template.py +++ b/examples/new_project_templates/trainer_gpu_cluster_template.py @@ -101,6 +101,7 @@ def main(hparams, cluster, results_dict): checkpoint_callback=checkpoint, early_stop_callback=early_stop, gpus=gpu_list, + nb_gpu_nodes=1 ) # train model diff --git a/pytorch_lightning/models/trainer.py b/pytorch_lightning/models/trainer.py index dab2815a..dcf100b0 100644 --- a/pytorch_lightning/models/trainer.py +++ b/pytorch_lightning/models/trainer.py @@ -41,6 +41,7 @@ class Trainer(TrainerIO): cluster=None, process_position=0, current_gpu_name=0, + nb_gpu_nodes=None, gpus=None, progress_bar=True, overfit_pct=0.0, @@ -58,6 +59,7 @@ class Trainer(TrainerIO): nb_sanity_val_steps=5): # Transfer params + self.nb_gpu_nodes = nb_gpu_nodes self.gradient_clip = gradient_clip self.check_val_every_n_epoch = check_val_every_n_epoch self.enable_early_stop = enable_early_stop @@ -304,7 +306,7 @@ class Trainer(TrainerIO): print('configuring server') rank = proc_rank * len(self.data_parallel_device_ids) + gpu_nb print(f"GPU: {gpu_nb} - Rank: {rank}") - world_size = self.cluster.per_experiment_nb_nodes * self.cluster.per_experiment_nb_gpus + world_size = self.nb_gpu_nodes * len(self.data_parallel_device_ids) dist.init_process_group("nccl", init_method=f'tcp://{ip}:12001', rank=rank, world_size=world_size) # copy model to each gpu