From c7fb5b1690fe685edbfdeacccc283986edc06e57 Mon Sep 17 00:00:00 2001 From: Disty0 Date: Thu, 5 Jun 2025 13:24:02 +0300 Subject: [PATCH] SDNQ fix VAE quant --- modules/processing_vae.py | 9 +++++---- 1 file changed, 5 insertions(+), 4 deletions(-) diff --git a/modules/processing_vae.py b/modules/processing_vae.py index 5a375adc2..2d3f8bb7a 100644 --- a/modules/processing_vae.py +++ b/modules/processing_vae.py @@ -125,10 +125,6 @@ def full_vae_decode(latents, model): model.vae.orig_dtype = model.vae.dtype model.vae = model.vae.to(dtype=torch.float32) latents = latents.to(devices.device) - if getattr(model.vae, "post_quant_conv", None) is not None: - latents = latents.to(next(iter(model.vae.post_quant_conv.parameters())).dtype) - else: - latents = latents.to(model.vae.dtype) # normalize latents latents_mean = model.vae.config.get("latents_mean", None) @@ -144,6 +140,11 @@ def full_vae_decode(latents, model): if shift_factor: latents = latents + shift_factor + if getattr(model.vae, "post_quant_conv", None) is not None and "VAE" not in shared.opts.sdnq_quantize_weights: + latents = latents.to(next(iter(model.vae.post_quant_conv.parameters())).dtype) + else: + latents = latents.to(model.vae.dtype) + log_debug(f'VAE config: {model.vae.config}') try: with devices.inference_context():