fixed nccl init

This commit is contained in:
William Falcon
2019-07-12 15:11:32 -04:00
parent 19391b1df1
commit 1952e9be49
+17 -13
View File
@@ -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):
"""