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
+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)