Initialize distributed training environment.
(args)
| 25 | |
| 26 | |
| 27 | def distributed_init(args): |
| 28 | """Initialize distributed training environment.""" |
| 29 | world_size = int(os.getenv('WORLD_SIZE', '1')) |
| 30 | rank = int(os.getenv('RANK', '0')) |
| 31 | local_rank = args.local_rank |
| 32 | |
| 33 | # Set environment variables for distributed training |
| 34 | os.environ['MASTER_ADDR'] = os.getenv('MASTER_ADDR', 'localhost') |
| 35 | os.environ['MASTER_PORT'] = str(args.master_port) |
| 36 | |
| 37 | return world_size, rank, local_rank |
| 38 | |
| 39 | |
| 40 | if __name__ == '__main__': |
no outgoing calls
no test coverage detected