From ad3f40f736f57d258fcd55192426e31ea2934c68 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Thu, 24 Oct 2024 13:14:43 -0400 Subject: [PATCH] improve sd3 loader Signed-off-by: Vladimir Mandic --- modules/model_flux.py | 16 ++--- modules/model_flux_nf4.py | 2 +- modules/model_sd3.py | 138 ++++++++++++++++++++++-------------- modules/postprocess/yolo.py | 2 +- modules/sd_models.py | 12 +++- 5 files changed, 106 insertions(+), 64 deletions(-) diff --git a/modules/model_flux.py b/modules/model_flux.py index d696f7df6..38207f73b 100644 --- a/modules/model_flux.py +++ b/modules/model_flux.py @@ -41,7 +41,7 @@ def load_flux_quanto(checkpoint_info): except Exception: shared.log.error(f"Load model: type=FLUX Failed to cast transformer to {devices.dtype}, set dtype to {transformer.dtype}") except Exception as e: - shared.log.error(f"Load model: type=FLUX Failed to load Quanto transformer: {e}") + shared.log.error(f"Load model: type=FLUX failed to load Quanto transformer: {e}") if debug: from modules import errors errors.display(e, 'FLUX Quanto:') @@ -68,7 +68,7 @@ def load_flux_quanto(checkpoint_info): except Exception: shared.log.error(f"Load model: type=FLUX Failed to cast text encoder to {devices.dtype}, set dtype to {text_encoder_2.dtype}") except Exception as e: - shared.log.error(f"Load model: type=FLUX Failed to load Quanto text encoder: {e}") + shared.log.error(f"Load model: type=FLUX failed to load Quanto text encoder: {e}") if debug: from modules import errors errors.display(e, 'FLUX Quanto:') @@ -100,7 +100,7 @@ def load_flux_bnb(checkpoint_info, diffusers_load_config): # pylint: disable=unu else: transformer = diffusers.FluxTransformer2DModel.from_single_file(repo_path, **diffusers_load_config) except Exception as e: - shared.log.error(f"Load model: type=FLUX Failed to load BnB transformer: {e}") + shared.log.error(f"Load model: type=FLUX failed to load BnB transformer: {e}") transformer, text_encoder_2 = None, None if debug: from modules import errors @@ -222,7 +222,7 @@ def load_flux(checkpoint_info, diffusers_load_config): # triggered by opts.sd_ch shared.opts.sd_unet = 'None' sd_unet.failed_unet.append(shared.opts.sd_unet) except Exception as e: - shared.log.error(f"Load model: type=FLUX Failed to load UNet: {e}") + shared.log.error(f"Load model: type=FLUX failed to load UNet: {e}") shared.opts.sd_unet = 'None' if debug: from modules import errors @@ -236,7 +236,7 @@ def load_flux(checkpoint_info, diffusers_load_config): # triggered by opts.sd_ch else: text_encoder_2 = load_t5(name=shared.opts.sd_text_encoder, cache_dir=shared.opts.diffusers_dir) except Exception as e: - shared.log.error(f"Load model: type=FLUX Failed to load T5: {e}") + shared.log.error(f"Load model: type=FLUX failed to load T5: {e}") shared.opts.sd_text_encoder = 'None' if debug: from modules import errors @@ -251,7 +251,7 @@ def load_flux(checkpoint_info, diffusers_load_config): # triggered by opts.sd_ch 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=FLUX Failed to load VAE: {e}") + shared.log.error(f"Load model: type=FLUX failed to load VAE: {e}") shared.opts.sd_vae = 'None' if debug: from modules import errors @@ -267,7 +267,7 @@ def load_flux(checkpoint_info, diffusers_load_config): # triggered by opts.sd_ch if _text_encoder is not None: text_encoder_2 = _text_encoder except Exception as e: - shared.log.error(f"Load model: type=FLUX Failed to load NF4 components: {e}") + shared.log.error(f"Load model: type=FLUX failed to load NF4 components: {e}") if debug: from modules import errors errors.display(e, 'FLUX NF4:') @@ -279,7 +279,7 @@ def load_flux(checkpoint_info, diffusers_load_config): # triggered by opts.sd_ch if _text_encoder is not None: text_encoder_2 = _text_encoder except Exception as e: - shared.log.error(f"Load model: type=FLUX Failed to load Quanto components: {e}") + shared.log.error(f"Load model: type=FLUX failed to load Quanto components: {e}") if debug: from modules import errors errors.display(e, 'FLUX Quanto:') diff --git a/modules/model_flux_nf4.py b/modules/model_flux_nf4.py index a1b46fd54..d023907d6 100644 --- a/modules/model_flux_nf4.py +++ b/modules/model_flux_nf4.py @@ -200,7 +200,7 @@ def load_flux_nf4(checkpoint_info): create_quantized_param(transformer, param, param_name, target_device=0, state_dict=original_state_dict, pre_quantized=True) except Exception as e: transformer, text_encoder_2 = None, None - shared.log.error(f"Load model: type=FLUX Failed to load UNET: {e}") + shared.log.error(f"Load model: type=FLUX failed to load UNET: {e}") if debug: from modules import errors errors.display(e, 'FLUX:') diff --git a/modules/model_sd3.py b/modules/model_sd3.py index 96f194c66..72f3a0c32 100644 --- a/modules/model_sd3.py +++ b/modules/model_sd3.py @@ -1,56 +1,49 @@ import os import diffusers import transformers +from modules import shared, devices, sd_models, sd_unet -default_repo_id = 'stabilityai/stable-diffusion-3-medium' +def load_overrides(kwargs, cache_dir): + if shared.opts.sd_unet != 'None': + try: + fn = sd_unet.unet_dict[shared.opts.sd_unet] + kwargs['transformer'] = diffusers.SD3Transformer2DModel.from_single_file(fn, cache_dir=cache_dir, torch_dtype=devices.dtype) + shared.log.debug(f'Load model: type=SD3 unet="{shared.opts.sd_unet}"') + except Exception as e: + shared.log.error(f"Load model: type=SD3 failed to load UNet: {e}") + shared.opts.sd_unet = 'None' + sd_unet.failed_unet.append(shared.opts.sd_unet) + if shared.opts.sd_text_encoder != 'None': + try: + from modules.model_te import load_t5, load_vit_l, load_vit_g + if 'vit-l' in shared.opts.sd_text_encoder.lower(): + kwargs['text_encoder'] = load_vit_l() + shared.log.debug(f'Load model: type=SD3 variant="vit-l" te="{shared.opts.sd_text_encoder}"') + elif 'vit-g' in shared.opts.sd_text_encoder.lower(): + kwargs['text_encoder_2'] = load_vit_g() + shared.log.debug(f'Load model: type=SD3 variant="vit-g" te="{shared.opts.sd_text_encoder}"') + else: + kwargs['text_encoder_3'] = load_t5(name=shared.opts.sd_text_encoder, cache_dir=shared.opts.diffusers_dir) + shared.log.debug(f'Load model: type=SD3 variant="t5" te="{shared.opts.sd_text_encoder}"') + except Exception as e: + shared.log.error(f"Load model: type=SD3 failed to load T5: {e}") + shared.opts.sd_text_encoder = 'None' + if shared.opts.sd_vae != 'None' and shared.opts.sd_vae != 'Automatic': + try: + from modules import sd_vae + 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') + kwargs['vae'] = diffusers.AutoencoderKL.from_single_file(vae_file, config=vae_config, cache_dir=cache_dir, torch_dtype=devices.dtype) + shared.log.debug(f'Load model: type=SD3 vae="{shared.opts.sd_vae}"') + except Exception as e: + shared.log.error(f"Load model: type=FLUX failed to load VAE: {e}") + shared.opts.sd_vae = 'None' + return kwargs -def load_sd3(checkpoint_info, cache_dir=None, config=None): - from modules import shared, devices, modelloader, sd_models - repo_id = sd_models.path_to_repo(checkpoint_info.name) - dtype = devices.dtype - kwargs = {} - if checkpoint_info.path is not None and checkpoint_info.path.endswith('.safetensors') and os.path.exists(checkpoint_info.path): - loader = diffusers.StableDiffusion3Pipeline.from_single_file - fn_size = os.path.getsize(checkpoint_info.path) - if fn_size < 5e9: - kwargs = { - 'text_encoder': transformers.CLIPTextModelWithProjection.from_pretrained( - default_repo_id, - subfolder='text_encoder', - cache_dir=cache_dir, - torch_dtype=dtype, - ), - 'text_encoder_2': transformers.CLIPTextModelWithProjection.from_pretrained( - default_repo_id, - subfolder='text_encoder_2', - cache_dir=cache_dir, - torch_dtype=dtype, - ), - 'tokenizer': transformers.CLIPTokenizer.from_pretrained( - default_repo_id, - subfolder='tokenizer', - cache_dir=cache_dir, - ), - 'tokenizer_2': transformers.CLIPTokenizer.from_pretrained( - default_repo_id, - subfolder='tokenizer_2', - cache_dir=cache_dir, - ), - 'text_encoder_3': None, - } - elif fn_size < 1e10: # if model is below 10gb it does not have te3 - kwargs = { - 'text_encoder_3': None, - } - else: - kwargs = {} - else: - modelloader.hf_login() - loader = diffusers.StableDiffusion3Pipeline.from_pretrained - kwargs['variant'] = 'fp16' - +def load_quants(kwargs, repo_id, cache_dir): if len(shared.opts.bnb_quantization) > 0: from modules.model_quant import load_bnb load_bnb('Load model: type=SD3') @@ -61,18 +54,57 @@ def load_sd3(checkpoint_info, cache_dir=None, config=None): bnb_4bit_quant_type=shared.opts.bnb_quantization_type, bnb_4bit_compute_dtype=devices.dtype ) - if 'Model' in shared.opts.bnb_quantization: - transformer = diffusers.SD3Transformer2DModel.from_pretrained(repo_id, subfolder="transformer", cache_dir=cache_dir, quantization_config=bnb_config, torch_dtype=devices.dtype) + if 'Model' in shared.opts.bnb_quantization and 'transformer' not in kwargs: + kwargs['transformer'] = diffusers.SD3Transformer2DModel.from_pretrained(repo_id, subfolder="transformer", cache_dir=cache_dir, quantization_config=bnb_config, torch_dtype=devices.dtype) shared.log.debug(f'Quantization: module=transformer type=bnb dtype={shared.opts.bnb_quantization_type} storage={shared.opts.bnb_quantization_storage}') - kwargs['transformer'] = transformer - if 'Text Encoder' in shared.opts.bnb_quantization: - te3 = transformers.T5EncoderModel.from_pretrained(repo_id, subfolder="text_encoder_3", variant='fp16', cache_dir=cache_dir, quantization_config=bnb_config, torch_dtype=devices.dtype) + if 'Text Encoder' in shared.opts.bnb_quantization and 'text_encoder_3' not in kwargs: + kwargs['text_encoder_3'] = transformers.T5EncoderModel.from_pretrained(repo_id, subfolder="text_encoder_3", variant='fp16', cache_dir=cache_dir, quantization_config=bnb_config, torch_dtype=devices.dtype) shared.log.debug(f'Quantization: module=t5 type=bnb dtype={shared.opts.bnb_quantization_type} storage={shared.opts.bnb_quantization_storage}') - kwargs['text_encoder_3'] = te3 + return kwargs + + +def load_missing(kwargs, fn, cache_dir): + keys = sd_models.get_safetensor_keys(fn) + size = os.stat(fn).st_size // 1024 // 1024 + if size > 15000: + repo_id = 'stabilityai/stable-diffusion-3.5-large' + else: + repo_id = 'stabilityai/stable-diffusion-3-medium' + if 'text_encoder' not in kwargs and 'text_encoder' not in keys: + kwargs['text_encoder'] = transformers.CLIPTextModelWithProjection.from_pretrained(repo_id, subfolder='text_encoder', cache_dir=cache_dir, torch_dtype=devices.dtype) + shared.log.debug(f'Load model: type=SD3 missing=te1 repo="{repo_id}"') + if 'text_encoder_2' not in kwargs and 'text_encoder_2' not in keys: + kwargs['text_encoder_2'] = transformers.CLIPTextModelWithProjection.from_pretrained(repo_id, subfolder='text_encoder_2', cache_dir=cache_dir, torch_dtype=devices.dtype) + shared.log.debug(f'Load model: type=SD3 missing=te2 repo="{repo_id}"') + if 'text_encoder_3' not in kwargs and 'text_encoder_3' not in keys: + kwargs['text_encoder_3'] = transformers.T5EncoderModel.from_pretrained(repo_id, subfolder="text_encoder_3", variant='fp16', cache_dir=cache_dir, torch_dtype=devices.dtype) + shared.log.debug(f'Load model: type=SD3 missing=te3 repo="{repo_id}"') + # if 'transformer' not in kwargs and 'transformer' not in keys: + # kwargs['transformer'] = diffusers.SD3Transformer2DModel.from_pretrained(default_repo_id, subfolder="transformer", cache_dir=cache_dir, torch_dtype=devices.dtype) + return kwargs + + +def load_sd3(checkpoint_info, cache_dir=None, config=None): + repo_id = sd_models.path_to_repo(checkpoint_info.name) + fn = checkpoint_info.path + + kwargs = {} + kwargs = load_overrides(kwargs, cache_dir) + kwargs = load_quants(kwargs, repo_id, cache_dir) + + if fn is not None and fn.endswith('.safetensors') and os.path.exists(fn): + kwargs = load_missing(kwargs, fn, cache_dir) + loader = diffusers.StableDiffusion3Pipeline.from_single_file + repo_id = fn + else: + loader = diffusers.StableDiffusion3Pipeline.from_pretrained + kwargs['variant'] = 'fp16' + + shared.log.debug(f'Load model: type=FLUX preloaded={list(kwargs)}') pipe = loader( repo_id, - torch_dtype=dtype, + torch_dtype=devices.dtype, cache_dir=cache_dir, config=config, **kwargs, diff --git a/modules/postprocess/yolo.py b/modules/postprocess/yolo.py index 2f0e12086..b162240e0 100644 --- a/modules/postprocess/yolo.py +++ b/modules/postprocess/yolo.py @@ -56,7 +56,7 @@ class YoloRestorer(Detailer): name = os.path.splitext(os.path.basename(f))[0] if name not in files: self.list[name] = os.path.join(shared.opts.yolo_dir, f) - shared.log.info(f'Available Yolo: path="{shared.opts.yolo_dir} items={len(list(self.list))} downloaded={downloaded}') + shared.log.info(f'Available Yolo: path="{shared.opts.yolo_dir}" items={len(list(self.list))} downloaded={downloaded}') return self.list def dependencies(self): diff --git a/modules/sd_models.py b/modules/sd_models.py index b9623f145..71a389c8a 100644 --- a/modules/sd_models.py +++ b/modules/sd_models.py @@ -417,6 +417,16 @@ def read_state_dict(checkpoint_file, map_location=None, what:str='model'): # pyl return sd +def get_safetensor_keys(filename): + keys = [] + try: + with safetensors.torch.safe_open(filename, framework="pt", device="cpu") as f: + keys = f.keys() + except Exception as e: + shared.log.error(f'Load dict: path="{filename}" {e}') + return keys + + def get_checkpoint_state_dict(checkpoint_info: CheckpointInfo, timer): if not os.path.isfile(checkpoint_info.filename): return None @@ -1088,7 +1098,7 @@ def load_diffuser_force(model_type, checkpoint_info, diffusers_load_config, op=' sd_model = load_flux(checkpoint_info, diffusers_load_config) elif model_type in ['Stable Diffusion 3']: from modules.model_sd3 import load_sd3 - shared.log.debug(f'Load {op}: model="Stable Diffusion 3" variant=medium') + shared.log.debug(f'Load {op}: model="Stable Diffusion 3"') shared.opts.scheduler = 'Default' sd_model = load_sd3(checkpoint_info, cache_dir=shared.opts.diffusers_dir, config=diffusers_load_config.get('config', None)) elif model_type in ['Meissonic']: # forced pipeline