(self, x, scale=None)
| 518 | return x_recon, mu, log_var |
| 519 | |
| 520 | def encode(self, x, scale=None): |
| 521 | self.clear_cache() |
| 522 | ## cache |
| 523 | t = x.shape[2] |
| 524 | iter_ = 1 + (t - 1) // 4 |
| 525 | if scale != None: |
| 526 | scale = [item.to(x.device, x.dtype) for item in scale] |
| 527 | ## 对encode输入的x,按时间拆分为1、4、4、4.... |
| 528 | for i in range(iter_): |
| 529 | self._enc_conv_idx = [0] |
| 530 | if i == 0: |
| 531 | out = self.encoder( |
| 532 | x[:, :, :1, :, :], |
| 533 | feat_cache=self._enc_feat_map, |
| 534 | feat_idx=self._enc_conv_idx) |
| 535 | else: |
| 536 | out_ = self.encoder( |
| 537 | x[:, :, 1 + 4 * (i - 1):1 + 4 * i, :, :], |
| 538 | feat_cache=self._enc_feat_map, |
| 539 | feat_idx=self._enc_conv_idx) |
| 540 | out = torch.cat([out, out_], 2) |
| 541 | mu, log_var = self.conv1(out).chunk(2, dim=1) |
| 542 | if scale != None: |
| 543 | if isinstance(scale[0], torch.Tensor): |
| 544 | mu = (mu - scale[0].view(1, self.z_dim, 1, 1, 1)) * scale[1].view( |
| 545 | 1, self.z_dim, 1, 1, 1) |
| 546 | else: |
| 547 | mu = (mu - scale[0]) * scale[1] |
| 548 | x = torch.cat([mu, log_var], dim = 1) |
| 549 | self.clear_cache() |
| 550 | return x |
| 551 | |
| 552 | def decode(self, z, scale=None): |
| 553 | self.clear_cache() |
no test coverage detected