(self, latents, parallel=False)
| 300 | |
| 301 | @torch.no_grad() |
| 302 | def decode(self, latents, parallel=False): |
| 303 | if latents.ndim == 4: |
| 304 | latents = latents.unsqueeze(0) |
| 305 | |
| 306 | if self.need_scaled: |
| 307 | latents_mean = torch.tensor(self.latents_mean).view(1, self.z_dim, 1, 1, 1).to(latents.device, latents.dtype) |
| 308 | latents_std = 1.0 / torch.tensor(self.latents_std).view(1, self.z_dim, 1, 1, 1).to(latents.device, latents.dtype) |
| 309 | latents = latents / latents_std + latents_mean |
| 310 | |
| 311 | return self.taehv.decode_video(latents.transpose(1, 2).to(self.dtype), parallel=parallel, show_progress_bar=False).transpose(1, 2).mul_(2).sub_(1) |
| 312 | |
| 313 | @torch.no_grad() |
| 314 | def encode_video(self, vid): |
nothing calls this directly
no test coverage detected