| 4 | |
| 5 | |
| 6 | class Exp_Basic(object): |
| 7 | def __init__(self, args): |
| 8 | self.args = args |
| 9 | self.model_dict = { |
| 10 | 'TimeMixer': TimeMixer, |
| 11 | } |
| 12 | self.device = self._acquire_device() |
| 13 | self.model = self._build_model().to(self.device) |
| 14 | |
| 15 | def _build_model(self): |
| 16 | raise NotImplementedError |
| 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 |
| 40 | |
| 41 | def vali(self): |
| 42 | pass |
| 43 | |
| 44 | def train(self): |
| 45 | pass |
| 46 | |
| 47 | def test(self): |
| 48 | pass |
nothing calls this directly
no outgoing calls
no test coverage detected