Encode the input video. Args: x (torch.Tensor): Input video tensor with shape (N, C, T, H, W). sample_posterior (bool): Whether to sample from the posterior. Returns: torch.Tensor: Encoded tensor with additional information.
(vae: callable, x: torch.Tensor)
| 74 | @staticmethod |
| 75 | @torch.no_grad() |
| 76 | def patch_vae_encode(vae: callable, x: torch.Tensor) -> torch.Tensor: |
| 77 | """ |
| 78 | Encode the input video. |
| 79 | |
| 80 | Args: |
| 81 | x (torch.Tensor): Input video tensor with shape (N, C, T, H, W). |
| 82 | sample_posterior (bool): Whether to sample from the posterior. |
| 83 | |
| 84 | Returns: |
| 85 | torch.Tensor: Encoded tensor with additional information. |
| 86 | """ |
| 87 | if not isinstance(x, torch.Tensor): |
| 88 | raise TypeError(f"Expected input x to be torch.Tensor, but got {type(x)}.") |
| 89 | if len(x.shape) != 5: |
| 90 | raise ValueError(f"Expected input tensor x to have shape (N, C, T, H, W), but got {x.shape}.") |
| 91 | |
| 92 | if not hasattr(vae, "encoder") or not callable(vae.encoder): |
| 93 | raise AttributeError("Encoder is not defined or callable. Please initialize 'self.encoder'.") |
| 94 | |
| 95 | # for setting vae encoding to deterministic |
| 96 | N, C, T, H, W = x.shape |
| 97 | if T == 1: |
| 98 | x = x.expand(-1, -1, 4, -1, -1) |
| 99 | x = vae.encoder(x) |
| 100 | posterior = DiagonalGaussianDistribution(x) |
| 101 | z = posterior.mode() |
| 102 | |
| 103 | return z[:, :, :1, :, :].type(x.dtype) |
| 104 | else: |
| 105 | x = vae.encoder(x) |
| 106 | posterior = DiagonalGaussianDistribution(x) |
| 107 | z = posterior.mode() |
| 108 | |
| 109 | return z.type(x.dtype) |
| 110 | |
| 111 | @staticmethod |
| 112 | def encode( |
nothing calls this directly
no test coverage detected