| 17 | return None |
| 18 | |
| 19 | def _acquire_device(self): |
| 20 | if self.args.use_gpu: |
| 21 | import platform |
| 22 | if platform.system() == 'Darwin': |
| 23 | device = torch.device('mps') |
| 24 | print('Use MPS') |
| 25 | return device |
| 26 | os.environ["CUDA_VISIBLE_DEVICES"] = str( |
| 27 | self.args.gpu) if not self.args.use_multi_gpu else self.args.devices |
| 28 | device = torch.device('cuda:{}'.format(self.args.gpu)) |
| 29 | if self.args.use_multi_gpu: |
| 30 | print('Use GPU: cuda{}'.format(self.args.device_ids)) |
| 31 | else: |
| 32 | print('Use GPU: cuda:{}'.format(self.args.gpu)) |
| 33 | else: |
| 34 | device = torch.device('cpu') |
| 35 | print('Use CPU') |
| 36 | return device |
| 37 | |
| 38 | def _get_data(self): |
| 39 | pass |