MCPcopy Create free account
hub / github.com/UVA-Computer-Vision-Lab/FrameINO / retrieve_latents

Function retrieve_latents

train_code/train_wan_motion.py:470–480  ·  view source on GitHub ↗
(
    encoder_output: torch.Tensor, generator: Optional[torch.Generator] = None, sample_mode: str = "sample"
)

Source from the content-addressed store, hash-verified

468
469# Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion_img2img.retrieve_latents
470def retrieve_latents(
471 encoder_output: torch.Tensor, generator: Optional[torch.Generator] = None, sample_mode: str = "sample"
472):
473 if hasattr(encoder_output, "latent_dist") and sample_mode == "sample":
474 return encoder_output.latent_dist.sample(generator)
475 elif hasattr(encoder_output, "latent_dist") and sample_mode == "argmax":
476 return encoder_output.latent_dist.mode()
477 elif hasattr(encoder_output, "latents"):
478 return encoder_output.latents
479 else:
480 raise AttributeError("Could not access latents of provided encoder_output")
481
482
483

Callers 2

Calls 1

sampleMethod · 0.45

Tested by

no test coverage detected