bug: fixed port

This commit is contained in:
Yann Ahlgrim
2026-06-09 01:07:32 +02:00
parent ca026ffecb
commit ccdddc47e5
+16 -8
View File
@@ -6,6 +6,7 @@
# #
import os import os
import socket
import torch import torch
import torch.distributed as dist import torch.distributed as dist
@@ -15,33 +16,40 @@ from logging import getLogger
logger = getLogger() logger = getLogger()
def init_distributed(port=40112, rank_and_world_size=(None, None)): def _find_free_port():
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
s.bind(('', 0))
return s.getsockname()[1]
def init_distributed(port=None, rank_and_world_size=(None, None)):
if dist.is_available() and dist.is_initialized(): if dist.is_available() and dist.is_initialized():
return dist.get_world_size(), dist.get_rank() return dist.get_world_size(), dist.get_rank()
rank, world_size = rank_and_world_size rank, world_size = rank_and_world_size
os.environ['MASTER_ADDR'] = 'localhost'
if (rank is None) or (world_size is None): if (rank is None) or (world_size is None):
try: try:
world_size = int(os.environ['SLURM_NTASKS']) world_size = int(os.environ['SLURM_NTASKS'])
rank = int(os.environ['SLURM_PROCID']) rank = int(os.environ['SLURM_PROCID'])
os.environ['MASTER_ADDR'] = os.environ['HOSTNAME']
except Exception: except Exception:
logger.info('SLURM vars not set (distributed training not available)') logger.info('SLURM vars not set (distributed training not available)')
world_size, rank = 1, 0 return 1, 0
return world_size, rank
os.environ['MASTER_ADDR'] = '127.0.0.1'
os.environ['MASTER_PORT'] = str(port if port is not None else _find_free_port())
os.environ['NCCL_SOCKET_IFNAME'] = 'lo'
os.environ['NCCL_NET_GIB_EXTRA_IFS'] = 'lo'
try: try:
os.environ['MASTER_PORT'] = str(port)
torch.distributed.init_process_group( torch.distributed.init_process_group(
backend='nccl', backend='nccl',
world_size=world_size, world_size=world_size,
rank=rank) rank=rank)
except Exception as e: except Exception as e:
world_size, rank = 1, 0 logger.warning(f'NCCL init failed ({e}); falling back to single-process')
logger.info(f'distributed training not available {e}') return 1, 0
return world_size, rank return world_size, rank