From 1952e9be491545fc59227001f1c0bc865672bba0 Mon Sep 17 00:00:00 2001 From: William Falcon Date: Fri, 12 Jul 2019 15:11:32 -0400 Subject: [PATCH] fixed nccl init --- pytorch_lightning/models/trainer.py | 30 ++++++++++++++++------------- 1 file changed, 17 insertions(+), 13 deletions(-) diff --git a/pytorch_lightning/models/trainer.py b/pytorch_lightning/models/trainer.py index fef2de82..3f7735bb 100644 --- a/pytorch_lightning/models/trainer.py +++ b/pytorch_lightning/models/trainer.py @@ -375,19 +375,23 @@ class Trainer(TrainerIO): :param tries: :return: """ - if tries > 20: - raise RuntimeError('Failed to connect using 20 different ip addresses') - - try: - root_node = os.environ['SLURM_NODELIST'].split(' ')[0] - os.environ['MASTER_ADDR'] = root_node - os.environ['MASTER_PORT'] = f'{port}' - dist.init_process_group("nccl", rank=self.proc_rank, world_size=self.world_size) - - except RuntimeError as e: - # port taken - warnings.warn(f'port {port} taken, trying port {port}...') - self.__init_tcp_connection(port + 1, tries + 1) + root_node = os.environ['SLURM_NODELIST'].split(' ')[0] + os.environ['MASTER_ADDR'] = root_node + os.environ['MASTER_PORT'] = f'{port}' + dist.init_process_group("nccl", rank=self.proc_rank, world_size=self.world_size) + # if tries > 20: + # raise RuntimeError('Failed to connect using 20 different ip addresses') + # + # try: + # root_node = os.environ['SLURM_NODELIST'].split(' ')[0] + # os.environ['MASTER_ADDR'] = root_node + # os.environ['MASTER_PORT'] = f'{port}' + # dist.init_process_group("nccl", rank=self.proc_rank, world_size=self.world_size) + # + # except RuntimeError as e: + # # port taken + # warnings.warn(f'port {port} taken, trying port {port}...') + # self.__init_tcp_connection(port + 1, tries + 1) def __run_pretrain_routine(self, model): """