cleanup multiple model loaders

Signed-off-by: Vladimir Mandic <mandic00@live.com>
This commit is contained in:
Vladimir Mandic
2025-04-11 22:16:05 -04:00
parent 78d8bfeba7
commit 0f595d4cc5
11 changed files with 142 additions and 108 deletions
+1 -1
View File
@@ -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
+6 -32
View File
@@ -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,
+33 -26
View File
@@ -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
+23 -5
View File
@@ -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),
+24 -18
View File
@@ -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
+29
View File
@@ -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
+1 -1
View File
@@ -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)
+8 -21
View File
@@ -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
+14 -3
View File
@@ -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
+1 -1
View File
@@ -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
+2
View File
@@ -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