(self, args)
| 1558 | return None |
| 1559 | |
| 1560 | def _set_distributed_vars(self, args): |
| 1561 | device_rank = args.device_rank if args is not None and hasattr(args, 'device_rank') else self.local_rank |
| 1562 | if device_rank >= 0: |
| 1563 | get_accelerator().set_device(device_rank) |
| 1564 | self.device = torch.device(get_accelerator().device_name(device_rank)) |
| 1565 | self.world_size = dist.get_world_size() |
| 1566 | self.global_rank = dist.get_rank() |
| 1567 | else: |
| 1568 | self.world_size = 1 |
| 1569 | self.global_rank = 0 |
| 1570 | self.device = get_accelerator().device() |
| 1571 | |
| 1572 | # Configure based on command line arguments |
| 1573 | def _configure_with_arguments(self, args, mpu): |
no test coverage detected