MCPcopy Create free account
hub / github.com/kwuking/TimeMixer / __init__

Method __init__

models/TimeMixer.py:188–266  ·  view source on GitHub ↗
(self, configs)

Source from the content-addressed store, hash-verified

186class Model(nn.Module):
187
188 def __init__(self, configs):
189 super(Model, self).__init__()
190 self.configs = configs
191 self.task_name = configs.task_name
192 self.seq_len = configs.seq_len
193 self.label_len = configs.label_len
194 self.pred_len = configs.pred_len
195 self.down_sampling_window = configs.down_sampling_window
196 self.channel_independence = configs.channel_independence
197 self.pdm_blocks = nn.ModuleList([PastDecomposableMixing(configs)
198 for _ in range(configs.e_layers)])
199
200 self.preprocess = series_decomp(configs.moving_avg)
201 self.enc_in = configs.enc_in
202 self.use_future_temporal_feature = configs.use_future_temporal_feature
203
204 if self.channel_independence == 1:
205 self.enc_embedding = DataEmbedding_wo_pos(1, configs.d_model, configs.embed, configs.freq,
206 configs.dropout)
207 else:
208 self.enc_embedding = DataEmbedding_wo_pos(configs.enc_in, configs.d_model, configs.embed, configs.freq,
209 configs.dropout)
210
211 self.layer = configs.e_layers
212
213 self.normalize_layers = torch.nn.ModuleList(
214 [
215 Normalize(self.configs.enc_in, affine=True, non_norm=True if configs.use_norm == 0 else False)
216 for i in range(configs.down_sampling_layers + 1)
217 ]
218 )
219
220 if self.task_name == 'long_term_forecast' or self.task_name == 'short_term_forecast':
221 self.predict_layers = torch.nn.ModuleList(
222 [
223 torch.nn.Linear(
224 configs.seq_len // (configs.down_sampling_window ** i),
225 configs.pred_len,
226 )
227 for i in range(configs.down_sampling_layers + 1)
228 ]
229 )
230
231 if self.channel_independence == 1:
232 self.projection_layer = nn.Linear(
233 configs.d_model, 1, bias=True)
234 else:
235 self.projection_layer = nn.Linear(
236 configs.d_model, configs.c_out, bias=True)
237
238 self.out_res_layers = torch.nn.ModuleList([
239 torch.nn.Linear(
240 configs.seq_len // (configs.down_sampling_window ** i),
241 configs.seq_len // (configs.down_sampling_window ** i),
242 )
243 for i in range(configs.down_sampling_layers + 1)
244 ])
245

Callers 4

__init__Method · 0.45
__init__Method · 0.45
__init__Method · 0.45
__init__Method · 0.45

Calls 4

series_decompClass · 0.90
NormalizeClass · 0.90

Tested by

no test coverage detected