This commit is contained in:
Seunghoon Lee
2024-02-06 19:23:59 +09:00
parent 70a15defc5
commit 4bfa0a5650
2 changed files with 10 additions and 11 deletions
+10 -8
View File
@@ -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)
@@ -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)