fix custom vae loader

Signed-off-by: Vladimir Mandic <mandic00@live.com>
This commit is contained in:
Vladimir Mandic
2026-05-17 08:37:08 +02:00
parent 297ab4bb60
commit d8433e53cd
5 changed files with 34 additions and 21 deletions
+14 -15
View File
@@ -169,6 +169,7 @@ def load_vae(model_file, vae_file=None, vae_source="unknown-source"):
vae_config = sd_detect.get_load_config(model_file, model_type, config_type='json')
if vae_config is not None:
diffusers_load_config['config'] = os.path.join(vae_config, 'vae')
vae = None
try:
import diffusers
vae_class = None
@@ -177,38 +178,36 @@ def load_vae(model_file, vae_file=None, vae_source="unknown-source"):
vae_class = shared.sd_model.vae.__class__
vae_loader = vae_class.from_single_file if os.path.isfile(vae_file) else vae_class.from_pretrained
elif os.path.isfile(vae_file):
if os.path.getsize(vae_file) > 1310944880: # 1.3GB
size = os.path.getsize(vae_file)
if size > 1310944880: # 1.3GB
vae_class = diffusers.ConsistencyDecoderVAE
vae_loader = vae_class.from_pretrained
vae_file = 'openai/consistency-decoder'
elif os.path.getsize(vae_file) < 10000000: # 10MB
vae_class = diffusers.AutoencoderTiny
vae_loader = vae_class.from_single_file
elif size < 25000000: # 25MB
log.error(f'Load module: type=VAE file="{vae_file}" size={size} invalid')
vae_loader = None
vae_class = None
else: # fallback
vae_class = diffusers.AutoencoderKL
# if getattr(vae.config, 'scaling_factor', 0) == 0.18125 and shared.sd_model_type == 'sdxl':
# vae.config.scaling_factor = 0.13025
# log.debug('Setting model: component=VAE fix scaling factor')
vae_loader = vae_class.from_single_file
vae_loader = vae_class.from_single_file
else:
if 'consistency-decoder' in vae_file:
vae_class = diffusers.ConsistencyDecoderVAE
else: # fallback
vae_class = diffusers.AutoencoderKL
vae_loader = vae_class.from_pretrained
if vae_loader is not None:
log.info(f'Load module: type=VAE model="{vae_file}" source={vae_source} cls={vae_class.__name__} config={diffusers_load_config}')
vae = vae_loader(vae_file, **diffusers_load_config)
vae = vae.to(devices.dtype_vae)
global loaded_vae_file # pylint: disable=global-statement
loaded_vae_file = os.path.basename(vae_file)
# log.debug(f'Diffusers VAE config: {vae.config}')
if shared.opts.diffusers_offload_mode == 'none':
sd_models.move_model(vae, devices.device)
global loaded_vae_file # pylint: disable=global-statement
loaded_vae_file = os.path.basename(vae_file)
if shared.opts.diffusers_offload_mode == 'none':
sd_models.move_model(vae, devices.device)
return vae
except Exception as e:
log.error(f"Load VAE failed: model={vae_file} {e}")
log.error(f"Load module: type=VAE model={vae_file} {e}")
if debug:
errors.display(e, 'VAE')
return None
+7 -2
View File
@@ -4,7 +4,6 @@ import json
import torch
import requests
from PIL import Image
from safetensors.torch import _tobytes
from modules.logger import log
@@ -47,6 +46,12 @@ dtypes = {
}
def tensor_to_bytes(tensor: torch.Tensor) -> bytes:
# Convert tensor to a compact raw byte buffer for remote VAE binary endpoints.
tensor = tensor.detach().to(device="cpu").contiguous()
return tensor.view(torch.uint8).numpy().tobytes()
def remote_decode(latents: torch.Tensor, width: int = 0, height: int = 0, model_type: str | None = None):
from modules import devices, shared, errors, modelloader
tensors = []
@@ -103,7 +108,7 @@ def remote_decode(latents: torch.Tensor, width: int = 0, height: int = 0, model_
url=url,
headers=headers,
params=params,
data=_tobytes(latent, "tensor"),
data=tensor_to_bytes(latent),
timeout=300,
)
if not response.ok:
+11 -2
View File
@@ -86,9 +86,11 @@ def get_model(model_cls, variant=None):
return model_cls, variant
def load_model(model_type = 'decoder', variant = None):
def load_model(model_type = 'decoder', variant = None, vae_file: str | None = None):
global prev_cls, prev_type, prev_model, prev_warnings # pylint: disable=global-statement
model_cls = shared.sd_model_type
model_cls = shared.sd_model_type if shared.sd_loaded else None
if vae_file is not None and os.path.exists(vae_file):
model_cls = 'sdxl'
if model_cls is None or model_cls == 'none':
return None, variant
model_cls, variant = get_model(model_cls, variant)
@@ -131,10 +133,17 @@ def load_model(model_type = 'decoder', variant = None):
else:
from modules.taesd.taesd import TAESD
vae = TAESD(decoder_path=fn if model_type=='decoder' else None, encoder_path=fn if model_type=='encoder' else None)
"""
_vae = diffusers.AutoencoderKL()
from installer import Dot
_config = diffusers.AutoencoderKL().config.copy()
vae.config = Dot(_config) # set config for compatibility with standard vae
"""
if vae is not None:
prev_warnings = False # reset warnings for new model
vae = vae.to(devices.device, dtype=dtype)
TAESD_MODELS[variant]['model'] = vae
vae.config = {}
return vae, variant
elif variant.startswith('Hybrid'):
cfg = CQYAN_MODELS[variant].get(model_cls, None)