From 319af31d25ec0ed440aec3786a9aac1f29881083 Mon Sep 17 00:00:00 2001 From: Disty0 Date: Mon, 16 Jun 2025 13:28:30 +0300 Subject: [PATCH 1/2] Custom UNet loading support for Lumina 2 --- CHANGELOG.md | 1 + modules/model_lumina.py | 38 ++++++++++++++++++++++++++++++-------- modules/sd_unet.py | 4 ++-- 3 files changed, 33 insertions(+), 10 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index e9ce65586..520828acd 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -5,6 +5,7 @@ - **Feature** - Support for Python 3.13 - TeaCache support for Lumina 2 + - Custom UNet loading support for Lumina 2 - **Changes** - Increase the medvram mode threshold from 8GB to 12GB diff --git a/modules/model_lumina.py b/modules/model_lumina.py index f4f79d833..23f108623 100644 --- a/modules/model_lumina.py +++ b/modules/model_lumina.py @@ -2,7 +2,9 @@ import os import transformers import diffusers from huggingface_hub import repo_exists -from modules import sd_hijack_te +from modules import errors, shared, sd_unet, sd_hijack_te + +debug = shared.log.trace if os.environ.get('SD_LOAD_DEBUG', None) is not None else lambda *args, **kwargs: None def load_lumina(_checkpoint_info, diffusers_load_config={}): @@ -20,6 +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 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 @@ -30,13 +33,32 @@ def load_lumina2(checkpoint_info, diffusers_load_config={}): diffusers.Lumina2Transformer2DModel.forward = teacache.teacache_lumina2_forward # patch must be done before transformer is loaded load_config, quant_config = model_quant.get_dit_args(diffusers_load_config, module='Transformer') - transformer = diffusers.Lumina2Transformer2DModel.from_pretrained( - repo_id, - subfolder="transformer", - cache_dir=shared.opts.diffusers_dir, - **load_config, - **quant_config, - ) + if shared.opts.sd_unet != 'Default': + try: + debug(f'Load model: type=Lumina2 unet="{shared.opts.sd_unet}"') + transformer = diffusers.Lumina2Transformer2DModel.from_single_file( + sd_unet.unet_dict[shared.opts.sd_unet], + cache_dir=shared.opts.diffusers_dir, + **load_config, + **quant_config + ) + if transformer is None: + shared.opts.sd_unet = 'Default' + sd_unet.failed_unet.append(shared.opts.sd_unet) + except Exception as e: + shared.log.error(f"Load model: type=Lumina2 failed to load UNet: {e}") + shared.opts.sd_unet = 'Default' + if debug: + errors.display(e, 'Lumina2 UNet:') + + if transformer is None: + transformer = diffusers.Lumina2Transformer2DModel.from_pretrained( + repo_id, + subfolder="transformer", + cache_dir=shared.opts.diffusers_dir, + **load_config, + **quant_config, + ) load_config, quant_config = model_quant.get_dit_args(diffusers_load_config, module='TE', device_map=True) text_encoder = transformers.AutoModel.from_pretrained( diff --git a/modules/sd_unet.py b/modules/sd_unet.py index f5f24677c..ccf84b62d 100644 --- a/modules/sd_unet.py +++ b/modules/sd_unet.py @@ -34,7 +34,7 @@ def load_unet(model): if prior_text_encoder is not None: model.prior_pipe.text_encoder = None # Prevent OOM model.prior_pipe.text_encoder = prior_text_encoder.to(devices.device, dtype=devices.dtype) - elif "Flux" in model.__class__.__name__ or "StableDiffusion3" in model.__class__.__name__ or "HiDream" in model.__class__.__name__: + elif "Flux" in model.__class__.__name__ or "StableDiffusion3" in model.__class__.__name__ or "HiDream" in model.__class__.__name__ or "Lumina2" in model.__class__.__name__: loaded_unet = shared.opts.sd_unet sd_models.load_diffuser() # TODO model load: force-reloading entire model as loading transformers only leads to massive memory usage """ @@ -71,7 +71,7 @@ def load_unet(model): def refresh_unet_list(): unet_dict.clear() - for file in files_cache.list_files(shared.opts.unet_dir, ext_filter=[".safetensors", ".gguf"]): + for file in files_cache.list_files(shared.opts.unet_dir, ext_filter=[".safetensors", ".gguf", ".pth"]): basename = os.path.basename(file) name = os.path.splitext(basename)[0] if ".safetensors" in basename else basename unet_dict[name] = file From bddd0913004101d67c4657c4a8c0fc1330a26de5 Mon Sep 17 00:00:00 2001 From: Disty0 Date: Mon, 16 Jun 2025 22:34:02 +0300 Subject: [PATCH 2/2] 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)