MCPcopy Create free account
hub / github.com/BorealisAI/scaleformer / __init__

Method __init__

models/FEDformerMS.py:40–142  ·  view source on GitHub ↗
(self, configs)

Source from the content-addressed store, hash-verified

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 [

Callers 1

__init__Method · 0.45

Calls 14

series_decomp_multiClass · 0.90
series_decompClass · 0.90
DataEmbedding_mineClass · 0.90
MultiWaveletCrossClass · 0.90
FourierBlockClass · 0.90
EncoderClass · 0.90
EncoderLayerClass · 0.90
my_LayernormClass · 0.90
DecoderClass · 0.90

Tested by

no test coverage detected