MCPcopy Create free account
hub / github.com/kwuking/TimeMixer / Exp_Basic

Class Exp_Basic

exp/exp_basic.py:6–48  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

4
5
6class 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

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected