| 542 | return mu |
| 543 | |
| 544 | def decode(self, z, scale): |
| 545 | self.clear_cache() |
| 546 | # z: [b,c,t,h,w] |
| 547 | if isinstance(scale[0], torch.Tensor): |
| 548 | z = z / scale[1].view(1, self.z_dim, 1, 1, 1) + scale[0].view( |
| 549 | 1, self.z_dim, 1, 1, 1) |
| 550 | else: |
| 551 | z = z / scale[1] + scale[0] |
| 552 | iter_ = z.shape[2] |
| 553 | x = self.conv2(z) |
| 554 | for i in range(iter_): |
| 555 | self._conv_idx = [0] |
| 556 | if i == 0: |
| 557 | out = self.decoder( |
| 558 | x[:, :, i:i + 1, :, :], |
| 559 | feat_cache=self._feat_map, |
| 560 | feat_idx=self._conv_idx) |
| 561 | else: |
| 562 | out_ = self.decoder( |
| 563 | x[:, :, i:i + 1, :, :], |
| 564 | feat_cache=self._feat_map, |
| 565 | feat_idx=self._conv_idx) |
| 566 | out = torch.cat([out, out_], 2) |
| 567 | self.clear_cache() |
| 568 | return out |
| 569 | |
| 570 | def reparameterize(self, mu, log_var): |
| 571 | std = torch.exp(0.5 * log_var) |