MCPcopy Create free account
hub / github.com/Anoise/WTFlib / __init__

Method __init__

LDPS_Graph/models/FEDformer.py:22–137  ·  view source on GitHub ↗
(self, configs)

Source from the content-addressed store, hash-verified

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,

Callers

nothing calls this directly

Calls 15

series_decomp_multiClass · 0.90
series_decompClass · 0.90
DataEmbeddingClass · 0.90
MultiWaveletCrossClass · 0.90
FourierBlockClass · 0.90
EncoderClass · 0.90
EncoderLayerClass · 0.90

Tested by

no test coverage detected