forked from ll121202/HyperD
-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathtrain.py
More file actions
28 lines (20 loc) · 820 Bytes
/
Copy pathtrain.py
File metadata and controls
28 lines (20 loc) · 820 Bytes
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
# Run a baseline model in BasicTS framework.
# pylint: disable=wrong-import-position
import os
import sys
from argparse import ArgumentParser
# sys.path.append(os.path.abspath(__file__ + '/../..'))
# os.chdir(os.path.abspath(os.path.join(os.path.dirname(__file__), '..')))
import torch
import basicts
torch.set_num_threads(4) # aviod high cpu avg usage
def parse_args():
parser = ArgumentParser(description='Run time series forecasting model in BasicTS framework!')
parser.add_argument('-c', '--cfg', default='baselines/HyperD/PEMS04.py', help='training config')
parser.add_argument('-g', '--gpus', default='0', help='visible gpus')
return parser.parse_args()
def main():
args = parse_args()
basicts.launch_training(args.cfg, args.gpus, node_rank=0)
if __name__ == '__main__':
main()