bug: fixed port
This commit is contained in:
@@ -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
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user