(self, x_enc, x_mark_enc, x_dec, x_mark_dec)
| 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 = [] |
no test coverage detected