mirror of
https://github.com/vladmandic/automatic
synced 2026-09-18 16:54:33 +02:00
cleanup
This commit is contained in:
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user