diff --git a/TODO.md b/TODO.md index 386e98558..6da94054e 100644 --- a/TODO.md +++ b/TODO.md @@ -4,7 +4,7 @@ Main ToDo list can be found at [GitHub projects](https://github.com/users/vladma ## Current -- HiDream: configurable LLM +- Simplify Flux.1 loader ### Issues/Limitations diff --git a/modules/model_hidream.py b/modules/model_hidream.py index 82cc38218..e8e734477 100644 --- a/modules/model_hidream.py +++ b/modules/model_hidream.py @@ -1,4 +1,3 @@ -import os import time import transformers import diffusers @@ -15,37 +14,11 @@ def hijack_encode_prompt(*args, **kwargs): return res -def get_args(load_config:dict={}, module:str=None, device_map:bool=False): - config = load_config.copy() - modelloader.hf_login() - if 'torch_dtype' not in config: - config['torch_dtype'] = devices.dtype - if 'low_cpu_mem_usage' in config: - del config['low_cpu_mem_usage'] - if 'load_connected_pipeline' in config: - del config['load_connected_pipeline'] - if 'safety_checker' in config: - del config['safety_checker'] - if 'requires_safety_checker' in config: - del config['requires_safety_checker'] - if device_map: - if shared.opts.device_map == 'cpu': - config['device_map'] = 'cpu' - if shared.opts.device_map == 'gpu': - config['device_map'] = devices.device - if devices.backend == "ipex" and os.environ.get('UR_L0_ENABLE_RELAXED_ALLOCATION_LIMITS', '0') != '1' and module in {'TE', 'LLM'}: - # Alchemis GPUs hits the 4GB allocation limit with transformers - # UR_L0_ENABLE_RELAXED_ALLOCATION_LIMITS emulates above 4GB allocations - config['device_map'] = 'cpu' - quant_args = model_quant.create_config(module=module) - return config, quant_args - - def load_hidream(checkpoint_info, diffusers_load_config={}): + modelloader.hf_login() repo_id = sd_models.path_to_repo(checkpoint_info.name) - shared.log.debug(f'Load model: type=HiDream model="{checkpoint_info.name}" repo="{repo_id}" offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype}') - load_args, quant_args = get_args(diffusers_load_config, module='Transformer', device_map=True) + load_args, quant_args = model_quant.get_dit_args(diffusers_load_config, module='Transformer', device_map=True) shared.log.debug(f'Load model: type=HiDream transformer="{repo_id}" quant="{model_quant.get_quant_type(quant_args)}" args={load_args}') transformer = diffusers.HiDreamImageTransformer2DModel.from_pretrained( repo_id, @@ -57,7 +30,7 @@ def load_hidream(checkpoint_info, diffusers_load_config={}): if shared.opts.diffusers_offload_mode != 'none': transformer = transformer.to(devices.cpu) - load_args, quant_args = get_args(diffusers_load_config, module='TE', device_map=True) + load_args, quant_args = model_quant.get_dit_args(diffusers_load_config, module='TE', device_map=True) shared.log.debug(f'Load model: type=HiDream te3="{repo_id}" quant="{model_quant.get_quant_type(quant_args)}" args={load_args}') text_encoder_3 = transformers.T5EncoderModel.from_pretrained( repo_id, @@ -69,7 +42,7 @@ def load_hidream(checkpoint_info, diffusers_load_config={}): if shared.opts.diffusers_offload_mode != 'none': text_encoder_3 = text_encoder_3.to(devices.cpu) - load_args, quant_args = get_args(diffusers_load_config, module='LLM', device_map=True) + load_args, quant_args = model_quant.get_dit_args(diffusers_load_config, module='LLM', device_map=True) shared.log.debug(f'Load model: type=HiDream te4="{shared.opts.model_h1_llama_repo}" quant="{model_quant.get_quant_type(quant_args)}" args={load_args}') tokenizer_4 = transformers.PreTrainedTokenizerFast.from_pretrained( shared.opts.model_h1_llama_repo, @@ -87,7 +60,8 @@ def load_hidream(checkpoint_info, diffusers_load_config={}): if shared.opts.diffusers_offload_mode != 'none': text_encoder_4 = text_encoder_4.to(devices.cpu) - load_args, quant_args = get_args(diffusers_load_config, module='Model') + load_args, _quant_args = model_quant.get_dit_args(diffusers_load_config, module='Model') + shared.log.debug(f'Load model: type=HiDream model="{checkpoint_info.name}" repo="{repo_id}" offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype} args={load_args}') pipe = diffusers.HiDreamImagePipeline.from_pretrained( repo_id, text_encoder_3=text_encoder_3, diff --git a/modules/model_lumina.py b/modules/model_lumina.py index f19fcd7da..f9d3b9abd 100644 --- a/modules/model_lumina.py +++ b/modules/model_lumina.py @@ -3,23 +3,13 @@ import diffusers def load_lumina(_checkpoint_info, diffusers_load_config={}): - from modules import shared, devices, modelloader + from modules import shared, devices, modelloader, model_quant modelloader.hf_login() - # {'low_cpu_mem_usage': True, 'torch_dtype': torch.float16, 'load_connected_pipeline': True, 'safety_checker': None, 'requires_safety_checker': False} - if 'torch_dtype' not in diffusers_load_config: - diffusers_load_config['torch_dtype'] = 'torch.float16' - if 'low_cpu_mem_usage' in diffusers_load_config: - del diffusers_load_config['low_cpu_mem_usage'] - if 'load_connected_pipeline' in diffusers_load_config: - del diffusers_load_config['load_connected_pipeline'] - if 'safety_checker' in diffusers_load_config: - del diffusers_load_config['safety_checker'] - if 'requires_safety_checker' in diffusers_load_config: - del diffusers_load_config['requires_safety_checker'] + load_config, _quant_config = model_quant.get_dit_args(diffusers_load_config, allow_quant=False) pipe = diffusers.LuminaText2ImgPipeline.from_pretrained( 'Alpha-VLLM/Lumina-Next-SFT-diffusers', cache_dir = shared.opts.diffusers_dir, - **diffusers_load_config, + **load_config, ) devices.torch_gc(force=True) return pipe @@ -27,18 +17,35 @@ 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 - quant_args = {} - quant_args = model_quant.create_bnb_config(quant_args) - if quant_args: - model_quant.load_bnb(f'Load model: type=Lumina quant={quant_args}') - if not quant_args: - quant_args = model_quant.create_config() - kwargs = {} repo_id = sd_models.path_to_repo(checkpoint_info.name) - if (('Model' in shared.opts.bnb_quantization or 'Model' in shared.opts.torchao_quantization or 'Model' in shared.opts.quanto_quantization) or ('Transformer' in shared.opts.bnb_quantization or 'Transformer' in shared.opts.torchao_quantization or 'Transformer' in shared.opts.quanto_quantization)): - kwargs['transformer'] = diffusers.Lumina2Transformer2DModel.from_pretrained(repo_id, subfolder="transformer", cache_dir=shared.opts.diffusers_dir, torch_dtype=devices.dtype, **quant_args) - if ('TE' in shared.opts.bnb_quantization or 'TE' in shared.opts.torchao_quantization or 'TE' in shared.opts.quanto_quantization): - kwargs['text_encoder'] = transformers.AutoModel.from_pretrained(repo_id, subfolder="text_encoder", cache_dir=shared.opts.diffusers_dir, torch_dtype=devices.dtype, **quant_args) - sd_model = diffusers.Lumina2Text2ImgPipeline.from_pretrained(checkpoint_info.path, cache_dir=shared.opts.diffusers_dir, **diffusers_load_config, **quant_args, **kwargs) + + 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.hfcache_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( + repo_id, + subfolder="text_encoder", + cache_dir=shared.opts.hfcache_dir, + torch_dtype=devices.dtype, + **load_config, + **quant_config, + ) + + load_config, quant_config = model_quant.get_dit_args(diffusers_load_config, allow_quant=False) + pipe = diffusers.Lumina2Text2ImgPipeline.from_pretrained( + repo_id, + cache_dir=shared.opts.diffusers_dir, + text_encoder=text_encoder, + transformer=transformer, + **load_config, + ) + devices.torch_gc(force=True) - return sd_model + return pipe diff --git a/modules/model_meissonic.py b/modules/model_meissonic.py index 153f377d1..2be5c3bbe 100644 --- a/modules/model_meissonic.py +++ b/modules/model_meissonic.py @@ -17,11 +17,29 @@ def load_meissonic(checkpoint_info, diffusers_load_config={}): diffusers_load_config['variant'] = 'fp16' diffusers_load_config['trust_remote_code'] = True - model = TransformerMeissonic.from_pretrained(fn, subfolder="transformer", cache_dir=cache_dir, **diffusers_load_config) - vqvae = diffusers.VQModel.from_pretrained(fn, subfolder="vqvae", cache_dir=cache_dir, **diffusers_load_config) - text_encoder = transformers.CLIPTextModelWithProjection.from_pretrained(fn, subfolder="text_encoder", cache_dir=cache_dir) - # text_encoder = transformers.CLIPTextModelWithProjection.from_pretrained("laion/CLIP-ViT-H-14-laion2B-s32B-b79K", cache_dir=cache_dir) - tokenizer = transformers.CLIPTokenizer.from_pretrained(fn, subfolder="tokenizer", cache_dir=cache_dir) + + model = TransformerMeissonic.from_pretrained( + fn, + subfolder="transformer", + cache_dir=cache_dir, + **diffusers_load_config, + ) + vqvae = diffusers.VQModel.from_pretrained( + fn, + subfolder="vqvae", + cache_dir=cache_dir, + **diffusers_load_config, + ) + text_encoder = transformers.CLIPTextModelWithProjection.from_pretrained( + fn, + subfolder="text_encoder", + cache_dir=cache_dir, + ) + tokenizer = transformers.CLIPTokenizer.from_pretrained( + fn, + subfolder="tokenizer", + cache_dir=cache_dir, + ) scheduler = MeissonicScheduler.from_pretrained(fn, subfolder="scheduler", cache_dir=cache_dir) pipe = PipelineMeissonic( vqvae=vqvae.to(devices.dtype), diff --git a/modules/model_pixart.py b/modules/model_pixart.py index 0757a1216..58d3dbfa7 100644 --- a/modules/model_pixart.py +++ b/modules/model_pixart.py @@ -1,30 +1,36 @@ +import transformers import diffusers def load_pixart(checkpoint_info, diffusers_load_config={}): - from modules import shared, devices, modelloader, model_te + from modules import shared, devices, modelloader, sd_models, model_quant modelloader.hf_login() - # shared.opts.data['cuda_dtype'] = 'FP32' # override - # shared.opts.data['diffusers_offload_mode}'] = "model" # override - # devices.set_cuda_params() - fn = checkpoint_info.path.replace('huggingface/', '') - t5 = model_te.load_t5(name=shared.opts.sd_text_encoder, cache_dir=shared.opts.diffusers_dir) + repo_id = sd_models.path_to_repo(checkpoint_info.name) + + load_args, quant_args = model_quant.get_dit_args(diffusers_load_config, module='Transformer') transformer = diffusers.PixArtTransformer2DModel.from_pretrained( - fn, - subfolder = 'transformer', - cache_dir = shared.opts.diffusers_dir, - **diffusers_load_config, + repo_id, + subfolder='transformer', + cache_dir=shared.opts.hfcache_dir, + **load_args, + **quant_args, ) - transformer.to(devices.device) - kwargs = { 'transformer': transformer } - if t5 is not None: - kwargs['text_encoder'] = t5 - diffusers_load_config.pop('variant', None) + load_args, quant_args = model_quant.get_dit_args(diffusers_load_config, module='TE', device_map=True) + text_encoder = transformers.T5EncoderModel.from_pretrained( + repo_id, + subfolder="text_encoder", + cache_dir=shared.opts.hfcache_dir, + **load_args, + **quant_args, + ) + + load_args, _quant_args = model_quant.get_dit_args(diffusers_load_config, allow_quant=False) pipe = diffusers.PixArtSigmaPipeline.from_pretrained( 'PixArt-alpha/PixArt-Sigma-XL-2-1024-MS', - cache_dir = shared.opts.diffusers_dir, - **kwargs, - **diffusers_load_config, + cache_dir=shared.opts.diffusers_dir, + transformer=transformer, + text_encoder=text_encoder, + **load_args, ) devices.torch_gc(force=True) return pipe diff --git a/modules/model_quant.py b/modules/model_quant.py index 370235e43..a32a8e353 100644 --- a/modules/model_quant.py +++ b/modules/model_quant.py @@ -455,3 +455,32 @@ def torchao_quantization(sd_model): log.error(f"Quantization: type=TorchAO {e}") setup_logging() # torchao uses dynamo which messes with logging so reset is needed return sd_model + + +def get_dit_args(load_config:dict={}, module:str=None, device_map:bool=False, allow_quant:bool=True): + from modules import shared, devices + config = load_config.copy() + if 'torch_dtype' not in config: + config['torch_dtype'] = devices.dtype + if 'low_cpu_mem_usage' in config: + del config['low_cpu_mem_usage'] + if 'load_connected_pipeline' in config: + del config['load_connected_pipeline'] + if 'safety_checker' in config: + del config['safety_checker'] + if 'requires_safety_checker' in config: + del config['requires_safety_checker'] + if 'variant' in config: + del config['variant'] + if device_map: + if shared.opts.device_map == 'cpu': + config['device_map'] = 'cpu' + if shared.opts.device_map == 'gpu': + config['device_map'] = devices.device + if devices.backend == "ipex" and os.environ.get('UR_L0_ENABLE_RELAXED_ALLOCATION_LIMITS', '0') != '1' and module in {'TE', 'LLM'}: + config['device_map'] = 'cpu' # alchemist gpus hits the 4GB allocation limit with transformers, UR_L0_ENABLE_RELAXED_ALLOCATION_LIMITS emulates above 4GB allocations + if allow_quant: + quant_args = create_config(module=module) + else: + quant_args = {} + return config, quant_args diff --git a/modules/model_sana.py b/modules/model_sana.py index 7f39f17e0..c2bc39119 100644 --- a/modules/model_sana.py +++ b/modules/model_sana.py @@ -20,9 +20,9 @@ def load_quants(kwargs, repo_id, cache_dir): def load_sana(checkpoint_info, kwargs={}): modelloader.hf_login() - fn = checkpoint_info if isinstance(checkpoint_info, str) else checkpoint_info.path repo_id = sd_models.path_to_repo(fn) + kwargs.pop('load_connected_pipeline', None) kwargs.pop('safety_checker', None) kwargs.pop('requires_safety_checker', None) diff --git a/modules/model_sd3.py b/modules/model_sd3.py index 3551e4f2b..c2ac92143 100644 --- a/modules/model_sd3.py +++ b/modules/model_sd3.py @@ -14,7 +14,6 @@ def load_overrides(kwargs, cache_dir): shared.log.debug(f'Load model: type=SD3 unet="{shared.opts.sd_unet}" fmt=safetensors') elif fn.endswith('.gguf'): from modules import ggml - # kwargs = load_gguf(kwargs, fn) kwargs['transformer'] = ggml.load_gguf(fn, cls=diffusers.SD3Transformer2DModel, compute_dtype=devices.dtype) sd_unet.loaded_unet = shared.opts.sd_unet shared.log.debug(f'Load model: type=SD3 unet="{shared.opts.sd_unet}" fmt=gguf') @@ -23,6 +22,7 @@ def load_overrides(kwargs, cache_dir): errors.display(e, 'UNet') shared.opts.sd_unet = 'Default' sd_unet.failed_unet.append(shared.opts.sd_unet) + if shared.opts.sd_text_encoder != 'Default': try: from modules.model_te import load_t5, load_vit_l, load_vit_g @@ -39,6 +39,7 @@ def load_overrides(kwargs, cache_dir): shared.log.error(f"Load model: type=SD3 failed to load T5: {e}") errors.display(e, 'TE') shared.opts.sd_text_encoder = 'Default' + if shared.opts.sd_vae != 'Default' and shared.opts.sd_vae != 'Automatic': try: from modules import sd_vae @@ -55,12 +56,11 @@ def load_overrides(kwargs, cache_dir): def load_quants(kwargs, repo_id, cache_dir): - quant_args = model_quant.create_config() - if not quant_args: - return kwargs - if 'transformer' not in kwargs and (('Model' in shared.opts.bnb_quantization or 'Model' in shared.opts.torchao_quantization or 'Model' in shared.opts.quanto_quantization) or ('Transformer' in shared.opts.bnb_quantization or 'Transformer' in shared.opts.torchao_quantization or 'Transformer' in shared.opts.quanto_quantization)): + quant_args = model_quant.create_config(module='Transformer') + if quant_args and 'quantization_config' in quant_args: kwargs['transformer'] = diffusers.SD3Transformer2DModel.from_pretrained(repo_id, subfolder="transformer", cache_dir=cache_dir, torch_dtype=devices.dtype, **quant_args) - if 'text_encoder_3' not in kwargs and ('TE' in shared.opts.bnb_quantization or 'TE' in shared.opts.torchao_quantization or 'TE' in shared.opts.quanto_quantization): + quant_args = model_quant.create_config(module='TE') + if quant_args and 'quantization_config' in quant_args: kwargs['text_encoder_3'] = transformers.T5EncoderModel.from_pretrained(repo_id, subfolder="text_encoder_3", variant='fp16', cache_dir=cache_dir, torch_dtype=devices.dtype, **quant_args) return kwargs @@ -79,13 +79,12 @@ def load_missing(kwargs, fn, cache_dir): 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) + load_args, quant_args = model_quant.get_dit_args({}, module='TE', device_map=True) + kwargs['text_encoder_3'] = transformers.T5EncoderModel.from_pretrained(repo_id, subfolder="text_encoder_3", variant='fp16', cache_dir=cache_dir, **load_args, **quant_args) shared.log.debug(f'Load model: type=SD3 missing=te3 repo="{repo_id}"') if 'vae' not in kwargs and 'vae' not in keys: kwargs['vae'] = diffusers.AutoencoderKL.from_pretrained(repo_id, subfolder='vae', cache_dir=cache_dir, torch_dtype=devices.dtype) shared.log.debug(f'Load model: type=SD3 missing=vae 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 @@ -93,11 +92,6 @@ def load_sd3(checkpoint_info, cache_dir=None, config=None): repo_id = sd_models.path_to_repo(checkpoint_info.name) fn = checkpoint_info.path - # unload current model - sd_models.unload_model_weights() - shared.sd_model = None - devices.torch_gc(force=True) - kwargs = {} kwargs = load_overrides(kwargs, cache_dir) if (fn is None) or (not os.path.exists(fn) or os.path.isdir(fn)): @@ -107,16 +101,10 @@ def load_sd3(checkpoint_info, cache_dir=None, config=None): if fn is not None and os.path.exists(fn) and os.path.isfile(fn): if fn.endswith('.safetensors'): loader = diffusers.StableDiffusion3Pipeline.from_single_file - # required_modules = model_tools.get_modules(diffusers.StableDiffusion3Pipeline) - # have_modules = model_tools.get_safetensor_keys(fn) - # loaded_modules = model_tools.load_modules('stabilityai/stable-diffusion-3.5-medium', required_modules) - # kwargs = {**kwargs, **loaded_modules} - # kwargs = load_missing(kwargs, fn, cache_dir) repo_id = fn elif fn.endswith('.gguf'): from modules import ggml kwargs['transformer'] = ggml.load_gguf(fn, cls=diffusers.SD3Transformer2DModel, compute_dtype=devices.dtype) - # kwargs = load_gguf(kwargs, fn) kwargs = load_missing(kwargs, fn, cache_dir) kwargs['variant'] = 'fp16' else: @@ -124,7 +112,6 @@ def load_sd3(checkpoint_info, cache_dir=None, config=None): shared.log.debug(f'Load model: type=SD3 kwargs={list(kwargs)} repo="{repo_id}"') - kwargs = model_quant.create_config(kwargs) if shared.opts.model_sd3_disable_te5: shared.log.debug('Load model: type=SD3 option="disable-te5"') kwargs['text_encoder_3'] = None diff --git a/modules/model_te.py b/modules/model_te.py index 2418558f2..698d4506e 100644 --- a/modules/model_te.py +++ b/modules/model_te.py @@ -20,12 +20,14 @@ def load_t5(name=None, cache_dir=None): modelloader.hf_login() repo_id = 'stabilityai/stable-diffusion-3-medium-diffusers' fn = te_dict.get(name) if name in te_dict else None + if fn is not None and name.lower().endswith('gguf'): from modules import ggml ggml.install_gguf() with open(os.path.join('configs', 'flux', 'text_encoder_2', 'config.json'), encoding='utf8') as f: t5_config = transformers.T5Config(**json.load(f)) t5 = transformers.T5EncoderModel.from_pretrained(None, gguf_file=fn, config=t5_config, device_map="auto", cache_dir=cache_dir, torch_dtype=devices.dtype) + elif fn is not None and 'fp8' in name.lower(): from accelerate.utils import set_module_tensor_to_device with open(os.path.join('configs', 'flux', 'text_encoder_2', 'config.json'), encoding='utf8') as f: @@ -45,28 +47,34 @@ def load_t5(name=None, cache_dir=None): try: t5 = t5.to(dtype=devices.dtype) except Exception: - shared.log.error(f"FLUX: Failed to cast text encoder to {devices.dtype}, set dtype to {t5.dtype}") + shared.log.error(f"T5: Failed to cast text encoder to {devices.dtype}, set dtype to {t5.dtype}") raise + elif fn is not None: with open(os.path.join('configs', 'flux', 'text_encoder_2', 'config.json'), encoding='utf8') as f: t5_config = transformers.T5Config(**json.load(f)) state_dict = load_file(fn) t5 = transformers.T5EncoderModel.from_pretrained(None, state_dict=state_dict, config=t5_config) + elif 'fp16' in name.lower(): t5 = transformers.T5EncoderModel.from_pretrained(repo_id, subfolder='text_encoder_3', cache_dir=cache_dir, torch_dtype=devices.dtype) + elif 'fp4' in name.lower(): model_quant.load_bnb('Load model: type=T5') 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) + elif 'fp8' in name.lower(): model_quant.load_bnb('Load model: type=T5') 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) + elif 'qint8' in name.lower(): model_quant.load_quanto('Load model: type=T5') from modules.model_quant import optimum_quanto_model 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 name.lower(): install('nncf==2.7.0', quiet=True) from modules.model_quant import nncf_compress_model @@ -78,12 +86,15 @@ def load_t5(name=None, cache_dir=None): dtype=torch.float32 if devices.dtype != torch.bfloat16 else torch.bfloat16 ) t5 = nncf_compress_model(t5) + elif '/' in name: shared.log.debug(f'Load model: type=T5 repo={name}') quant_config = model_quant.create_config(module='TE') t5 = transformers.T5EncoderModel.from_pretrained(name, cache_dir=cache_dir, torch_dtype=devices.dtype, **quant_config) + else: t5 = None + if t5 is not None: loaded_te = name return t5 @@ -124,8 +135,8 @@ def load_vit_l(): config = transformers.PretrainedConfig.from_json_file('configs/sdxl/text_encoder/config.json') state_dict = load_file(os.path.join(shared.opts.te_dir, f'{shared.opts.sd_text_encoder}.safetensors')) te = transformers.CLIPTextModel.from_pretrained(pretrained_model_name_or_path=None, state_dict=state_dict, config=config) - loaded_te = shared.opts.sd_text_encoder te = te.to(dtype=devices.dtype) + loaded_te = shared.opts.sd_text_encoder return te @@ -134,8 +145,8 @@ def load_vit_g(): config = transformers.PretrainedConfig.from_json_file('configs/sdxl/text_encoder_2/config.json') state_dict = load_file(os.path.join(shared.opts.te_dir, f'{shared.opts.sd_text_encoder}.safetensors')) te = transformers.CLIPTextModelWithProjection.from_pretrained(pretrained_model_name_or_path=None, state_dict=state_dict, config=config) - loaded_te = shared.opts.sd_text_encoder te = te.to(dtype=devices.dtype) + loaded_te = shared.opts.sd_text_encoder return te diff --git a/modules/modelloader.py b/modules/modelloader.py index 7d12646d2..edb26d84f 100644 --- a/modules/modelloader.py +++ b/modules/modelloader.py @@ -365,7 +365,7 @@ def find_diffuser(name: str, full=False): if len(models) == 0: models = list(hf_api.list_models(model_name=name, full=True, limit=20, sort="downloads", direction=-1)) # widen search models = [m for m in models if m.id.startswith(name)] # filter exact - shared.log.debug(f'Searching diffusers models: {name} {len(models) > 0}') + shared.log.debug(f'Search model: repo="{name}" {len(models) > 0}') if len(models) > 0: if not full: return models[0].id diff --git a/modules/sd_models.py b/modules/sd_models.py index cc6d8086f..e724d2ad7 100644 --- a/modules/sd_models.py +++ b/modules/sd_models.py @@ -277,6 +277,8 @@ def load_diffuser_initial(diffusers_load_config, op='model'): def load_diffuser_force(model_type, checkpoint_info, diffusers_load_config, op='model'): sd_model = None + unload_model_weights() + shared.sd_model = None try: if model_type in ['Stable Cascade']: # forced pipeline from modules.model_stablecascade import load_cascade_combined