mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 17:24:32 +02:00
cleanup multiple model loaders
Signed-off-by: Vladimir Mandic <mandic00@live.com>
This commit is contained in:
@@ -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
|
||||
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user