(self, x_enc, x_mark_enc)
| 286 | return (out1_list, out2_list) |
| 287 | |
| 288 | def __multi_scale_process_inputs(self, x_enc, x_mark_enc): |
| 289 | if self.configs.down_sampling_method == 'max': |
| 290 | down_pool = torch.nn.MaxPool1d(self.configs.down_sampling_window, return_indices=False) |
| 291 | elif self.configs.down_sampling_method == 'avg': |
| 292 | down_pool = torch.nn.AvgPool1d(self.configs.down_sampling_window) |
| 293 | elif self.configs.down_sampling_method == 'conv': |
| 294 | padding = 1 if torch.__version__ >= '1.5.0' else 2 |
| 295 | down_pool = nn.Conv1d(in_channels=self.configs.enc_in, out_channels=self.configs.enc_in, |
| 296 | kernel_size=3, padding=padding, |
| 297 | stride=self.configs.down_sampling_window, |
| 298 | padding_mode='circular', |
| 299 | bias=False) |
| 300 | else: |
| 301 | return x_enc, x_mark_enc |
| 302 | # B,T,C -> B,C,T |
| 303 | x_enc = x_enc.permute(0, 2, 1) |
| 304 | |
| 305 | x_enc_ori = x_enc |
| 306 | x_mark_enc_mark_ori = x_mark_enc |
| 307 | |
| 308 | x_enc_sampling_list = [] |
| 309 | x_mark_sampling_list = [] |
| 310 | x_enc_sampling_list.append(x_enc.permute(0, 2, 1)) |
| 311 | x_mark_sampling_list.append(x_mark_enc) |
| 312 | |
| 313 | for i in range(self.configs.down_sampling_layers): |
| 314 | x_enc_sampling = down_pool(x_enc_ori) |
| 315 | |
| 316 | x_enc_sampling_list.append(x_enc_sampling.permute(0, 2, 1)) |
| 317 | x_enc_ori = x_enc_sampling |
| 318 | |
| 319 | if x_mark_enc_mark_ori is not None: |
| 320 | x_mark_sampling_list.append(x_mark_enc_mark_ori[:, ::self.configs.down_sampling_window, :]) |
| 321 | x_mark_enc_mark_ori = x_mark_enc_mark_ori[:, ::self.configs.down_sampling_window, :] |
| 322 | |
| 323 | x_enc = x_enc_sampling_list |
| 324 | if x_mark_enc_mark_ori is not None: |
| 325 | x_mark_enc = x_mark_sampling_list |
| 326 | else: |
| 327 | x_mark_enc = x_mark_enc |
| 328 | |
| 329 | return x_enc, x_mark_enc |
| 330 | |
| 331 | def forecast(self, x_enc, x_mark_enc, x_dec, x_mark_dec): |
| 332 |
no outgoing calls
no test coverage detected