(self, configs)
| 20 | FEDformer performs the attention mechanism on frequency domain and achieved O(N) complexity |
| 21 | """ |
| 22 | def __init__(self, configs): |
| 23 | super(Model, self).__init__() |
| 24 | self.version = configs.version |
| 25 | self.mode_select = configs.mode_select |
| 26 | self.modes = configs.modes |
| 27 | self.seq_len = configs.seq_len |
| 28 | self.label_len = configs.label_len |
| 29 | self.pred_len = configs.pred_len |
| 30 | self.output_attention = configs.output_attention |
| 31 | |
| 32 | # Decomp |
| 33 | kernel_size = configs.moving_avg |
| 34 | if isinstance(kernel_size, list): |
| 35 | self.decomp = series_decomp_multi(kernel_size) |
| 36 | else: |
| 37 | self.decomp = series_decomp(kernel_size) |
| 38 | |
| 39 | # Embedding |
| 40 | # The series-wise connection inherently contains the sequential information. |
| 41 | # Thus, we can discard the position embedding of transformers. |
| 42 | # self.enc_embedding = DataEmbedding_wo_pos(configs.enc_in, configs.d_model, configs.embed, configs.freq, |
| 43 | # configs.dropout) |
| 44 | # self.dec_embedding = DataEmbedding_wo_pos(configs.dec_in, configs.d_model, configs.embed, configs.freq, |
| 45 | # configs.dropout) |
| 46 | if configs.embed_type == 0: |
| 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 | elif configs.embed_type == 1: |
| 52 | self.enc_embedding = DataEmbedding(configs.enc_in, configs.d_model, configs.embed, configs.freq, |
| 53 | configs.dropout) |
| 54 | self.dec_embedding = DataEmbedding(configs.dec_in, configs.d_model, configs.embed, configs.freq, |
| 55 | configs.dropout) |
| 56 | elif configs.embed_type == 2: |
| 57 | self.enc_embedding = DataEmbedding_wo_pos_temp(configs.enc_in, configs.d_model, configs.embed, configs.freq, |
| 58 | configs.dropout) |
| 59 | self.dec_embedding = DataEmbedding_wo_pos_temp(configs.dec_in, configs.d_model, configs.embed, configs.freq, |
| 60 | configs.dropout) |
| 61 | elif configs.embed_type == 3: |
| 62 | self.enc_embedding = DataEmbedding_wo_temp(configs.enc_in, configs.d_model, configs.embed, configs.freq, |
| 63 | configs.dropout) |
| 64 | self.dec_embedding = DataEmbedding_wo_temp(configs.dec_in, configs.d_model, configs.embed, configs.freq, |
| 65 | configs.dropout) |
| 66 | |
| 67 | if configs.version == 'Wavelets': |
| 68 | encoder_self_att = MultiWaveletTransform(ich=configs.d_model, L=configs.L, base=configs.base) |
| 69 | decoder_self_att = MultiWaveletTransform(ich=configs.d_model, L=configs.L, base=configs.base) |
| 70 | decoder_cross_att = MultiWaveletCross(in_channels=configs.d_model, |
| 71 | out_channels=configs.d_model, |
| 72 | seq_len_q=self.seq_len // 2 + self.pred_len, |
| 73 | seq_len_kv=self.seq_len, |
| 74 | modes=configs.modes, |
| 75 | ich=configs.d_model, |
| 76 | base=configs.base, |
| 77 | activation=configs.cross_activation) |
| 78 | else: |
| 79 | encoder_self_att = FourierBlock(in_channels=configs.d_model, |
nothing calls this directly
no test coverage detected