From bddd0913004101d67c4657c4a8c0fc1330a26de5 Mon Sep 17 00:00:00 2001 From: Disty0 Date: Mon, 16 Jun 2025 22:34:02 +0300 Subject: [PATCH] Custom VAE loading support for Lumina 2 --- CHANGELOG.md | 3 ++- modules/model_lumina.py | 19 ++++++++++++++++++- modules/sd_models.py | 2 +- 3 files changed, 21 insertions(+), 3 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 520828acd..ed2ebb762 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -5,7 +5,7 @@ - **Feature** - Support for Python 3.13 - TeaCache support for Lumina 2 - - Custom UNet loading support for Lumina 2 + - Custom UNet and VAE loading support for Lumina 2 - **Changes** - Increase the medvram mode threshold from 8GB to 12GB @@ -44,6 +44,7 @@ - VAE Tiling with non-default tile sizes - Lumina 2 with IPEX - Nunchaku updated repo + - Double loading of models with custom UNets ## Update for 2025-06-02 diff --git a/modules/model_lumina.py b/modules/model_lumina.py index 23f108623..d817fa48c 100644 --- a/modules/model_lumina.py +++ b/modules/model_lumina.py @@ -22,7 +22,7 @@ def load_lumina(_checkpoint_info, diffusers_load_config={}): def load_lumina2(checkpoint_info, diffusers_load_config={}): from modules import shared, devices, sd_models, model_quant - transformer, text_encoder = None, None + transformer, text_encoder, vae = None, None, None repo_id = sd_models.path_to_repo(checkpoint_info.name) if os.path.isdir(checkpoint_info.filename) and not repo_exists(repo_id): repo_id = checkpoint_info.filename @@ -51,6 +51,21 @@ def load_lumina2(checkpoint_info, diffusers_load_config={}): if debug: errors.display(e, 'Lumina2 UNet:') + if shared.opts.sd_vae != 'Default' and shared.opts.sd_vae != 'Automatic': + try: + debug(f'Load model: type=Lumina2 vae="{shared.opts.sd_vae}"') + from modules import sd_vae + # vae = sd_vae.load_vae_diffusers(None, sd_vae.vae_dict[shared.opts.sd_vae], 'override') + vae_file = sd_vae.vae_dict[shared.opts.sd_vae] + if os.path.exists(vae_file): + vae_config = os.path.join('configs', 'flux', 'vae', 'config.json') + vae = diffusers.AutoencoderKL.from_single_file(vae_file, config=vae_config, **diffusers_load_config) + except Exception as e: + shared.log.error(f"Load model: type=Lumina2 failed to load VAE: {e}") + shared.opts.sd_vae = 'Default' + if debug: + errors.display(e, 'Lumina2 VAE:') + if transformer is None: transformer = diffusers.Lumina2Transformer2DModel.from_pretrained( repo_id, @@ -70,6 +85,8 @@ def load_lumina2(checkpoint_info, diffusers_load_config={}): ) load_config, quant_config = model_quant.get_dit_args(diffusers_load_config, allow_quant=False) + if vae is not None: + load_config['vae'] = vae pipe = diffusers.Lumina2Pipeline.from_pretrained( repo_id, cache_dir=shared.opts.diffusers_dir, diff --git a/modules/sd_models.py b/modules/sd_models.py index 0bea615ff..d8aad13a3 100644 --- a/modules/sd_models.py +++ b/modules/sd_models.py @@ -591,7 +591,7 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No if "Kandinsky" in sd_model.__class__.__name__: # need a special case sd_model.scheduler.name = 'DDIM' - if model_type not in ['Stable Cascade']: # need a special-case + if hasattr(sd_model, "unet") and model_type not in ['Stable Cascade']: # others calls load_diffuser again sd_unet.load_unet(sd_model) add_noise_pred_to_diffusers_callback(sd_model)