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

Method forecast

models/TimeMixer.py:331–383  ·  view source on GitHub ↗
(self, x_enc, x_mark_enc, x_dec, x_mark_dec)

Source from the content-addressed store, hash-verified

329 return x_enc, x_mark_enc
330
331 def forecast(self, x_enc, x_mark_enc, x_dec, x_mark_dec):
332
333 if self.use_future_temporal_feature:
334 if self.channel_independence == 1:
335 B, T, N = x_enc.size()
336 x_mark_dec = x_mark_dec.repeat(N, 1, 1)
337 self.x_mark_dec = self.enc_embedding(None, x_mark_dec)
338 else:
339 self.x_mark_dec = self.enc_embedding(None, x_mark_dec)
340
341 x_enc, x_mark_enc = self.__multi_scale_process_inputs(x_enc, x_mark_enc)
342
343 x_list = []
344 x_mark_list = []
345 if x_mark_enc is not None:
346 for i, x, x_mark in zip(range(len(x_enc)), x_enc, x_mark_enc):
347 B, T, N = x.size()
348 x = self.normalize_layers[i](x, 'norm')
349 if self.channel_independence == 1:
350 x = x.permute(0, 2, 1).contiguous().reshape(B * N, T, 1)
351 x_mark = x_mark.repeat(N, 1, 1)
352 x_list.append(x)
353 x_mark_list.append(x_mark)
354 else:
355 for i, x in zip(range(len(x_enc)), x_enc, ):
356 B, T, N = x.size()
357 x = self.normalize_layers[i](x, 'norm')
358 if self.channel_independence == 1:
359 x = x.permute(0, 2, 1).contiguous().reshape(B * N, T, 1)
360 x_list.append(x)
361
362 # embedding
363 enc_out_list = []
364 x_list = self.pre_enc(x_list)
365 if x_mark_enc is not None:
366 for i, x, x_mark in zip(range(len(x_list[0])), x_list[0], x_mark_list):
367 enc_out = self.enc_embedding(x, x_mark) # [B,T,C]
368 enc_out_list.append(enc_out)
369 else:
370 for i, x in zip(range(len(x_list[0])), x_list[0]):
371 enc_out = self.enc_embedding(x, None) # [B,T,C]
372 enc_out_list.append(enc_out)
373
374 # Past Decomposable Mixing as encoder for past
375 for i in range(self.layer):
376 enc_out_list = self.pdm_blocks[i](enc_out_list)
377
378 # Future Multipredictor Mixing as decoder for future
379 dec_out_list = self.future_multi_mixing(B, enc_out_list, x_list)
380
381 dec_out = torch.stack(dec_out_list, dim=-1).sum(-1)
382 dec_out = self.normalize_layers[0](dec_out, 'denorm')
383 return dec_out
384
385 def future_multi_mixing(self, B, enc_out_list, x_list):
386 dec_out_list = []

Callers 1

forwardMethod · 0.95

Calls 3

pre_encMethod · 0.95
future_multi_mixingMethod · 0.95

Tested by

no test coverage detected