mirror of
https://github.com/vladmandic/automatic
synced 2026-09-18 16:54:33 +02:00
@@ -5,6 +5,7 @@
|
||||
- **Feature**
|
||||
- Support for Python 3.13
|
||||
- TeaCache support for Lumina 2
|
||||
- Custom UNet and VAE loading support for Lumina 2
|
||||
|
||||
- **Changes**
|
||||
- Increase the medvram mode threshold from 8GB to 12GB
|
||||
@@ -43,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
|
||||
|
||||
|
||||
+47
-8
@@ -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, 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
|
||||
@@ -30,13 +33,47 @@ 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 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,
|
||||
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(
|
||||
@@ -48,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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
+2
-2
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user