From 4bfa0a5650b145f5d4addc075e8e4b1c9630ffe2 Mon Sep 17 00:00:00 2001 From: Seunghoon Lee Date: Tue, 6 Feb 2024 19:23:59 +0900 Subject: [PATCH] cleanup --- modules/onnx_impl/__init__.py | 18 ++++++++++-------- .../onnx_stable_diffusion_inpaint_pipeline.py | 3 --- 2 files changed, 10 insertions(+), 11 deletions(-) diff --git a/modules/onnx_impl/__init__.py b/modules/onnx_impl/__init__.py index 0a819c3aa..abdea8e59 100644 --- a/modules/onnx_impl/__init__.py +++ b/modules/onnx_impl/__init__.py @@ -127,20 +127,22 @@ class VAE(TorchCompatibleModule): def device(self): return self.pipeline.vae_decoder.device - def _run(self, model: Callable, latent_sample: torch.Tensor, *_, **__): + def encode(self, sample: torch.Tensor, *_, **__): + sample_np = sample.cpu().numpy() + return [ + torch.from_numpy(np.concatenate( + [self.pipeline.vae_encoder(sample=sample_np[i : i + 1])[0] for i in range(sample_np.shape[0])] + )).to(sample.device) + ] + + def decode(self, latent_sample: torch.Tensor, *_, **__): latents_np = latent_sample.cpu().numpy() return [ torch.from_numpy(np.concatenate( - [model(latent_sample=latents_np[i : i + 1])[0] for i in range(latents_np.shape[0])] + [self.pipeline.vae_decoder(latent_sample=latents_np[i : i + 1])[0] for i in range(latents_np.shape[0])] )).to(latent_sample.device) ] - def encode(self, *args, **kwargs): - return self._run(self.pipeline.vae_encoder, *args, **kwargs) - - def decode(self, *args, **kwargs): - return self._run(self.pipeline.vae_decoder, *args, **kwargs) - def to(self, *args, **kwargs): self.pipeline.vae_encoder = self.pipeline.vae_encoder.to(*args, **kwargs) self.pipeline.vae_decoder = self.pipeline.vae_decoder.to(*args, **kwargs) diff --git a/modules/onnx_impl/pipelines/onnx_stable_diffusion_inpaint_pipeline.py b/modules/onnx_impl/pipelines/onnx_stable_diffusion_inpaint_pipeline.py index c553ae997..dccfb808d 100644 --- a/modules/onnx_impl/pipelines/onnx_stable_diffusion_inpaint_pipeline.py +++ b/modules/onnx_impl/pipelines/onnx_stable_diffusion_inpaint_pipeline.py @@ -28,8 +28,6 @@ class OnnxStableDiffusionInpaintPipeline(diffusers.OnnxStableDiffusionInpaintPip ): super().__init__(vae_encoder, vae_decoder, text_encoder, tokenizer, unet, scheduler, safety_checker, feature_extractor, requires_safety_checker) - self.vae_scale_factor = 2 ** (len(self.vae_decoder.config.get("block_out_channels")) - 1) - @torch.no_grad() def __call__( self, @@ -97,7 +95,6 @@ class OnnxStableDiffusionInpaintPipeline(diffusers.OnnxStableDiffusionInpaintPip generator, latents, num_channels_latents, - self.vae_scale_factor, ) scaling_factor = self.vae_decoder.config.get("scaling_factor", 0.18215)