custom vae loader

Co-authored-by: Copilot <copilot@github.com>
Signed-off-by: Vladimir Mandic <mandic00@live.com>
This commit is contained in:
Vladimir Mandic
2026-05-02 10:22:43 +02:00
parent e74a9a60e5
commit ee5978f72b
16 changed files with 89 additions and 13 deletions
+2
View File
@@ -9,6 +9,8 @@
- option *inputs -> skip processing* to force images to passed to model as-is without any pre-processing
examples of models that support multi-inputs: *qwen-image-edit, flux.2, google-gemini*
- **Anima** support for *img2img* and *inpaint* workflows
- custom **VAE** loader for all pipelines
*note*: vae still needs to be compatible with the model
- **UI**
- add button to manually reorient input/output panels
- all ui panels can be minimized/maximized by clicking on their header
+26 -13
View File
@@ -169,25 +169,38 @@ 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')
log.info(f'Load module: type=VAE model="{vae_file}" source={vae_source} config={diffusers_load_config}')
try:
import diffusers
if os.path.isfile(vae_file):
vae_class = None
vae_loader = None
if shared.sd_loaded and getattr(shared.sd_model, 'vae', None) is not None:
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
vae = diffusers.ConsistencyDecoderVAE.from_pretrained('openai/consistency-decoder', **diffusers_load_config) # consistency decoder does not have from single file, so we'll just download it once more
vae_class = diffusers.ConsistencyDecoderVAE
vae_loader = vae_class.from_pretrained
vae_file = 'openai/consistency-decoder'
elif os.path.getsize(vae_file) < 10000000: # 10MB
vae = diffusers.AutoencoderTiny.from_single_file(vae_file, **diffusers_load_config)
else:
vae = diffusers.AutoencoderKL.from_single_file(vae_file, **diffusers_load_config)
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 = vae.to(devices.dtype_vae)
vae_class = diffusers.AutoencoderTiny
vae_loader = vae_class.from_single_file
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
else:
if 'consistency-decoder' in vae_file:
vae = diffusers.ConsistencyDecoderVAE.from_pretrained(vae_file, **diffusers_load_config)
else:
vae = diffusers.AutoencoderKL.from_pretrained(vae_file, **diffusers_load_config)
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}')
+30
View File
@@ -250,3 +250,33 @@ def load_text_encoder(repo_id, cls_name, load_config=None, subfolder="text_encod
devices.torch_gc()
shared.state.end(jobid)
return text_encoder
def load_vae_override(pipe, load_config=None, override_cls=None, override_args={}):
if (shared.opts.sd_vae in [None, 'None', 'Default', 'Automatic']):
return
if (pipe is None) or (getattr(pipe, 'vae', None) is None):
return
if load_config is None:
load_config = {}
cls = override_cls or pipe.vae.__class__
if not hasattr(cls, 'from_single_file'):
log.error(f'Load model: vae="{shared.opts.sd_vae}" cls={cls.__name__} safetensors=unsupported')
return
load_args, quant_args = model_quant.get_dit_args(load_config, module='VAE')
log.info(f'Load model: vae="{shared.opts.sd_vae}" cls={cls.__name__} args={load_args} quant={quant_args}')
try:
fn = os.path.join(shared.opts.vae_dir, shared.opts.sd_vae)
vae = cls.from_single_file(
fn,
cache_dir=shared.opts.hfcache_dir,
**override_args,
**load_args,
**quant_args,
)
if vae is not None:
pipe.vae = vae
except Exception as e:
log.error(f'Load model: vae="{shared.opts.sd_vae}" cls={cls.__name__} {e}')
# errors.display(e, 'Load')
+3
View File
@@ -125,6 +125,9 @@ def load_anima(checkpoint_info, diffusers_load_config=None):
**load_args,
)
# generic.load_vae_override(pipe, diffusers_load_config, override_cls=diffusers.AutoencoderKLQwenImage, override_args={'low_cpu_mem_usage': False, 'ignore_mismatched_sizes': True})
generic.load_vae_override(pipe, diffusers_load_config)
del text_encoder
del transformer
del llm_adapter
+2
View File
@@ -25,6 +25,8 @@ def load_auraflow(checkpoint_info, diffusers_load_config=None):
**load_args,
)
generic.load_vae_override(pipe, diffusers_load_config)
del text_encoder
del transformer
sd_hijack_te.init_hijack(pipe)
+3
View File
@@ -28,6 +28,9 @@ def load_chroma(checkpoint_info, diffusers_load_config=None):
diffusers.pipelines.auto_pipeline.AUTO_TEXT2IMAGE_PIPELINES_MAPPING["chroma"] = diffusers.ChromaPipeline
diffusers.pipelines.auto_pipeline.AUTO_IMAGE2IMAGE_PIPELINES_MAPPING["chroma"] = diffusers.ChromaImg2ImgPipeline
diffusers.pipelines.auto_pipeline.AUTO_INPAINT_PIPELINES_MAPPING["chroma"] = diffusers.ChromaInpaintPipeline
generic.load_vae_override(pipe, diffusers_load_config)
del text_encoder
del transformer
sd_hijack_te.init_hijack(pipe)
+2
View File
@@ -40,6 +40,8 @@ def load_ernie_image(checkpoint_info, diffusers_load_config=None):
'use_pe': shared.opts.model_ernie_enable_pe,
}
generic.load_vae_override(pipe, diffusers_load_config)
del transformer
del text_encoder
sd_hijack_te.init_hijack(pipe)
+2
View File
@@ -60,6 +60,8 @@ def load_flux(checkpoint_info, diffusers_load_config=None):
**load_args,
)
generic.load_vae_override(pipe, diffusers_load_config)
if os.environ.get('SD_REMOTE_T5', None) is not None:
from modules import sd_te_remote
log.warning('Remote-TE: applying patch')
+2
View File
@@ -31,6 +31,8 @@ def load_flux2(checkpoint_info, diffusers_load_config=None):
diffusers.pipelines.auto_pipeline.AUTO_IMAGE2IMAGE_PIPELINES_MAPPING["flux2"] = diffusers.Flux2Pipeline
diffusers.pipelines.auto_pipeline.AUTO_INPAINT_PIPELINES_MAPPING["flux2"] = diffusers.Flux2Pipeline
generic.load_vae_override(pipe, diffusers_load_config)
from pipelines.flux import flux2_lora
flux2_lora.apply_patch()
+2
View File
@@ -34,6 +34,8 @@ def load_flux2_klein(checkpoint_info, diffusers_load_config=None):
diffusers.pipelines.auto_pipeline.AUTO_IMAGE2IMAGE_PIPELINES_MAPPING["flux2klein"] = diffusers.Flux2KleinPipeline
diffusers.pipelines.auto_pipeline.AUTO_INPAINT_PIPELINES_MAPPING["flux2klein"] = diffusers.Flux2KleinPipeline
generic.load_vae_override(pipe, diffusers_load_config)
from pipelines.flux import flux2_lora
flux2_lora.apply_patch()
+5
View File
@@ -18,6 +18,9 @@ def load_lumina(checkpoint_info, diffusers_load_config=None):
cache_dir = shared.opts.diffusers_dir,
**load_config,
)
generic.load_vae_override(pipe, diffusers_load_config)
sd_hijack_te.init_hijack(pipe)
devices.torch_gc(force=True, reason='load')
return pipe
@@ -47,6 +50,8 @@ def load_lumina2(checkpoint_info, diffusers_load_config=None):
**load_config,
)
generic.load_vae_override(pipe, diffusers_load_config)
del transformer
del text_encoder
sd_hijack_te.init_hijack(pipe)
+2
View File
@@ -42,6 +42,8 @@ def load_nucleus(checkpoint_info, diffusers_load_config=None):
'output_type': 'np',
}
generic.load_vae_override(pipe, diffusers_load_config)
del transformer
del text_encoder
del processor
+2
View File
@@ -35,6 +35,8 @@ def load_pixart(checkpoint_info, diffusers_load_config=None):
**load_args,
)
generic.load_vae_override(pipe, diffusers_load_config)
del text_encoder
del transformer
sd_hijack_te.init_hijack(pipe)
+2
View File
@@ -88,6 +88,8 @@ def load_qwen(checkpoint_info, diffusers_load_config=None):
pipe.task_args['layers'] = shared.opts.model_qwen_layers
pipe.task_args['resolution'] = 640
generic.load_vae_override(pipe, diffusers_load_config)
del text_encoder
del transformer
sd_hijack_te.init_hijack(pipe)
+2
View File
@@ -32,6 +32,8 @@ def load_sd3(checkpoint_info, diffusers_load_config=None):
**load_args,
)
generic.load_vae_override(pipe, diffusers_load_config)
del text_encoder_3
del transformer
sd_hijack_te.init_hijack(pipe)
+2
View File
@@ -52,6 +52,8 @@ def load_z_image(checkpoint_info, diffusers_load_config=None):
**load_args,
)
generic.load_vae_override(pipe, diffusers_load_config)
del transformer
del text_encoder
sd_hijack_te.init_hijack(pipe)