| 550 | return x |
| 551 | |
| 552 | def decode(self, z, scale=None): |
| 553 | self.clear_cache() |
| 554 | # z: [b,c,t,h,w] |
| 555 | if scale != None: |
| 556 | scale = [item.to(z.device, z.dtype) for item in scale] |
| 557 | if isinstance(scale[0], torch.Tensor): |
| 558 | z = z / scale[1].view(1, self.z_dim, 1, 1, 1) + scale[0].view( |
| 559 | 1, self.z_dim, 1, 1, 1) |
| 560 | else: |
| 561 | z = z / scale[1] + scale[0] |
| 562 | iter_ = z.shape[2] |
| 563 | x = self.conv2(z) |
| 564 | for i in range(iter_): |
| 565 | self._conv_idx = [0] |
| 566 | if i == 0: |
| 567 | out = self.decoder( |
| 568 | x[:, :, i:i + 1, :, :], |
| 569 | feat_cache=self._feat_map, |
| 570 | feat_idx=self._conv_idx) |
| 571 | else: |
| 572 | out_ = self.decoder( |
| 573 | x[:, :, i:i + 1, :, :], |
| 574 | feat_cache=self._feat_map, |
| 575 | feat_idx=self._conv_idx) |
| 576 | out = torch.cat([out, out_], 2) |
| 577 | self.clear_cache() |
| 578 | return out |
| 579 | |
| 580 | def reparameterize(self, mu, log_var): |
| 581 | std = torch.exp(0.5 * log_var) |