(self, configs)
| 38 | FEDformer performs the attention mechanism on frequency domain and achieved O(N) complexity |
| 39 | """ |
| 40 | def __init__(self, configs): |
| 41 | super(Model, self).__init__() |
| 42 | self.version = configs.version |
| 43 | self.mode_select = configs.mode_select |
| 44 | self.modes = configs.modes |
| 45 | self.seq_len = configs.seq_len |
| 46 | self.label_len = configs.label_len |
| 47 | self.pred_len = configs.pred_len |
| 48 | self.output_attention = configs.output_attention |
| 49 | |
| 50 | # Decomp |
| 51 | kernel_size = configs.moving_avg |
| 52 | if isinstance(kernel_size, list): |
| 53 | self.decomp = series_decomp_multi(kernel_size) |
| 54 | else: |
| 55 | self.decomp = series_decomp(kernel_size) |
| 56 | |
| 57 | # Embedding |
| 58 | # The series-wise connection inherently contains the sequential information. |
| 59 | # Thus, we can discard the position embedding of transformers. |
| 60 | self.enc_embedding = DataEmbedding_mine(configs.enc_in, configs.d_model, configs.embed, configs.freq, configs.dropout) |
| 61 | self.dec_embedding = DataEmbedding_mine(configs.dec_in, configs.d_model, configs.embed, configs.freq, configs.dropout, is_decoder=True) |
| 62 | |
| 63 | if configs.version == 'Wavelets': |
| 64 | encoder_self_att = MultiWaveletTransform(ich=configs.d_model, L=configs.L, base=configs.base) |
| 65 | decoder_self_att = MultiWaveletTransform(ich=configs.d_model, L=configs.L, base=configs.base) |
| 66 | decoder_cross_att = MultiWaveletCross(in_channels=configs.d_model, |
| 67 | out_channels=configs.d_model, |
| 68 | seq_len_q=self.seq_len // 2 + self.pred_len, |
| 69 | seq_len_kv=self.seq_len, |
| 70 | modes=configs.modes, |
| 71 | ich=configs.d_model, |
| 72 | base=configs.base, |
| 73 | activation=configs.cross_activation) |
| 74 | else: |
| 75 | encoder_self_att = FourierBlock(in_channels=configs.d_model, |
| 76 | out_channels=configs.d_model, |
| 77 | seq_len=self.seq_len, |
| 78 | modes=configs.modes, |
| 79 | mode_select_method=configs.mode_select) |
| 80 | decoder_self_att = FourierBlock(in_channels=configs.d_model, |
| 81 | out_channels=configs.d_model, |
| 82 | seq_len=self.seq_len//2+self.pred_len, |
| 83 | modes=configs.modes, |
| 84 | mode_select_method=configs.mode_select) |
| 85 | decoder_cross_att = FourierCrossAttention(in_channels=configs.d_model, |
| 86 | out_channels=configs.d_model, |
| 87 | seq_len_q=self.seq_len//2+self.pred_len, |
| 88 | seq_len_kv=self.seq_len, |
| 89 | modes=configs.modes, |
| 90 | mode_select_method=configs.mode_select) |
| 91 | # Encoder |
| 92 | enc_modes = int(min(configs.modes, configs.seq_len//2)) |
| 93 | dec_modes = int(min(configs.modes, (configs.seq_len//2+configs.pred_len)//2)) |
| 94 | print('enc_modes: {}, dec_modes: {}'.format(enc_modes, dec_modes)) |
| 95 | |
| 96 | self.encoder = Encoder( |
| 97 | [ |
no test coverage detected