diff --git a/CHANGELOG.md b/CHANGELOG.md index b8444c168..5fd63d597 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,5 +1,24 @@ # Change Log for SD.Next +## Update for 2024-09-01 + +- flux improve logging, warn when attempting to load unet as base model +- flux unet support fp8/fp4 quantization +- flux vae support fp16 +- flux lora support additional training tools (*1) +- flux model support loading all-in-one safetensors (*1) + not recommended due to massive duplication of components, but added due to popular demand +- taesd configurable number of layers + can be used to speed-up taesd decoding by reducing number of ops + e.g. if generating 1024px image, reducing layers by 1 will result in preview being 512px + set via *settings -> live preview -> taesd decode layers* +- xhinker prompt parser handle offloaded models +- t5 enum manually downloaded models (*2) + +*notes*: +- (*1) requires `diffusers==0.31.0.dev0` +- (*2) work-in-progress + ## Update for 2024-08-31 ### Highlights for 2024-08-31 diff --git a/modules/model_flux.py b/modules/model_flux.py index 4940a659e..93cf1a641 100644 --- a/modules/model_flux.py +++ b/modules/model_flux.py @@ -5,16 +5,46 @@ import diffusers import transformers from safetensors.torch import load_file from huggingface_hub import hf_hub_download +from accelerate.utils import compute_module_sizes from modules import shared, devices -def load_quanto_transformer(checkpoint_info): - from optimum.quanto import requantize # pylint: disable=no-name-in-module - repo_path = checkpoint_info.path +debug = os.environ.get('SD_LOAD_DEBUG', None) is not None + + +def get_quant(file_path): + if "qint8" in file_path.lower(): + return 'qint8' + if "qint4" in file_path.lower(): + return 'qint4' + if "fp8" in file_path.lower(): + return 'fp8' + if "fp4" in file_path.lower(): + return 'fp4' + if "nf4" in file_path.lower(): + return 'nf4' + return 'none' + + +def load_flux_quanto(checkpoint_info, diffusers_load_config, transformer_only=False): + from installer import install + install('optimum-quanto', quiet=True) + try: + from optimum import quanto # pylint: disable=no-name-in-module + from optimum.quanto import requantize # pylint: disable=no-name-in-module + except Exception as e: + shared.log.error(f"FLUX: Failed to import optimum-quanto: {e}") + raise + quanto.tensor.qbits.QBitsTensor.create = lambda *args, **kwargs: quanto.tensor.qbits.QBitsTensor(*args, **kwargs) + + if isinstance(checkpoint_info, str): + repo_path = checkpoint_info + else: + repo_path = checkpoint_info.path quantization_map = os.path.join(repo_path, "transformer", "quantization_map.json") if not os.path.exists(quantization_map): repo_id = checkpoint_info.name.replace('Diffusers/', '') - quantization_map = hf_hub_download(repo_id, subfolder='transformer', filename='quantization_map.json', cache_dir=shared.opts.diffusers_dir) + quantization_map = hf_hub_download(repo_id, subfolder='transformer', filename='quantization_map.json', **diffusers_load_config) with open(quantization_map, "r", encoding='utf8') as f: quantization_map = json.load(f) state_dict = load_file(os.path.join(repo_path, "transformer", "diffusion_pytorch_model.safetensors")) @@ -23,16 +53,19 @@ def load_quanto_transformer(checkpoint_info): transformer = diffusers.FluxTransformer2DModel.from_config(os.path.join(repo_path, "transformer", "config.json")).to(dtype=dtype) requantize(transformer, state_dict, quantization_map, device=torch.device("cpu")) transformer.eval() - return transformer + if transformer.dtype != devices.dtype: + try: + transformer = transformer.to(dtype=devices.dtype) + except Exception: + shared.log.error(f"FLUX: Failed to cast transformer to {devices.dtype}, set dtype to {transformer.dtype}") + raise + if transformer_only: + return transformer, None - -def load_quanto_text_encoder_2(checkpoint_info): - from optimum.quanto import requantize # pylint: disable=no-name-in-module - repo_path = checkpoint_info.path quantization_map = os.path.join(repo_path, "text_encoder_2", "quantization_map.json") if not os.path.exists(quantization_map): repo_id = checkpoint_info.name.replace('Diffusers/', '') - quantization_map = hf_hub_download(repo_id, subfolder='text_encoder_2', filename='quantization_map.json', cache_dir=shared.opts.diffusers_dir) + quantization_map = hf_hub_download(repo_id, subfolder='text_encoder_2', filename='quantization_map.json', **diffusers_load_config) with open(quantization_map, "r", encoding='utf8') as f: quantization_map = json.load(f) with open(os.path.join(repo_path, "text_encoder_2", "config.json"), encoding='utf8') as f: @@ -43,71 +76,73 @@ def load_quanto_text_encoder_2(checkpoint_info): text_encoder_2 = transformers.T5EncoderModel(t5_config).to(dtype=dtype) requantize(text_encoder_2, state_dict, quantization_map, device=torch.device("cpu")) text_encoder_2.eval() - return text_encoder_2 + if text_encoder_2.dtype != devices.dtype: + try: + text_encoder_2 = text_encoder_2.to(dtype=devices.dtype) + except Exception: + shared.log.error(f"FLUX: Failed to cast text encoder to {devices.dtype}, set dtype to {text_encoder_2.dtype}") + raise + return transformer, text_encoder_2 -def load_transformer(file_path): +def load_flux_bnb(checkpoint_info, diffusers_load_config, transformer_only=False): + if isinstance(checkpoint_info, str): + repo_path = checkpoint_info + else: + repo_path = checkpoint_info.path + from installer import install + install('bitsandbytes', quiet=True) + from diffusers import FluxTransformer2DModel + quant = get_quant(repo_path) + if quant == 'fp8': + quantization_config = transformers.BitsAndBytesConfig(load_in_8bit=True) + transformer = FluxTransformer2DModel.from_single_file(repo_path, **diffusers_load_config, quantization_config=quantization_config) + elif quant == 'fp4': + quantization_config = transformers.BitsAndBytesConfig(load_in_4bit=True) + transformer = FluxTransformer2DModel.from_single_file(repo_path, **diffusers_load_config, quantization_config=quantization_config) + else: + transformer = FluxTransformer2DModel.from_single_file(repo_path, **diffusers_load_config) + if transformer_only: + return transformer, None + # TODO load text_encoder_2 + + +def load_transformer(file_path): # triggered by opts.sd_unet change + quant = get_quant(file_path) diffusers_load_config = { "low_cpu_mem_usage": True, "torch_dtype": devices.dtype, "cache_dir": shared.opts.hfcache_dir, } - from diffusers import FluxTransformer2DModel - transformer = FluxTransformer2DModel.from_single_file(file_path, **diffusers_load_config) + shared.log.info(f'Loading UNet: type=FLUX file="{file_path}" offload={shared.opts.diffusers_offload_mode} quant={quant} dtype={devices.dtype}') + if 'nf4' in quant: + from modules.model_flux_nf4 import load_flux_nf4 + transformer = load_flux_nf4(file_path, diffusers_load_config, transformer_only=True) + elif quant == 'qint8' or quant == 'qint4': + transformer, _ = load_flux_quanto(file_path, diffusers_load_config, transformer_only=True) + elif quant == 'fp8' or quant == 'fp4': + transformer, _ = load_flux_bnb(file_path, diffusers_load_config, transformer_only=True) + else: + from diffusers import FluxTransformer2DModel + transformer = FluxTransformer2DModel.from_single_file(file_path, **diffusers_load_config) if transformer is None: shared.log.error('Failed to load UNet model') + if debug: + shared.log.debug(f'FLUX transformer: size={round(compute_module_sizes(transformer)[""] / 1024 / 1204)}') return transformer -def load_flux(checkpoint_info, diffusers_load_config): - if "qint8" in checkpoint_info.path.lower(): - quant = 'qint8' - elif "qint4" in checkpoint_info.path.lower(): - quant = 'qint4' - elif "nf4" in checkpoint_info.path.lower(): - quant = 'nf4' - else: - quant = None - shared.log.debug(f'Loading FLUX: model="{checkpoint_info.name}" quant={quant}') +def load_flux(checkpoint_info, diffusers_load_config): # triggered by opts.sd_checkpoint change + quant = get_quant(checkpoint_info.path) + shared.log.debug(f'Loading FLUX: model="{checkpoint_info.name}" quant={quant} offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype}') if quant == 'nf4': - from installer import install - install('bitsandbytes', quiet=True) - try: - import bitsandbytes # pylint: disable=unused-import - except Exception as e: - shared.log.error(f"FLUX: Failed to import bitsandbytes: {e}") - raise from modules.model_flux_nf4 import load_flux_nf4 pipe = load_flux_nf4(checkpoint_info, diffusers_load_config) elif quant == 'qint8' or quant == 'qint4': - from installer import install - install('optimum-quanto', quiet=True) - try: - from optimum import quanto # pylint: disable=no-name-in-module - except Exception as e: - shared.log.error(f"FLUX: Failed to import optimum-quanto: {e}") - raise - quanto.tensor.qbits.QBitsTensor.create = lambda *args, **kwargs: quanto.tensor.qbits.QBitsTensor(*args, **kwargs) pipe = diffusers.FluxPipeline.from_pretrained(checkpoint_info.path, cache_dir=shared.opts.diffusers_dir, transformer=None, text_encoder_2=None, **diffusers_load_config) - pipe.transformer = load_quanto_transformer(checkpoint_info) - pipe.text_encoder_2 = load_quanto_text_encoder_2(checkpoint_info) - if pipe.transformer.dtype != devices.dtype: - try: - pipe.transformer = pipe.transformer.to(dtype=devices.dtype) - except Exception: - shared.log.error(f"FLUX: Failed to cast transformer to {devices.dtype}, set dtype to {pipe.transformer.dtype}") - raise - if pipe.text_encoder_2.dtype != devices.dtype: - try: - pipe.text_encoder_2 = pipe.text_encoder_2.to(dtype=devices.dtype) - except Exception: - shared.log.error(f"FLUX: Failed to cast text encoder to {devices.dtype}, set dtype to {pipe.text_encoder_2.dtype}") - raise + pipe.transformer, pipe.text_encoder_2 = load_flux_quanto(checkpoint_info, diffusers_load_config) else: pipe = diffusers.FluxPipeline.from_pretrained(checkpoint_info.path, cache_dir=shared.opts.diffusers_dir, **diffusers_load_config) - if devices.dtype == torch.float16 and not shared.opts.no_half_vae: - shared.log.warning("FLUX: does not support FP16 VAE, enabling no-half-vae") - shared.opts.no_half_vae = True - # from accelerate.utils import compute_module_sizes - # shared.log.debug(f'FLUX computed size: {round(compute_module_sizes(pipe.transformer)[""] / 1024 / 1204)}') + if debug: + shared.log.debug(f'FLUX transformer: size={round(compute_module_sizes(pipe.transformer)[""] / 1024 / 1204)}') return pipe diff --git a/modules/model_flux_nf4.py b/modules/model_flux_nf4.py index a28dc3526..61d852862 100644 --- a/modules/model_flux_nf4.py +++ b/modules/model_flux_nf4.py @@ -5,7 +5,6 @@ Copied from: https://github.com/huggingface/diffusers/issues/9165 import os import torch import torch.nn as nn -import bitsandbytes as bnb from transformers.quantizers.quantizers_utils import get_module_from_name from huggingface_hub import hf_hub_download from accelerate import init_empty_weights @@ -16,6 +15,22 @@ import safetensors.torch from modules import shared, devices +bnb = None +debug = os.environ.get('SD_LOAD_DEBUG', None) is not None + + +def load_bnb(): + from installer import install + install('bitsandbytes', quiet=True) + try: + import bitsandbytes + global bnb # pylint: disable=global-statement + bnb = bitsandbytes + except Exception as e: + shared.log.error(f"FLUX: Failed to import bitsandbytes: {e}") + raise + + def _replace_with_bnb_linear( model, method="nf4", @@ -148,25 +163,31 @@ def create_quantized_param( module._parameters[tensor_name] = new_value # pylint: disable=protected-access -def load_flux_nf4(checkpoint_info, diffusers_load_config): - repo_path = checkpoint_info.path +def load_flux_nf4(checkpoint_info, diffusers_load_config, transformer_only=False): + load_bnb() + if isinstance(checkpoint_info, str): + repo_path = checkpoint_info + else: + repo_path = checkpoint_info.path if os.path.exists(repo_path) and os.path.isfile(repo_path): ckpt_path = repo_path - if os.path.exists(repo_path) and os.path.isdir(repo_path) and os.path.exists(os.path.join(repo_path, "diffusion_pytorch_model.safetensors")): + elif os.path.exists(repo_path) and os.path.isdir(repo_path) and os.path.exists(os.path.join(repo_path, "diffusion_pytorch_model.safetensors")): ckpt_path = os.path.join(repo_path, "diffusion_pytorch_model.safetensors") else: ckpt_path = hf_hub_download(repo_path, filename="diffusion_pytorch_model.safetensors", cache_dir=shared.opts.diffusers_dir) original_state_dict = safetensors.torch.load_file(ckpt_path) - if 'sayakpaul' in checkpoint_info.path: + if 'sayakpaul' in repo_path: converted_state_dict = original_state_dict # already converted else: try: converted_state_dict = convert_flux_transformer_checkpoint_to_diffusers(original_state_dict) except Exception as e: - from modules import errors - errors.display(e, 'FLUX convert:') - raise + shared.log.error(f"FLUX: Failed to convert UNET: {e}") + if debug: + from modules import errors + errors.display(e, 'FLUX convert:') + converted_state_dict = original_state_dict with init_empty_weights(): config = FluxTransformer2DModel.load_config("black-forest-labs/flux.1-dev", subfolder="transformer") @@ -187,6 +208,9 @@ def load_flux_nf4(checkpoint_info, diffusers_load_config): create_quantized_param(model, param, param_name, target_device=0, state_dict=original_state_dict, pre_quantized=True) del original_state_dict - pipe = FluxPipeline.from_pretrained("black-forest-labs/flux.1-dev", transformer=model, cache_dir=shared.opts.diffusers_dir, **diffusers_load_config) devices.torch_gc(force=True) - return pipe + if transformer_only: + return model + else: + pipe = FluxPipeline.from_pretrained("black-forest-labs/flux.1-dev", transformer=model, cache_dir=shared.opts.diffusers_dir, **diffusers_load_config) + return pipe diff --git a/modules/model_t5.py b/modules/model_t5.py index 2e1ff06c9..fda3bc6fe 100644 --- a/modules/model_t5.py +++ b/modules/model_t5.py @@ -1,54 +1,40 @@ +import os import torch import transformers +from modules import shared, devices, files_cache + + +t5_dict = {} def load_t5(t5=None, cache_dir=None): - from modules import devices, modelloader + from modules import modelloader repo_id = 'stabilityai/stable-diffusion-3-medium-diffusers' - if 'fp16' in t5.lower(): + fn = t5_dict.get(t5) if t5 in t5_dict else None + if fn is not None: + shared.log.error(f'Loading T5: file="{fn}" unsupported') + t5 = None + elif 'fp16' in t5.lower(): modelloader.hf_login() - t5 = transformers.T5EncoderModel.from_pretrained( - repo_id, - subfolder='text_encoder_3', - # torch_dtype=dtype, - cache_dir=cache_dir, - torch_dtype=devices.dtype, - ) + t5 = transformers.T5EncoderModel.from_pretrained(repo_id, subfolder='text_encoder_3', cache_dir=cache_dir, torch_dtype=devices.dtype) elif 'fp4' in t5.lower(): modelloader.hf_login() from installer import install install('bitsandbytes', quiet=True) quantization_config = transformers.BitsAndBytesConfig(load_in_4bit=True) - t5 = transformers.T5EncoderModel.from_pretrained( - repo_id, - subfolder='text_encoder_3', - quantization_config=quantization_config, - cache_dir=cache_dir, - torch_dtype=devices.dtype, - ) + t5 = transformers.T5EncoderModel.from_pretrained(repo_id, subfolder='text_encoder_3', quantization_config=quantization_config, cache_dir=cache_dir, torch_dtype=devices.dtype) elif 'fp8' in t5.lower(): modelloader.hf_login() from installer import install install('bitsandbytes', quiet=True) quantization_config = transformers.BitsAndBytesConfig(load_in_8bit=True) - t5 = transformers.T5EncoderModel.from_pretrained( - repo_id, - subfolder='text_encoder_3', - quantization_config=quantization_config, - cache_dir=cache_dir, - torch_dtype=devices.dtype, - ) + t5 = transformers.T5EncoderModel.from_pretrained(repo_id, subfolder='text_encoder_3', quantization_config=quantization_config, cache_dir=cache_dir, torch_dtype=devices.dtype) elif 'qint8' in t5.lower(): modelloader.hf_login() from installer import install install('optimum-quanto', quiet=True) from modules.sd_models_compile import optimum_quanto_model - t5 = transformers.T5EncoderModel.from_pretrained( - repo_id, - subfolder='text_encoder_3', - cache_dir=cache_dir, - torch_dtype=devices.dtype, - ) + t5 = transformers.T5EncoderModel.from_pretrained(repo_id, subfolder='text_encoder_3', cache_dir=cache_dir, torch_dtype=devices.dtype) t5 = optimum_quanto_model(t5, weights="qint8", activations="none") elif 'int8' in t5.lower(): modelloader.hf_login() @@ -56,12 +42,7 @@ def load_t5(t5=None, cache_dir=None): install('nncf==2.7.0', quiet=True) from modules.sd_models_compile import nncf_compress_model from modules.sd_hijack import NNCF_T5DenseGatedActDense - t5 = transformers.T5EncoderModel.from_pretrained( - repo_id, - subfolder='text_encoder_3', - cache_dir=cache_dir, - torch_dtype=devices.dtype, - ) + t5 = transformers.T5EncoderModel.from_pretrained(repo_id, subfolder='text_encoder_3', cache_dir=cache_dir, torch_dtype=devices.dtype) for i in range(len(t5.encoder.block)): t5.encoder.block[i].layer[1].DenseReluDense = NNCF_T5DenseGatedActDense( t5.encoder.block[i].layer[1].DenseReluDense, @@ -74,10 +55,11 @@ def load_t5(t5=None, cache_dir=None): def set_t5(pipe, module, t5=None, cache_dir=None): - from modules import devices, shared if pipe is None or not hasattr(pipe, module): return pipe t5 = load_t5(t5=t5, cache_dir=cache_dir) + if module == "text_encoder_2" and t5 is None: # do not unload te2 + return setattr(pipe, module, t5) if shared.opts.diffusers_offload_mode == "sequential": from accelerate import cpu_offload @@ -90,3 +72,11 @@ def set_t5(pipe, module, t5=None, cache_dir=None): pipe.maybe_free_model_hooks() devices.torch_gc() return pipe + + +def refresh_t5_list(): + t5_dict.clear() + for file in files_cache.list_files(shared.opts.t5_dir, ext_filter=[".safetensors"]): + name = os.path.splitext(os.path.basename(file))[0] + t5_dict[name] = file + shared.log.debug(f'Available T5s: path="{shared.opts.t5_dir}" items={len(t5_dict)}') diff --git a/modules/paths.py b/modules/paths.py index 763eddd31..e71e8e917 100644 --- a/modules/paths.py +++ b/modules/paths.py @@ -101,6 +101,7 @@ def create_paths(opts): create_path(fix_path('diffusers_dir')) create_path(fix_path('vae_dir')) create_path(fix_path('unet_dir')) + create_path(fix_path('t5_dir')) create_path(fix_path('lora_dir')) create_path(fix_path('embeddings_dir')) create_path(fix_path('hypernetwork_dir')) diff --git a/modules/processing_vae.py b/modules/processing_vae.py index e5108f0d3..60cf50444 100644 --- a/modules/processing_vae.py +++ b/modules/processing_vae.py @@ -99,7 +99,7 @@ def taesd_vae_decode(latents): debug(f'VAE decode: name=TAESD images={len(latents)} latents={latents.shape} slicing={shared.opts.diffusers_vae_slicing}') if len(latents) == 0: return [] - if shared.opts.diffusers_vae_slicing: + if shared.opts.diffusers_vae_slicing and len(latents) > 1: decoded = torch.zeros((len(latents), 3, latents.shape[2] * 8, latents.shape[3] * 8), dtype=devices.dtype_vae, device=devices.device) for i in range(latents.shape[0]): decoded[i] = sd_vae_taesd.decode(latents[i]) diff --git a/modules/prompt_parser_diffusers.py b/modules/prompt_parser_diffusers.py index 09c8c1899..31cb7c68f 100644 --- a/modules/prompt_parser_diffusers.py +++ b/modules/prompt_parser_diffusers.py @@ -463,13 +463,13 @@ def get_xhinker_text_embeddings(pipe, prompt: str = "", neg_prompt: str = "", cl te1_device, te2_device, te3_device = None, None, None if hasattr(pipe, "text_encoder") and pipe.text_encoder.device != devices.device: te1_device = pipe.text_encoder.device - pipe.text_encoder = pipe.text_encoder.to(devices.device) + sd_models.move_model(pipe.text_encoder, devices.device) if hasattr(pipe, "text_encoder_2") and pipe.text_encoder_2.device != devices.device: te2_device = pipe.text_encoder_2.device - pipe.text_encoder_2 = pipe.text_encoder_2.to(devices.device) + sd_models.move_model(pipe.text_encoder_2, devices.device) if hasattr(pipe, "text_encoder_3") and pipe.text_encoder_3.device != devices.device: te3_device = pipe.text_encoder_3.device - pipe.text_encoder_3 = pipe.text_encoder_3.to(devices.device) + sd_models.move_model(pipe.text_encoder_3, devices.device) if is_sd3: prompt_embed, negative_embed, positive_pooled, negative_pooled = get_weighted_text_embeddings_sd3(pipe=pipe, prompt=prompt, neg_prompt=neg_prompt, use_t5_encoder=bool(pipe.text_encoder_3)) @@ -481,10 +481,10 @@ def get_xhinker_text_embeddings(pipe, prompt: str = "", neg_prompt: str = "", cl prompt_embed, negative_embed = get_weighted_text_embeddings_sd15(pipe=pipe, prompt=prompt, neg_prompt=neg_prompt, clip_skip=clip_skip) if te1_device is not None: - pipe.text_encoder = pipe.text_encoder.to(te1_device) + sd_models.move_model(pipe.text_encoder, te1_device) if te2_device is not None: - pipe.text_encoder_2 = pipe.text_encoder_2.to(te2_device) + sd_models.move_model(pipe.text_encoder_2, te1_device) if te3_device is not None: - pipe.text_encoder_3 = pipe.text_encoder_3.to(te3_device) + sd_models.move_model(pipe.text_encoder_3, te1_device) return prompt_embed, positive_pooled, negative_embed, negative_pooled diff --git a/modules/prompt_parser_xhinker.py b/modules/prompt_parser_xhinker.py index 6e43c8860..b1c3efc69 100644 --- a/modules/prompt_parser_xhinker.py +++ b/modules/prompt_parser_xhinker.py @@ -269,12 +269,12 @@ def get_weighted_text_embeddings_sd15( # get positive prompt embeddings with weights token_tensor = torch.tensor( [prompt_token_groups[i]] - , dtype=torch.long, device=pipe.device + , dtype=torch.long, device=pipe.text_encoder.device ) weight_tensor = torch.tensor( prompt_weight_groups[i] , dtype=torch.float16 - , device=pipe.device + , device=pipe.text_encoder.device ) token_embedding = pipe.text_encoder(token_tensor)[0].squeeze(0) @@ -286,12 +286,12 @@ def get_weighted_text_embeddings_sd15( # get negative prompt embeddings with weights neg_token_tensor = torch.tensor( [neg_prompt_token_groups[i]] - , dtype=torch.long, device=pipe.device + , dtype=torch.long, device=pipe.text_encoder.device ) neg_weight_tensor = torch.tensor( neg_prompt_weight_groups[i] , dtype=torch.float16 - , device=pipe.device + , device=pipe.text_encoder.device ) neg_token_embedding = pipe.text_encoder(neg_token_tensor)[0].squeeze(0) for z in range(len(neg_weight_tensor)): @@ -449,36 +449,36 @@ def get_weighted_text_embeddings_sdxl( # get positive prompt embeddings with weights token_tensor = torch.tensor( [prompt_token_groups[i]] - , dtype=torch.long, device=pipe.device + , dtype=torch.long, device=pipe.text_encoder.device ) weight_tensor = torch.tensor( prompt_weight_groups[i] , dtype=torch.float16 - , device=pipe.device + , device=pipe.text_encoder.device ) token_tensor_2 = torch.tensor( [prompt_token_groups_2[i]] - , dtype=torch.long, device=pipe.device + , dtype=torch.long, device=pipe.text_encoder_2.device ) # use first text encoder prompt_embeds_1 = pipe.text_encoder( - token_tensor.to(pipe.device) + token_tensor.to(pipe.text_encoder.device) , output_hidden_states=True ) prompt_embeds_1_hidden_states = prompt_embeds_1.hidden_states[-2] # use second text encoder prompt_embeds_2 = pipe.text_encoder_2( - token_tensor_2.to(pipe.device) + token_tensor_2.to(pipe.text_encoder_2.device) , output_hidden_states=True ) prompt_embeds_2_hidden_states = prompt_embeds_2.hidden_states[-2] pooled_prompt_embeds = prompt_embeds_2[0] prompt_embeds_list = [prompt_embeds_1_hidden_states, prompt_embeds_2_hidden_states] - token_embedding = torch.concat(prompt_embeds_list, dim=-1).squeeze(0).to(pipe.device) + token_embedding = torch.concat(prompt_embeds_list, dim=-1).squeeze(0).to(pipe.text_encoder.device) for j in range(len(weight_tensor)): if weight_tensor[j] != 1.0: @@ -509,35 +509,35 @@ def get_weighted_text_embeddings_sdxl( # get negative prompt embeddings with weights neg_token_tensor = torch.tensor( [neg_prompt_token_groups[i]] - , dtype=torch.long, device=pipe.device + , dtype=torch.long, device=pipe.text_encoder.device ) neg_token_tensor_2 = torch.tensor( [neg_prompt_token_groups_2[i]] - , dtype=torch.long, device=pipe.device + , dtype=torch.long, device=pipe.text_encoder_2.device ) neg_weight_tensor = torch.tensor( neg_prompt_weight_groups[i] , dtype=torch.float16 - , device=pipe.device + , device=pipe.text_encoder.device ) # use first text encoder neg_prompt_embeds_1 = pipe.text_encoder( - neg_token_tensor.to(pipe.device) + neg_token_tensor.to(pipe.text_encoder.device) , output_hidden_states=True ) neg_prompt_embeds_1_hidden_states = neg_prompt_embeds_1.hidden_states[-2] # use second text encoder neg_prompt_embeds_2 = pipe.text_encoder_2( - neg_token_tensor_2.to(pipe.device) + neg_token_tensor_2.to(pipe.text_encoder_2.device) , output_hidden_states=True ) neg_prompt_embeds_2_hidden_states = neg_prompt_embeds_2.hidden_states[-2] negative_pooled_prompt_embeds = neg_prompt_embeds_2[0] neg_prompt_embeds_list = [neg_prompt_embeds_1_hidden_states, neg_prompt_embeds_2_hidden_states] - neg_token_embedding = torch.concat(neg_prompt_embeds_list, dim=-1).squeeze(0).to(pipe.device) + neg_token_embedding = torch.concat(neg_prompt_embeds_list, dim=-1).squeeze(0).to(pipe.text_encoder.device) for z in range(len(neg_weight_tensor)): if neg_weight_tensor[z] != 1.0: @@ -657,18 +657,18 @@ def get_weighted_text_embeddings_sdxl_refiner( # get positive prompt embeddings with weights token_tensor_2 = torch.tensor( [prompt_token_groups_2[i]] - , dtype=torch.long, device=pipe.device + , dtype=torch.long, device=pipe.text_encoder_2.device ) weight_tensor_2 = torch.tensor( prompt_weight_groups_2[i] , dtype=torch.float16 - , device=pipe.device + , device=pipe.text_encoder_2.device ) # use second text encoder prompt_embeds_2 = pipe.text_encoder_2( - token_tensor_2.to(pipe.device) + token_tensor_2.to(pipe.text_encoder_2.device) , output_hidden_states=True ) prompt_embeds_2_hidden_states = prompt_embeds_2.hidden_states[-2] @@ -703,17 +703,17 @@ def get_weighted_text_embeddings_sdxl_refiner( # get negative prompt embeddings with weights neg_token_tensor_2 = torch.tensor( [neg_prompt_token_groups_2[i]] - , dtype=torch.long, device=pipe.device + , dtype=torch.long, device=pipe.text_encoder_2.device ) neg_weight_tensor_2 = torch.tensor( neg_prompt_weight_groups_2[i] , dtype=torch.float16 - , device=pipe.device + , device=pipe.text_encoder_2.device ) # use second text encoder neg_prompt_embeds_2 = pipe.text_encoder_2( - neg_token_tensor_2.to(pipe.device) + neg_token_tensor_2.to(pipe.text_encoder_2.device) , output_hidden_states=True ) neg_prompt_embeds_2_hidden_states = neg_prompt_embeds_2.hidden_states[-2] @@ -787,8 +787,6 @@ def get_weighted_text_embeddings_sdxl_2p( """ prompt_2 = prompt_2 or prompt neg_prompt_2 = neg_prompt_2 or neg_prompt - - import math eos = pipe.tokenizer.eos_token_id # tokenizer 1 @@ -907,33 +905,33 @@ def get_weighted_text_embeddings_sdxl_2p( # get positive prompt embeddings with weights token_tensor = torch.tensor( [prompt_token_groups[i]] - , dtype=torch.long, device=pipe.device + , dtype=torch.long, device=pipe.text_encoder.device ) weight_tensor = torch.tensor( prompt_weight_groups[i] - , device=pipe.device + , device=pipe.text_encoder.device ) token_tensor_2 = torch.tensor( [prompt_token_groups_2[i]] - , device=pipe.device + , device=pipe.text_encoder_2.device ) weight_tensor_2 = torch.tensor( prompt_weight_groups_2[i] - , device=pipe.device + , device=pipe.text_encoder_2.device ) # use first text encoder prompt_embeds_1 = pipe.text_encoder( - token_tensor.to(pipe.device) + token_tensor.to(pipe.text_encoder.device) , output_hidden_states=True ) prompt_embeds_1_hidden_states = prompt_embeds_1.hidden_states[-2] # use second text encoder prompt_embeds_2 = pipe.text_encoder_2( - token_tensor_2.to(pipe.device) + token_tensor_2.to(pipe.text_encoder_2.device) , output_hidden_states=True ) prompt_embeds_2_hidden_states = prompt_embeds_2.hidden_states[-2] @@ -966,31 +964,31 @@ def get_weighted_text_embeddings_sdxl_2p( # get negative prompt embeddings with weights neg_token_tensor = torch.tensor( [neg_prompt_token_groups[i]] - , device=pipe.device + , device=pipe.text_encoder.device ) neg_token_tensor_2 = torch.tensor( [neg_prompt_token_groups_2[i]] - , device=pipe.device + , device=pipe.text_encoder_2.device ) neg_weight_tensor = torch.tensor( neg_prompt_weight_groups[i] - , device=pipe.device + , device=pipe.text_encoder.device ) neg_weight_tensor_2 = torch.tensor( neg_prompt_weight_groups_2[i] - , device=pipe.device + , device=pipe.text_encoder_2.device ) # use first text encoder neg_prompt_embeds_1 = pipe.text_encoder( - neg_token_tensor.to(pipe.device) + neg_token_tensor.to(pipe.text_encoder.device) , output_hidden_states=True ) neg_prompt_embeds_1_hidden_states = neg_prompt_embeds_1.hidden_states[-2] # use second text encoder neg_prompt_embeds_2 = pipe.text_encoder_2( - neg_token_tensor_2.to(pipe.device) + neg_token_tensor_2.to(pipe.text_encoder_2.device) , output_hidden_states=True ) neg_prompt_embeds_2_hidden_states = neg_prompt_embeds_2.hidden_states[-2] @@ -1049,7 +1047,6 @@ def get_weighted_text_embeddings_sd3( pooled_prompt_embeds (torch.Tensor) negative_pooled_prompt_embeds (torch.Tensor) """ - import math eos = pipe.tokenizer.eos_token_id # tokenizer 1 @@ -1161,22 +1158,22 @@ def get_weighted_text_embeddings_sd3( # get positive prompt embeddings with weights token_tensor = torch.tensor( [prompt_token_groups[i]] - , dtype=torch.long, device=pipe.device + , dtype=torch.long, device=pipe.text_encoder.device ) weight_tensor = torch.tensor( prompt_weight_groups[i] , dtype=torch.float16 - , device=pipe.device + , device=pipe.text_encoder.device ) token_tensor_2 = torch.tensor( [prompt_token_groups_2[i]] - , dtype=torch.long, device=pipe.device + , dtype=torch.long, device=pipe.text_encoder_2.device ) # use first text encoder prompt_embeds_1 = pipe.text_encoder( - token_tensor.to(pipe.device) + token_tensor.to(pipe.text_encoder.device) , output_hidden_states=True ) prompt_embeds_1_hidden_states = prompt_embeds_1.hidden_states[-2] @@ -1184,14 +1181,14 @@ def get_weighted_text_embeddings_sd3( # use second text encoder prompt_embeds_2 = pipe.text_encoder_2( - token_tensor_2.to(pipe.device) + token_tensor_2.to(pipe.text_encoder_2.device) , output_hidden_states=True ) prompt_embeds_2_hidden_states = prompt_embeds_2.hidden_states[-2] pooled_prompt_embeds_2 = prompt_embeds_2[0] prompt_embeds_list = [prompt_embeds_1_hidden_states, prompt_embeds_2_hidden_states] - token_embedding = torch.concat(prompt_embeds_list, dim=-1).squeeze(0).to(pipe.device) + token_embedding = torch.concat(prompt_embeds_list, dim=-1).squeeze(0).to(pipe.text_encoder.device) for j in range(len(weight_tensor)): if weight_tensor[j] != 1.0: @@ -1222,21 +1219,21 @@ def get_weighted_text_embeddings_sd3( # get negative prompt embeddings with weights neg_token_tensor = torch.tensor( [neg_prompt_token_groups[i]] - , dtype=torch.long, device=pipe.device + , dtype=torch.long, device=pipe.text_encoder.device ) neg_token_tensor_2 = torch.tensor( [neg_prompt_token_groups_2[i]] - , dtype=torch.long, device=pipe.device + , dtype=torch.long, device=pipe.text_encoder_2.device ) neg_weight_tensor = torch.tensor( neg_prompt_weight_groups[i] , dtype=torch.float16 - , device=pipe.device + , device=pipe.text_encoder.device ) # use first text encoder neg_prompt_embeds_1 = pipe.text_encoder( - neg_token_tensor.to(pipe.device) + neg_token_tensor.to(pipe.text_encoder.device) , output_hidden_states=True ) neg_prompt_embeds_1_hidden_states = neg_prompt_embeds_1.hidden_states[-2] @@ -1244,14 +1241,14 @@ def get_weighted_text_embeddings_sd3( # use second text encoder neg_prompt_embeds_2 = pipe.text_encoder_2( - neg_token_tensor_2.to(pipe.device) + neg_token_tensor_2.to(pipe.text_encoder_2.device) , output_hidden_states=True ) neg_prompt_embeds_2_hidden_states = neg_prompt_embeds_2.hidden_states[-2] negative_pooled_prompt_embeds_2 = neg_prompt_embeds_2[0] neg_prompt_embeds_list = [neg_prompt_embeds_1_hidden_states, neg_prompt_embeds_2_hidden_states] - neg_token_embedding = torch.concat(neg_prompt_embeds_list, dim=-1).squeeze(0).to(pipe.device) + neg_token_embedding = torch.concat(neg_prompt_embeds_list, dim=-1).squeeze(0).to(pipe.text_encoder.device) for z in range(len(neg_weight_tensor)): if neg_weight_tensor[z] != 1.0: @@ -1286,8 +1283,8 @@ def get_weighted_text_embeddings_sd3( # ----------------- generate positive t5 embeddings -------------------- prompt_tokens_3 = torch.tensor([prompt_tokens_3], dtype=torch.long) - t5_prompt_embeds = pipe.text_encoder_3(prompt_tokens_3.to(pipe.device))[0].squeeze(0) - t5_prompt_embeds = t5_prompt_embeds.to(device=pipe.device) + t5_prompt_embeds = pipe.text_encoder_3(prompt_tokens_3.to(pipe.text_encoder_3.device))[0].squeeze(0) + t5_prompt_embeds = t5_prompt_embeds.to(device=pipe.text_encoder_3.device) # add weight to t5 prompt for z in range(len(prompt_weights_3)): @@ -1296,7 +1293,7 @@ def get_weighted_text_embeddings_sd3( t5_prompt_embeds = t5_prompt_embeds.unsqueeze(0) else: t5_prompt_embeds = torch.zeros(1, 4096, dtype=prompt_embeds.dtype).unsqueeze(0) - t5_prompt_embeds = t5_prompt_embeds.to(device=pipe.device) + t5_prompt_embeds = t5_prompt_embeds.to(device=pipe.text_encoder_3.device) # merge with the clip embedding 1 and clip embedding 2 clip_prompt_embeds = torch.nn.functional.pad( @@ -1308,8 +1305,8 @@ def get_weighted_text_embeddings_sd3( # ---------------------- get neg t5 embeddings ------------------------- neg_prompt_tokens_3 = torch.tensor([neg_prompt_tokens_3], dtype=torch.long) - t5_neg_prompt_embeds = pipe.text_encoder_3(neg_prompt_tokens_3.to(pipe.device))[0].squeeze(0) - t5_neg_prompt_embeds = t5_neg_prompt_embeds.to(device=pipe.device) + t5_neg_prompt_embeds = pipe.text_encoder_3(neg_prompt_tokens_3.to(pipe.pipe.text_encoder_3.device))[0].squeeze(0) + t5_neg_prompt_embeds = t5_neg_prompt_embeds.to(device=pipe.text_encoder_3.device) # add weight to neg t5 embeddings for z in range(len(neg_prompt_weights_3)): @@ -1318,7 +1315,7 @@ def get_weighted_text_embeddings_sd3( t5_neg_prompt_embeds = t5_neg_prompt_embeds.unsqueeze(0) else: t5_neg_prompt_embeds = torch.zeros(1, 4096, dtype=prompt_embeds.dtype).unsqueeze(0) - t5_neg_prompt_embeds = t5_prompt_embeds.to(device=pipe.device) + t5_neg_prompt_embeds = t5_prompt_embeds.to(device=pipe.text_encoder_3.device) clip_neg_prompt_embeds = torch.nn.functional.pad( negative_prompt_embeds, (0, t5_neg_prompt_embeds.shape[-1] - negative_prompt_embeds.shape[-1]) @@ -1359,7 +1356,7 @@ def get_weighted_text_embeddings_flux1( """ prompt2 = prompt if prompt2 is None else prompt2 if device is None: - device = pipe.device + device = pipe.text_encoder.device # tokenizer 1 - openai/clip-vit-large-patch14 prompt_tokens, prompt_weights = get_prompts_tokens_with_weights( diff --git a/modules/sd_models.py b/modules/sd_models.py index 529924976..5b1b9f021 100644 --- a/modules/sd_models.py +++ b/modules/sd_models.py @@ -566,7 +566,7 @@ def detect_pipeline(f: str, op: str = 'model', warning=True, quiet=False): # guess by size if os.path.isfile(f) and f.endswith('.safetensors'): size = round(os.path.getsize(f) / 1024 / 1024) - if size < 128: + if (size > 0 and size < 128): warn(f'Model size smaller than expected: {f} size={size} MB') elif (size >= 316 and size <= 324) or (size >= 156 and size <= 164): # 320 or 160 warn(f'Model detected as VAE model, but attempting to load as model: {op}={f} size={size} MB') @@ -591,6 +591,8 @@ def detect_pipeline(f: str, op: str = 'model', warning=True, quiet=False): guess = 'Stable Diffusion XL' elif (size > 5692 and size < 5698) or (size > 4134 and size < 4138) or (size > 10362 and size < 10366) or (size > 15028 and size < 15228): guess = 'Stable Diffusion 3' + elif (size > 20000 and size < 40000): + guess = 'FLUX' # guess by name """ if 'LCM_' in f.upper() or 'LCM-' in f.upper() or '_LCM' in f.upper() or '-LCM' in f.upper(): @@ -620,8 +622,10 @@ def detect_pipeline(f: str, op: str = 'model', warning=True, quiet=False): guess = 'Kolors' if 'auraflow' in f.lower(): guess = 'AuraFlow' - if 'flux.1' in f.lower() or 'flux1' in f.lower(): + if 'flux' in f.lower(): guess = 'FLUX' + if size > 11000 and size < 20000: + warn(f'Model detected as FLUX UNET model, but attempting to load a base model: {op}={f} size={size} MB') # switch for specific variant if guess == 'Stable Diffusion' and 'inpaint' in f.lower(): guess = 'Stable Diffusion Inpaint' diff --git a/modules/sd_samplers.py b/modules/sd_samplers.py index 29034d63f..04d3e9d37 100644 --- a/modules/sd_samplers.py +++ b/modules/sd_samplers.py @@ -71,7 +71,7 @@ def create_sampler(name, model): sampler = config.constructor(model) if shared.sd_model_type == 'f1': if 'base_image_seq_len' not in sampler.sampler.config or 'max_image_seq_len' not in sampler.sampler.config or 'base_shift' not in sampler.sampler.config or 'max_shift' not in sampler.sampler.config: - shared.log.warning('FLUX sampler: attempting to use a non compatible scheduler') + shared.log.warning(f'FLUX: sampler="{name}" non compatible') return None if not hasattr(model, 'scheduler_config'): model.scheduler_config = sampler.sampler.config.copy() diff --git a/modules/sd_unet.py b/modules/sd_unet.py index c948c223b..16d942a22 100644 --- a/modules/sd_unet.py +++ b/modules/sd_unet.py @@ -1,8 +1,9 @@ import os -from modules import shared, devices, files_cache +from modules import shared, devices, files_cache, sd_models unet_dict = {} +debug = os.environ.get('SD_LOAD_DEBUG', None) is not None def load_unet(model): @@ -28,15 +29,13 @@ def load_unet(model): model.prior_pipe.text_encoder = None # Prevent OOM model.prior_pipe.text_encoder = prior_text_encoder.to(devices.device, dtype=devices.dtype) if "Flux" in model.__class__.__name__: - shared.log.info(f'Loading UNet: name="{shared.opts.sd_unet}" file="{unet_dict[shared.opts.sd_unet]}" offload={shared.opts.diffusers_offload_mode}') from modules.model_flux import load_transformer transformer = load_transformer(unet_dict[shared.opts.sd_unet]) if transformer is not None: model.transformer = None if shared.opts.diffusers_offload_mode == 'none': - model.transformer = transformer.to(devices.device, devices.dtype) - else: - model.transformer = transformer + sd_models.move_model(transformer, devices.device) + model.transformer = transformer from modules.sd_models import set_diffuser_offload set_diffuser_offload(model, 'model') else: @@ -52,6 +51,9 @@ def load_unet(model): model.unet = unet.to(devices.device, devices.dtype_unet) except Exception as e: shared.log.error(f'Failed to load UNet model: {e}') + if debug: + from modules import errors + errors.display(e, 'UNet load:') return devices.torch_gc() diff --git a/modules/sd_vae.py b/modules/sd_vae.py index 0b5993256..698f89d76 100644 --- a/modules/sd_vae.py +++ b/modules/sd_vae.py @@ -233,6 +233,8 @@ def load_vae_diffusers(model_file, vae_file=None, vae_source="unknown-source"): global loaded_vae_file # pylint: disable=global-statement loaded_vae_file = os.path.basename(vae_file) # shared.log.debug(f'Diffusers VAE config: {vae.config}') + if shared.opts.diffusers_offload_mode == 'none': + sd_models.move_model(vae, devices.device) return vae except Exception as e: shared.log.error(f"Loading VAE failed: model={vae_file} {e}") diff --git a/modules/sd_vae_taesd.py b/modules/sd_vae_taesd.py index 3400f05f9..5cd7fab7c 100644 --- a/modules/sd_vae_taesd.py +++ b/modules/sd_vae_taesd.py @@ -55,6 +55,8 @@ def Decoder(latent_channels=4): return nn.Sequential( Clamp(), conv(latent_channels, 64), nn.ReLU(), Block(64, 64), Block(64, 64), Block(64, 64), nn.Upsample(scale_factor=2), conv(64, 64, bias=False), + Block(64, 64), Block(64, 64), Block(64, 64), nn.Identity(), conv(64, 64, bias=False), + Block(64, 64), Block(64, 64), Block(64, 64), nn.Identity(), conv(64, 64, bias=False), Block(64, 64), conv(64, 3), ) elif shared.opts.live_preview_taesd_layers == 2: @@ -62,6 +64,7 @@ def Decoder(latent_channels=4): Clamp(), conv(latent_channels, 64), nn.ReLU(), Block(64, 64), Block(64, 64), Block(64, 64), nn.Upsample(scale_factor=2), conv(64, 64, bias=False), Block(64, 64), Block(64, 64), Block(64, 64), nn.Upsample(scale_factor=2), conv(64, 64, bias=False), + Block(64, 64), Block(64, 64), Block(64, 64), nn.Identity(), conv(64, 64, bias=False), Block(64, 64), conv(64, 3), ) else: @@ -86,9 +89,9 @@ class TAESD(nn.Module): # pylint: disable=abstract-method self.encoder = Encoder(latent_channels) self.decoder = Decoder(latent_channels) if encoder_path is not None: - self.encoder.load_state_dict(torch.load(encoder_path, map_location="cpu")) + self.encoder.load_state_dict(torch.load(encoder_path, map_location="cpu"), strict=False) if decoder_path is not None: - self.decoder.load_state_dict(torch.load(decoder_path, map_location="cpu")) + self.decoder.load_state_dict(torch.load(decoder_path, map_location="cpu"), strict=False) def guess_latent_channels(self, decoder_path, encoder_path): """guess latent channel count based on encoder filename""" diff --git a/modules/shared.py b/modules/shared.py index e1aafca8a..0fceaa7da 100644 --- a/modules/shared.py +++ b/modules/shared.py @@ -406,7 +406,8 @@ options_templates.update(options_section(('sd', "Execution & Models"), { "sd_model_refiner": OptionInfo('None', "Refiner model", gr.Dropdown, lambda: {"choices": ['None'] + list_checkpoint_tiles()}, refresh=refresh_checkpoints), "sd_vae": OptionInfo("Automatic", "VAE model", gr.Dropdown, lambda: {"choices": shared_items.sd_vae_items()}, refresh=shared_items.refresh_vae_list), "sd_unet": OptionInfo("None", "UNET model", gr.Dropdown, lambda: {"choices": shared_items.sd_unet_items()}, refresh=shared_items.refresh_unet_list), - "sd_text_encoder": OptionInfo('None', "Text encoder model", gr.Dropdown, lambda: {"choices": ['None', 'T5 FP4', 'T5 FP8', 'T5 INT8', 'T5 QINT8', 'T5 FP16']}), + # "sd_text_encoder": OptionInfo('None', "Text encoder model", gr.Dropdown, lambda: {"choices": ['None', 'T5 FP4', 'T5 FP8', 'T5 INT8', 'T5 QINT8', 'T5 FP16']}), + "sd_text_encoder": OptionInfo('None', "Text encoder model", gr.Dropdown, lambda: {"choices": shared_items.sd_t5_items()}, refresh=shared_items.refresh_t5_list), "sd_model_dict": OptionInfo('None', "Use separate base dict", gr.Dropdown, lambda: {"choices": ['None'] + list_checkpoint_tiles()}, refresh=refresh_checkpoints), "sd_checkpoint_autoload": OptionInfo(True, "Model autoload on start"), "sd_textencoder_cache": OptionInfo(True, "Cache text encoder results"), @@ -575,6 +576,7 @@ options_templates.update(options_section(('system-paths', "System Paths"), { "hfcache_dir": OptionInfo(os.path.join(os.path.expanduser('~'), '.cache', 'huggingface', 'hub'), "Folder for Huggingface cache", folder=True), "vae_dir": OptionInfo(os.path.join(paths.models_path, 'VAE'), "Folder with VAE files", folder=True), "unet_dir": OptionInfo(os.path.join(paths.models_path, 'UNET'), "Folder with UNET files", folder=True), + "t5_dir": OptionInfo(os.path.join(paths.models_path, 'T5'), "Folder with T5 files", folder=True), "sd_lora": OptionInfo("", "Add LoRA to prompt", gr.Textbox, {"visible": False}), "lora_dir": OptionInfo(os.path.join(paths.models_path, 'Lora'), "Folder with LoRA network(s)", folder=True), "lyco_dir": OptionInfo(os.path.join(paths.models_path, 'LyCORIS'), "Folder with LyCORIS network(s)", gr.Text, {"visible": False}), diff --git a/modules/shared_items.py b/modules/shared_items.py index 1b0077ef8..9f110f413 100644 --- a/modules/shared_items.py +++ b/modules/shared_items.py @@ -23,6 +23,17 @@ def refresh_unet_list(): modules.sd_unet.refresh_unet_list() +def sd_t5_items(): + import modules.model_t5 + predefined = ['None', 'T5 FP4', 'T5 FP8', 'T5 INT8', 'T5 QINT8', 'T5 FP16'] + return predefined + list(modules.model_t5.t5_dict) + + +def refresh_t5_list(): + import modules.model_t5 + modules.model_t5.refresh_t5_list() + + def list_crossattention(diffusers=False): if diffusers: return [ diff --git a/webui.py b/webui.py index c9af90e75..210f9e00c 100644 --- a/webui.py +++ b/webui.py @@ -24,6 +24,7 @@ import modules.scripts import modules.sd_models import modules.sd_vae import modules.sd_unet +import modules.model_t5 import modules.progress import modules.ui import modules.txt2img @@ -90,6 +91,9 @@ def initialize(): modules.sd_unet.refresh_unet_list() timer.startup.record("unet") + modules.model_t5.refresh_t5_list() + timer.startup.record("unet") + extensions.list_extensions() timer.startup.record("extensions")