| 16 | import numpy as np |
| 17 | |
| 18 | def main(): |
| 19 | fix_seed = 2021 |
| 20 | random.seed(fix_seed) |
| 21 | torch.manual_seed(fix_seed) |
| 22 | np.random.seed(fix_seed) |
| 23 | |
| 24 | parser = argparse.ArgumentParser(description='Autoformer & Transformer family for Time Series Forecasting') |
| 25 | |
| 26 | # basic config |
| 27 | parser.add_argument('--is_training', type=int, default=1, help='status') |
| 28 | parser.add_argument('--use_multi_scale', action='store_true', help='using mult-scale') |
| 29 | parser.add_argument('--prob_forecasting', action='store_true', help='using probabilistic forecasting') |
| 30 | parser.add_argument('--scales', default=[16, 8, 4, 2, 1], help='scales in mult-scale') |
| 31 | parser.add_argument('--scale_factor', type=int, default=2, help='scale factor for upsample') |
| 32 | parser.add_argument('--model', type=str, required=True, default='Autoformer', |
| 33 | help='model name, options: [Autoformer, Informer, Transformer, Reformer, FEDformer] and their MS versions: [AutoformerMS, InformerMS, etc]') |
| 34 | |
| 35 | # data loader |
| 36 | parser.add_argument('--data', type=str, default='custom', help='dataset type') |
| 37 | parser.add_argument('--root_path', type=str, default='./data/ETT/', help='root path of the data file') |
| 38 | parser.add_argument('--data_path', type=str, default='ETTh1.csv', help='data file') |
| 39 | parser.add_argument('--features', type=str, default='M', |
| 40 | help='forecasting task, options:[M, S, MS]; M:multivariate predict multivariate, S:univariate predict univariate, MS:multivariate predict univariate') |
| 41 | parser.add_argument('--target', type=str, default='OT', help='target feature in S or MS task') |
| 42 | parser.add_argument('--freq', type=str, default='h', |
| 43 | help='freq for time features encoding, options:[s:secondly, t:minutely, h:hourly, d:daily, b:business days, w:weekly, m:monthly], you can also use more detailed freq like 15min or 3h') |
| 44 | parser.add_argument('--checkpoints', type=str, default='./checkpoints/', help='location of model checkpoints') |
| 45 | |
| 46 | # forecasting task |
| 47 | parser.add_argument('--seq_len', type=int, default=96, help='input sequence length') |
| 48 | parser.add_argument('--label_len', type=int, default=48, help='start token length') |
| 49 | parser.add_argument('--pred_len', type=int, default=96, help='prediction sequence length') |
| 50 | |
| 51 | # supplementary config for FiLM model |
| 52 | parser.add_argument('--modes1', type=int, default=64, help='modes to be selected random 64') |
| 53 | parser.add_argument('--mode_type',type=int,default=0) |
| 54 | |
| 55 | # supplementary config for FEDformer model |
| 56 | parser.add_argument('--version', type=str, default='Wavelets', |
| 57 | help='for FEDformer, there are two versions to choose, options: [Fourier, Wavelets]') |
| 58 | parser.add_argument('--mode_select', type=str, default='low', |
| 59 | help='for FEDformer, there are two mode selection method, options: [random, low]') |
| 60 | parser.add_argument('--modes', type=int, default=64, help='modes to be selected random 64') |
| 61 | parser.add_argument('--L', type=int, default=3, help='ignore level') |
| 62 | parser.add_argument('--base', type=str, default='legendre', help='mwt base') |
| 63 | parser.add_argument('--cross_activation', type=str, default='tanh', |
| 64 | help='mwt cross atention activation function tanh or softmax') |
| 65 | |
| 66 | # supplementary config for Reformer model |
| 67 | parser.add_argument('--bucket_size', type=int, default=4, help='for Reformer') |
| 68 | parser.add_argument('--n_hashes', type=int, default=4, help='for Reformer') |
| 69 | parser.add_argument('--film_ours', default=True, action='store_true') |
| 70 | parser.add_argument('--ab', type=int, default=2, help='ablation version') |
| 71 | parser.add_argument('--ratio', type=float, default=0.5, help='dropout') |
| 72 | parser.add_argument('--film_version', type=int, default=0, help='compression') |
| 73 | |
| 74 | # model define |
| 75 | parser.add_argument('--enc_in', type=int, default=7, help='encoder input size') |