fix loading custom t5

Signed-off-by: Vladimir Mandic <mandic00@live.com>
This commit is contained in:
Vladimir Mandic
2025-08-13 16:06:20 -04:00
parent 8ca74d0cd2
commit 41ae06bd90
3 changed files with 18 additions and 10 deletions
+9 -2
View File
@@ -15,10 +15,14 @@ def load_t5(name=None, cache_dir=None):
global loaded_te # pylint: disable=global-statement
if name is None:
return None
cache_dir = cache_dir or shared.opts.hfcache_dir
from modules import modelloader
modelloader.hf_login()
repo_id = 'stabilityai/stable-diffusion-3-medium-diffusers'
fn = te_dict.get(name) if name in te_dict else None
if os.path.exists(name):
fn = name
else:
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
@@ -46,12 +50,13 @@ def load_t5(name=None, cache_dir=None):
except Exception:
shared.log.error(f"T5: Failed to cast text encoder to {devices.dtype}, set dtype to {t5.dtype}")
raise
del state_dict
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)
t5 = transformers.T5EncoderModel.from_pretrained(None, state_dict=state_dict, config=t5_config, torch_dtype=devices.dtype)
elif 'fp16' in name.lower():
t5 = transformers.T5EncoderModel.from_pretrained(repo_id, subfolder='text_encoder_3', cache_dir=cache_dir, torch_dtype=devices.dtype)
@@ -141,6 +146,7 @@ def load_vit_l():
te = transformers.CLIPTextModel.from_pretrained(pretrained_model_name_or_path=None, state_dict=state_dict, config=config)
te = te.to(dtype=devices.dtype)
loaded_te = shared.opts.sd_text_encoder
del state_dict
return te
@@ -151,6 +157,7 @@ def load_vit_g():
te = transformers.CLIPTextModelWithProjection.from_pretrained(pretrained_model_name_or_path=None, state_dict=state_dict, config=config)
te = te.to(dtype=devices.dtype)
loaded_te = shared.opts.sd_text_encoder
del state_dict
return te
+2
View File
@@ -1074,6 +1074,8 @@ def reload_text_encoder(initial=False):
from modules.model_te import set_t5
shared.log.debug(f'Load module: type=t5 path="{shared.opts.sd_text_encoder}" module="text_encoder_3"')
set_t5(pipe=shared.sd_model, module='text_encoder_3', t5=shared.opts.sd_text_encoder, cache_dir=shared.opts.diffusers_dir)
clear_caches()
apply_balanced_offload(shared.sd_model)
def reload_model_weights(sd_model=None, info=None, op='model', force=False, revision=None):
+7 -8
View File
@@ -90,6 +90,7 @@ def load_text_encoder(repo_id, cls_name, load_config={}, subfolder="text_encoder
# load from local file gguf
if local_file is not None and local_file.lower().endswith('.gguf'):
shared.log.debug(f'Load model: text_encoder="{local_file}" cls={cls_name.__name__} quant="{quant_type}"')
"""
from modules import ggml
ggml.install_gguf()
text_encoder = cls_name.from_pretrained(
@@ -99,17 +100,15 @@ def load_text_encoder(repo_id, cls_name, load_config={}, subfolder="text_encoder
**load_args,
)
text_encoder = model_quant.do_post_load_quant(text_encoder, allow=quant_type is not None)
"""
text_encoder = model_te.load_t5(local_file)
text_encoder = model_quant.do_post_load_quant(text_encoder, allow=quant_type is not None)
# load from local file safetensors
elif local_file is not None and local_file.lower().endswith('.safetensors'):
shared.log.debug(f'Load model: text_encoder="{local_file}" cls={cls_name.__name__} quant="{quant_type}"')
if dtype is not None:
load_args['torch_dtype'] = dtype
text_encoder = cls_name.from_pretrained(
local_file,
cache_dir=shared.opts.hfcache_dir,
**load_args,
**quant_args,
)
from modules import model_te
text_encoder = model_te.load_t5(local_file)
text_encoder = model_quant.do_post_load_quant(text_encoder, allow=quant_type is not None)
# use shared t5 if possible
elif cls_name == transformers.T5EncoderModel and allow_shared:
with open(os.path.join('configs', 'flux', 'text_encoder_2', 'config.json'), encoding='utf8') as f: