(
encoder_output: torch.Tensor, generator: Optional[torch.Generator] = None, sample_mode: str = "sample"
)
| 468 | |
| 469 | # Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion_img2img.retrieve_latents |
| 470 | def 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 |
no test coverage detected