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