From ee5978f72bc81c0246b474f72756fce2e07b974b Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Sat, 2 May 2026 10:22:43 +0200 Subject: [PATCH] custom vae loader Co-authored-by: Copilot Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 2 ++ modules/sd_vae.py | 39 ++++++++++++++++++++++------------ pipelines/generic.py | 30 ++++++++++++++++++++++++++ pipelines/model_anima.py | 3 +++ pipelines/model_auraflow.py | 2 ++ pipelines/model_chroma.py | 3 +++ pipelines/model_ernie.py | 2 ++ pipelines/model_flux.py | 2 ++ pipelines/model_flux2.py | 2 ++ pipelines/model_flux2_klein.py | 2 ++ pipelines/model_lumina.py | 5 +++++ pipelines/model_nucleus.py | 2 ++ pipelines/model_pixart.py | 2 ++ pipelines/model_qwen.py | 2 ++ pipelines/model_sd3.py | 2 ++ pipelines/model_z_image.py | 2 ++ 16 files changed, 89 insertions(+), 13 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 89124bfc8..6dae657b5 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -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 diff --git a/modules/sd_vae.py b/modules/sd_vae.py index 8bad35232..92e7ca12a 100644 --- a/modules/sd_vae.py +++ b/modules/sd_vae.py @@ -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}') diff --git a/pipelines/generic.py b/pipelines/generic.py index 63196fe08..523d24ce3 100644 --- a/pipelines/generic.py +++ b/pipelines/generic.py @@ -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') diff --git a/pipelines/model_anima.py b/pipelines/model_anima.py index bd0fc1e7e..fa7b3b6c2 100644 --- a/pipelines/model_anima.py +++ b/pipelines/model_anima.py @@ -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 diff --git a/pipelines/model_auraflow.py b/pipelines/model_auraflow.py index 5d181ebce..597307c52 100644 --- a/pipelines/model_auraflow.py +++ b/pipelines/model_auraflow.py @@ -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) diff --git a/pipelines/model_chroma.py b/pipelines/model_chroma.py index 4472a1a85..1e7f9ea8a 100644 --- a/pipelines/model_chroma.py +++ b/pipelines/model_chroma.py @@ -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) diff --git a/pipelines/model_ernie.py b/pipelines/model_ernie.py index c5ca50691..d0eba17aa 100644 --- a/pipelines/model_ernie.py +++ b/pipelines/model_ernie.py @@ -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) diff --git a/pipelines/model_flux.py b/pipelines/model_flux.py index 23bd40455..17686e95e 100644 --- a/pipelines/model_flux.py +++ b/pipelines/model_flux.py @@ -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') diff --git a/pipelines/model_flux2.py b/pipelines/model_flux2.py index dd6c364b0..a203730cb 100644 --- a/pipelines/model_flux2.py +++ b/pipelines/model_flux2.py @@ -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() diff --git a/pipelines/model_flux2_klein.py b/pipelines/model_flux2_klein.py index 31e539128..3a238805b 100644 --- a/pipelines/model_flux2_klein.py +++ b/pipelines/model_flux2_klein.py @@ -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() diff --git a/pipelines/model_lumina.py b/pipelines/model_lumina.py index c55d8281b..c5bc0f460 100644 --- a/pipelines/model_lumina.py +++ b/pipelines/model_lumina.py @@ -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) diff --git a/pipelines/model_nucleus.py b/pipelines/model_nucleus.py index 1dd4ca89b..b7087fa59 100644 --- a/pipelines/model_nucleus.py +++ b/pipelines/model_nucleus.py @@ -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 diff --git a/pipelines/model_pixart.py b/pipelines/model_pixart.py index 2f91ac895..1ced8b659 100644 --- a/pipelines/model_pixart.py +++ b/pipelines/model_pixart.py @@ -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) diff --git a/pipelines/model_qwen.py b/pipelines/model_qwen.py index 36aa99888..52c9fd251 100644 --- a/pipelines/model_qwen.py +++ b/pipelines/model_qwen.py @@ -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) diff --git a/pipelines/model_sd3.py b/pipelines/model_sd3.py index c19cc1c11..5717dbc38 100644 --- a/pipelines/model_sd3.py +++ b/pipelines/model_sd3.py @@ -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) diff --git a/pipelines/model_z_image.py b/pipelines/model_z_image.py index 3081bb595..e619de619 100644 --- a/pipelines/model_z_image.py +++ b/pipelines/model_z_image.py @@ -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)