(self, x_enc, x_mark_enc, mask)
| 462 | return dec_out |
| 463 | |
| 464 | def imputation(self, x_enc, x_mark_enc, mask): |
| 465 | means = torch.sum(x_enc, dim=1) / torch.sum(mask == 1, dim=1) |
| 466 | means = means.unsqueeze(1).detach() |
| 467 | x_enc = x_enc - means |
| 468 | x_enc = x_enc.masked_fill(mask == 0, 0) |
| 469 | stdev = torch.sqrt(torch.sum(x_enc * x_enc, dim=1) / |
| 470 | torch.sum(mask == 1, dim=1) + 1e-5) |
| 471 | stdev = stdev.unsqueeze(1).detach() |
| 472 | x_enc /= stdev |
| 473 | |
| 474 | B, T, N = x_enc.size() |
| 475 | x_enc, x_mark_enc = self.__multi_scale_process_inputs(x_enc, x_mark_enc) |
| 476 | |
| 477 | x_list = [] |
| 478 | x_mark_list = [] |
| 479 | if x_mark_enc is not None: |
| 480 | for i, x, x_mark in zip(range(len(x_enc)), x_enc, x_mark_enc): |
| 481 | B, T, N = x.size() |
| 482 | if self.channel_independence == 1: |
| 483 | x = x.permute(0, 2, 1).contiguous().reshape(B * N, T, 1) |
| 484 | x_list.append(x) |
| 485 | x_mark = x_mark.repeat(N, 1, 1) |
| 486 | x_mark_list.append(x_mark) |
| 487 | else: |
| 488 | for i, x in zip(range(len(x_enc)), x_enc, ): |
| 489 | B, T, N = x.size() |
| 490 | if self.channel_independence == 1: |
| 491 | x = x.permute(0, 2, 1).contiguous().reshape(B * N, T, 1) |
| 492 | x_list.append(x) |
| 493 | |
| 494 | # embedding |
| 495 | enc_out_list = [] |
| 496 | for x in x_list: |
| 497 | enc_out = self.enc_embedding(x, None) # [B,T,C] |
| 498 | enc_out_list.append(enc_out) |
| 499 | |
| 500 | # MultiScale-CrissCrossAttention as encoder for past |
| 501 | for i in range(self.layer): |
| 502 | enc_out_list = self.pdm_blocks[i](enc_out_list) |
| 503 | |
| 504 | dec_out = self.projection_layer(enc_out_list[0]) |
| 505 | dec_out = dec_out.reshape(B, self.configs.c_out, -1).permute(0, 2, 1).contiguous() |
| 506 | |
| 507 | dec_out = dec_out * \ |
| 508 | (stdev[:, 0, :].unsqueeze(1).repeat(1, self.seq_len, 1)) |
| 509 | dec_out = dec_out + \ |
| 510 | (means[:, 0, :].unsqueeze(1).repeat(1, self.seq_len, 1)) |
| 511 | return dec_out |
| 512 | |
| 513 | def forward(self, x_enc, x_mark_enc, x_dec, x_mark_dec, mask=None): |
| 514 | if self.task_name == 'long_term_forecast' or self.task_name == 'short_term_forecast': |
no test coverage detected