mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 09:14:35 +02:00
fix custom vae loader
Signed-off-by: Vladimir Mandic <mandic00@live.com>
This commit is contained in:
+14
-15
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user