MCPcopy Create free account
hub / github.com/SwayStar123/SpeedrunDiT / init

Function init

preprocessing/torch_utils/distributed.py:19–44  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

17#----------------------------------------------------------------------------
18
19def init():
20 global _sync_device
21
22 if not torch.distributed.is_initialized():
23 # Setup some reasonable defaults for env-based distributed init if
24 # not set by the running environment.
25 if 'MASTER_ADDR' not in os.environ:
26 os.environ['MASTER_ADDR'] = 'localhost'
27 if 'MASTER_PORT' not in os.environ:
28 s = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
29 s.bind(('', 0))
30 s.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
31 os.environ['MASTER_PORT'] = str(s.getsockname()[1])
32 s.close()
33 if 'RANK' not in os.environ:
34 os.environ['RANK'] = '0'
35 if 'LOCAL_RANK' not in os.environ:
36 os.environ['LOCAL_RANK'] = '0'
37 if 'WORLD_SIZE' not in os.environ:
38 os.environ['WORLD_SIZE'] = '1'
39 backend = 'gloo' if os.name == 'nt' else 'nccl'
40 torch.distributed.init_process_group(backend=backend, init_method='env://')
41 torch.cuda.set_device(int(os.environ.get('LOCAL_RANK', '0')))
42
43 _sync_device = torch.device('cuda') if get_world_size() > 1 else None
44 training_stats.init_multiprocessing(rank=get_rank(), sync_device=_sync_device)
45
46#----------------------------------------------------------------------------
47

Callers

nothing calls this directly

Calls 3

get_world_sizeFunction · 0.85
get_rankFunction · 0.85
closeMethod · 0.80

Tested by

no test coverage detected