MCPcopy Create free account
hub / github.com/Francis-Rings/FlashPortrait / encode

Method encode

wan/models/wan_vae_tiled.py:264–280  ·  view source on GitHub ↗

High-level encode with tiling support

(self, x, return_dict=True)

Source from the content-addressed store, hash-verified

262 return dec
263
264 def encode(self, x, return_dict=True):
265 """High-level encode with tiling support"""
266 if x.shape[-2] >= self.tile_sample_min_height and x.shape[-1] >= self.tile_sample_min_width:
267 mu = self.tiled_encode(x)
268 # DiagonalGaussianDistribution expects [mu, logvar] concatenated along channel dim
269 # It will split by channel, so we need to provide 2*z_dim channels
270 # Use zeros for logvar since we're using deterministic encoding (mode)
271 logvar = torch.zeros_like(mu)
272 latent = torch.cat([mu, logvar], dim=1)
273 from diffusers.models.autoencoders.vae import DiagonalGaussianDistribution
274 from diffusers.models.modeling_outputs import AutoencoderKLOutput
275 posterior = DiagonalGaussianDistribution(latent)
276 if not return_dict:
277 return (posterior,)
278 return AutoencoderKLOutput(latent_dist=posterior)
279 else:
280 return self.vae.encode(x, return_dict=return_dict)
281
282 def decode(self, z, return_dict=True):
283 """High-level decode with tiling support"""

Callers

nothing calls this directly

Calls 1

tiled_encodeMethod · 0.95

Tested by

no test coverage detected