MCPcopy Create free account
hub / github.com/SandAI-org/MAGI-1 / patch_vae_encode

Method patch_vae_encode

inference/pipeline/video_process.py:76–109  ·  view source on GitHub ↗

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)

Source from the content-addressed store, hash-verified

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(

Callers

nothing calls this directly

Calls 2

modeMethod · 0.95

Tested by

no test coverage detected