mirror of
https://github.com/vladmandic/automatic
synced 2026-09-20 01:31:13 +02:00
fix loading custom t5
Signed-off-by: Vladimir Mandic <mandic00@live.com>
This commit is contained in:
+9
-2
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user