()
| 17 | #---------------------------------------------------------------------------- |
| 18 | |
| 19 | def 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 |
nothing calls this directly
no test coverage detected