refactor pipeline loaders to generic methods and introduce te_shared_t5 option

Signed-off-by: Vladimir Mandic <mandic00@live.com>
This commit is contained in:
Vladimir Mandic
2025-08-10 20:49:58 -04:00
parent f45e3342e6
commit 2a85c05689
22 changed files with 347 additions and 811 deletions
+6
View File
@@ -85,6 +85,10 @@ And (*as always*) many bugfixes and improvements to existing features!
- new `diffusers_offload_nonblocking` exerimental setting
instructs torch to use non-blocking move operations when possible
- **Features**
- new `T5: Use shared instance of text encoder` option
in *settings -> text encoder*
since a lot of new models use T5 text encoder, this option allows to share
the same instance across all models without duplicate downloads
- **Wan** select which stage to run: *first/second/both* with configurable *boundary ration* when running both stages
in settings -> model options
- prompt parser allow explict `BOS` and `EOS` tokens in prompt
@@ -97,6 +101,8 @@ And (*as always*) many bugfixes and improvements to existing features!
- remove `api-only` cli option
- **API**
- add `/sdapi/v1/checkpoint` POST endpoint to simply load a model
- **Refactor**
- new unified pipeline component loader in `pipelines/generic`
- **Fixes**
- refactor legacy processing loop
- fix settings components mismatch
+56 -46
View File
@@ -20,39 +20,47 @@ models = [
"sdxl-base-v10-vaefix",
"tempest-by-vlad-0.1",
"icbinpXL_v6",
"briaai/BRIA-3.2",
"Freepik/F-Lite",
"Freepik/F-Lite-Texture",
"ostris/Flex.2-preview",
"stabilityai/stable-diffusion-3.5-medium",
"stabilityai/stable-diffusion-3.5-large",
"fal/AuraFlow-v0.3",
"THUDM/CogView3-Plus-3B",
"THUDM/CogView4-6B",
"nvidia/Cosmos-Predict2-2B-Text2Image",
"nvidia/Cosmos-Predict2-14B-Text2Image",
"Qwen/Qwen-Image",
"Qwen/Qwen-Lightning",
"Shitao/OmniGen-v1-diffusers",
"OmniGen2/OmniGen2",
"HiDream-ai/HiDream-I1-Full",
"Kwai-Kolors/Kolors-diffusers",
"vladmandic/chroma-unlocked-v50",
"vladmandic/chroma-unlocked-v50-annealed",
"Alpha-VLLM/Lumina-Next-SFT-diffusers",
"Alpha-VLLM/Lumina-Image-2.0",
"MeissonFlow/Meissonic",
"Efficient-Large-Model/SANA1.5_1.6B_1024px_diffusers",
"Efficient-Large-Model/SANA1.5_4.8B_1024px_diffusers",
"PixArt-alpha/PixArt-XL-2-1024-MS",
"PixArt-alpha/PixArt-Sigma-XL-2-1024-MS",
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
"Wan-AI/Wan2.1-T2V-14B-Diffusers",
"stabilityai/stable-cascade",
]
models_tbd = [
"black-forest-labs/FLUX.1-dev",
"black-forest-labs/FLUX.1-Kontext-dev",
"black-forest-labs/FLUX.1-Krea-dev",
"vladmandic/chroma-unlocked-v50",
"vladmandic/chroma-unlocked-v50-annealed",
"Qwen/Qwen-Image",
"briaai/BRIA-3.2",
"stabilityai/stable-cascade",
"ostris/Flex.2-preview",
"OmniGen2/OmniGen2",
"Freepik/F-Lite",
"Freepik/F-Lite-Texture",
"HiDream-ai/HiDream-I1-Full",
"nvidia/Cosmos-Predict2-2B-Text2Image",
"nvidia/Cosmos-Predict2-14B-Text2Image",
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
"Wan-AI/Wan2.1-T2V-14B-Diffusers",
"Efficient-Large-Model/SANA1.5_1.6B_1024px_diffusers",
"Efficient-Large-Model/SANA1.5_4.8B_1024px_diffusers",
"fal/AuraFlow-v0.3",
"PixArt-alpha/PixArt-XL-2-1024-MS",
"PixArt-alpha/PixArt-Sigma-XL-2-1024-MS",
"Tencent-Hunyuan/HunyuanDiT-v1.2-Diffusers",
"Alpha-VLLM/Lumina-Next-SFT-diffusers",
"Alpha-VLLM/Lumina-Image-2.0",
"Kwai-Kolors/Kolors-diffusers",
"THUDM/CogView4-6B",
"kandinsky-community/kandinsky-3",
"Tencent-Hunyuan/HunyuanDiT-v1.2-Diffusers", # TODO
"kandinsky-community/kandinsky-3", # TODO
]
styles = [
'Fixed Astronaut',
]
styles_tbd = [
'Fixed Bear',
'Fixed Steampunk City',
'Fixed Road sign',
@@ -89,27 +97,29 @@ def generate(): # pylint: disable=redefined-outer-name
model_name = pathvalidate.sanitize_filename(model, replacement_text='_')
log.info(f'model: name="{model}" n={m+1}/{len(models)}')
for s, style in enumerate(styles):
model_name = pathvalidate.sanitize_filename(model, replacement_text='_')
style_name = pathvalidate.sanitize_filename(style, replacement_text='_')
fn = os.path.join(output_folder, f'{model_name}__{style_name}.jpg')
if os.path.exists(fn):
continue
request(f'/sdapi/v1/checkpoint?sd_model_checkpoint={model}', method='POST')
loaded = request('/sdapi/v1/checkpoint', method='GET')
if not (model in loaded.get('checkpoint') or model in loaded.get('title') or model in loaded.get('name')):
log.error(f' model: error="{model}"')
continue
log.info(f' style: name="{style}" n={s+1}/{len(styles)} fn="{fn}"')
t0 = time.time()
data = request('/sdapi/v1/txt2img', { 'styles': [style] })
t1 = time.time()
if 'images' in data and len(data['images']) > 0:
b64 = data['images'][0].split(',',1)[0]
image = Image.open(io.BytesIO(base64.b64decode(b64)))
info = data['info']
log.info(f' image: size={image.size} time={t1-t0:.2f} info="{len(info)}" fn="{fn}"')
image.save(fn)
try:
model_name = pathvalidate.sanitize_filename(model, replacement_text='_')
style_name = pathvalidate.sanitize_filename(style, replacement_text='_')
fn = os.path.join(output_folder, f'{model_name}__{style_name}.jpg')
if os.path.exists(fn):
continue
request(f'/sdapi/v1/checkpoint?sd_model_checkpoint={model}', method='POST')
loaded = request('/sdapi/v1/checkpoint', method='GET')
if not loaded or not (model in loaded.get('checkpoint') or model in loaded.get('title') or model in loaded.get('name')):
log.error(f' model: error="{model}"')
continue
log.info(f' style: name="{style}" n={s+1}/{len(styles)} fn="{fn}"')
t0 = time.time()
data = request('/sdapi/v1/txt2img', { 'styles': [style] })
t1 = time.time()
if 'images' in data and len(data['images']) > 0:
b64 = data['images'][0].split(',',1)[0]
image = Image.open(io.BytesIO(base64.b64decode(b64)))
info = data['info']
log.info(f' image: size={image.size} time={t1-t0:.2f} info="{len(info)}" fn="{fn}"')
image.save(fn)
except Exception as e:
log.error(f' model: error="{model}" style="{style}" exception="{e}"')
if __name__ == "__main__":
log.info('test-all-models')
+3 -3
View File
@@ -7,9 +7,9 @@ def get_lora(lora: str) -> dict:
if lora not in lora_load.available_networks:
raise HTTPException(status_code=404, detail=f"Lora '{lora}' not found")
obj = lora_load.available_networks[lora]
obj.meta = obj.get_metadata()
obj.info = obj.get_info()
obj.desc = obj.get_desc()
# obj.meta = obj.get_metadata()
# obj.info = obj.get_info()
# obj.desc = obj.get_desc()
return obj.__dict__
def get_loras():
+1 -1
View File
@@ -44,7 +44,7 @@ def setup_middleware(app: FastAPI, cmd_opts):
endpoint = req.scope.get('path', 'err')
token = req.cookies.get("access-token") or req.cookies.get("access-token-unsecure")
if (cmd_opts.api_log) and endpoint.startswith('/sdapi'):
if any([endpoint.startswith(x) for x in ignore_endpoints]):
if any([endpoint.startswith(x) for x in ignore_endpoints]): # noqa C419 # pylint: disable=use-a-generator
return res
log.info('API user={user} code={code} {prot}/{ver} {method} {endpoint} {cli} {duration}'.format( # pylint: disable=consider-using-f-string, logging-format-interpolation
user = app.tokens.get(token) if hasattr(app, 'tokens') else None,
+2 -2
View File
@@ -55,7 +55,7 @@ def ram_stats():
ram_total = min(ram_total, get_docker_limit(), get_runpod_limit())
ram['total'] = gb(ram_total)
ram['used'] = gb(res.rss)
ram['free'] = ram['total'] - ram['used']
ram['free'] = round(ram['total'] - ram['used'])
except Exception as e:
ram['total'] = 0
ram['used'] = 0
@@ -97,7 +97,7 @@ def memory_stats():
mem['job'] = shared.state.job
try:
mem['gpu']['swap'] = round(mem['gpu']['active'] - mem['gpu']['used']) if mem['gpu']['active'] > mem['gpu']['used'] else 0
except:
except Exception:
mem['gpu']['swap'] = 0
return mem
+5 -4
View File
@@ -10,7 +10,7 @@ import diffusers.loaders.single_file_utils
import torch
import huggingface_hub as hf
from installer import log
from modules import timer, paths, shared, shared_state, shared_items, modelloader, devices, script_callbacks, sd_vae, sd_unet, errors, sd_models_compile, sd_hijack_accelerate, sd_detect, model_quant, sd_hijack_te
from modules import timer, paths, shared, shared_items, modelloader, devices, script_callbacks, sd_vae, sd_unet, errors, sd_models_compile, sd_hijack_accelerate, sd_detect, model_quant, sd_hijack_te
from modules.memstats import memory_stats
from modules.modeldata import model_data
from modules.sd_checkpoint import CheckpointInfo, select_checkpoint, list_models, checkpoints_list, checkpoint_titles, get_closet_checkpoint_match, model_hash, update_model_hashes, setup_model, write_metadata, read_metadata_from_safetensors # pylint: disable=unused-import
@@ -222,6 +222,8 @@ def move_model(model, device=None, force=False):
pass # ignore model move if quantization is enabled
elif 'already been set to the correct devices' in str(e0):
pass # ignore errors on pre-quant models
elif 'Casting a quantized model to' in str(e0):
pass # ignore errors on quantized models
else:
raise e0
t1 = time.time()
@@ -328,7 +330,7 @@ def load_diffuser_force(model_type, checkpoint_info, diffusers_load_config, op='
allow_post_quant = False
elif model_type in ['Stable Diffusion 3']:
from pipelines.model_sd3 import load_sd3
sd_model = load_sd3(checkpoint_info, cache_dir=shared.opts.diffusers_dir, config=diffusers_load_config.get('config', None))
sd_model = load_sd3(checkpoint_info, diffusers_load_config)
allow_post_quant = False
elif model_type in ['CogView 3']: # forced pipeline
from pipelines.model_cogview import load_cogview3
@@ -467,7 +469,7 @@ def load_diffuser_file(model_type, pipeline, checkpoint_info, diffusers_load_con
diffusers_load_config['config'] = model_config
if model_type.startswith('Stable Diffusion 3'):
from pipelines.model_sd3 import load_sd3
sd_model = load_sd3(checkpoint_info=checkpoint_info, cache_dir=shared.opts.diffusers_dir, config=diffusers_load_config.get('config', None))
sd_model = load_sd3(checkpoint_info, diffusers_load_config)
elif hasattr(pipeline, 'from_single_file'):
diffusers.loaders.single_file_utils.CHECKPOINT_KEY_NAMES["clip"] = "cond_stage_model.transformer.text_model.embeddings.position_embedding.weight" # patch for diffusers==0.28.0
diffusers_load_config['use_safetensors'] = True
@@ -1057,7 +1059,6 @@ def reload_model_weights(sd_model=None, info=None, op='model', force=False, revi
unload_model_weights(op=op)
return None
orig_state = copy.deepcopy(shared.state)
# shared.state = shared_state.State()
shared.state.begin('Load')
if sd_model is None:
sd_model = model_data.sd_model if op == 'model' or op == 'dict' else model_data.sd_refiner
+1
View File
@@ -248,6 +248,7 @@ options_templates.update(options_section(('text_encoder', "Text Encoder"), {
"diffusers_zeros_prompt_pad": OptionInfo(False, "Use zeros for prompt padding", gr.Checkbox),
"te_hijack": OptionInfo(True, "Offload after prompt encode", gr.Checkbox),
"te_optional_sep": OptionInfo("<h2>Optional</h2>", "", gr.HTML),
"te_shared_t5": OptionInfo(False, "T5: Use shared instance of text encoder"),
"te_pooled_embeds": OptionInfo(False, "SDXL: Use weighted pooled embeds"),
"te_complex_human_instruction": OptionInfo(True, "Sana: Use complex human instructions"),
"te_use_mask": OptionInfo(True, "Lumina: Use mask in transformers"),
+132
View File
@@ -0,0 +1,132 @@
import os
import json
import diffusers
import transformers
from modules import shared, devices, sd_models, model_quant
debug = shared.log.trace if os.environ.get('SD_LOAD_DEBUG', None) is not None else lambda *args, **kwargs: None
def load_transformer(repo_id, cls_name, load_config={}, subfolder="transformer"):
load_args, quant_args = model_quant.get_dit_args(load_config, module='Model', device_map=True)
quant_type = model_quant.get_quant_type(quant_args)
local_file = None
if shared.opts.sd_unet is not None and shared.opts.sd_unet != 'Default':
from modules import sd_unet
if shared.opts.sd_unet not in list(sd_unet.unet_dict):
shared.log.error(f'Load module: type=transformer file="{shared.opts.sd_unet}" not found')
elif os.path.exists(sd_unet.unet_dict[shared.opts.sd_unet]):
local_file = sd_unet.unet_dict[shared.opts.sd_unet]
if local_file is not None and local_file.lower().endswith('.gguf'):
shared.log.debug(f'Load model: transformer="{local_file}" cls={cls_name.__name__} quant="{quant_type}" args={load_args}')
from modules import ggml
ggml.install_gguf()
loader = cls_name.from_single_file if hasattr(cls_name, 'from_single_file') else cls_name.from_pretrained
transformer = loader(
local_file,
quantization_config=diffusers.GGUFQuantizationConfig(compute_dtype=devices.dtype),
cache_dir=shared.opts.hfcache_dir,
**load_args,
)
transformer = model_quant.do_post_load_quant(transformer, allow=quant_type is not None)
elif local_file is not None and local_file.lower().endswith('.safetensors'):
shared.log.debug(f'Load model: transformer="{local_file}" cls={cls_name.__name__} quant="{quant_type}" args={load_args}')
loader = cls_name.from_single_file if hasattr(cls_name, 'from_single_file') else cls_name.from_pretrained
transformer = loader(
local_file,
cache_dir=shared.opts.hfcache_dir,
**load_args,
)
transformer = model_quant.do_post_load_quant(transformer, allow=quant_type is not None)
else:
shared.log.debug(f'Load model: transformer="{repo_id}" cls={cls_name.__name__} quant="{quant_type}" args={load_args}')
if subfolder is not None:
load_args['subfolder'] = subfolder
transformer = cls_name.from_pretrained(
repo_id,
cache_dir=shared.opts.hfcache_dir,
**load_args,
**quant_args,
)
if shared.opts.diffusers_offload_mode != 'none' and transformer is not None:
sd_models.move_model(transformer, devices.cpu)
return transformer
def load_text_encoder(repo_id, cls_name, load_config={}, subfolder="text_encoder"):
load_args, quant_args = model_quant.get_dit_args(load_config, module='TE', device_map=True)
quant_type = model_quant.get_quant_type(quant_args)
text_encoder = None
# load from local file if specified
local_file = None
if shared.opts.sd_text_encoder is not None and shared.opts.sd_text_encoder != 'Default':
from modules import model_te
if shared.opts.sd_text_encoder not in list(model_te.te_dict):
shared.log.error(f'Load module: type=te file="{shared.opts.sd_text_encoder}" not found')
elif os.path.exists(model_te.te_dict[shared.opts.sd_text_encoder]):
local_file = model_te.te_dict[shared.opts.sd_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(
gguf_file=local_file,
quantization_config=diffusers.GGUFQuantizationConfig(compute_dtype=devices.dtype),
cache_dir=shared.opts.hfcache_dir,
**load_args,
)
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}"')
text_encoder = cls_name.from_pretrained(
local_file,
cache_dir=shared.opts.hfcache_dir,
**load_args,
)
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:
with open(os.path.join('configs', 'flux', 'text_encoder_2', 'config.json'), encoding='utf8') as f:
load_args['config'] = transformers.T5Config(**json.load(f))
if model_quant.check_nunchaku('TE'):
import nunchaku
repo_id = 'nunchaku-tech/nunchaku-t5/awq-int4-flux.1-t5xxl.safetensors'
cls_name = nunchaku.NunchakuT5EncoderModel
shared.log.debug(f'Load model: text_encoder="{repo_id}" cls={cls_name.__name__} quant="SVDQuant"')
text_encoder = nunchaku.NunchakuT5EncoderModel.from_pretrained(
repo_id,
torch_dtype=devices.dtype,
)
text_encoder.quantization_method = 'SVDQuant'
elif shared.opts.te_shared_t5:
repo_id = 'Disty0/t5-xxl'
shared.log.debug(f'Load model: text_encoder="{repo_id}" cls={cls_name.__name__} quant="{quant_type}" shared={shared.opts.te_shared_t5}')
text_encoder = cls_name.from_pretrained(
repo_id,
cache_dir=shared.opts.hfcache_dir,
**load_args,
**quant_args,
)
# load from repo
if text_encoder is None:
shared.log.debug(f'Load model: text_encoder="{repo_id}" cls={cls_name.__name__} quant="{quant_type}" shared={shared.opts.te_shared_t5}')
if subfolder is not None:
load_args['subfolder'] = subfolder
text_encoder = cls_name.from_pretrained(
repo_id,
cache_dir=shared.opts.hfcache_dir,
**load_args,
**quant_args,
)
if shared.opts.diffusers_offload_mode != 'none' and text_encoder is not None:
sd_models.move_model(text_encoder, devices.cpu)
return text_encoder
+18 -11
View File
@@ -1,21 +1,28 @@
import os
import torch
import transformers
import diffusers
from modules import shared, sd_models, devices
debug = shared.log.trace if os.environ.get('SD_LOAD_DEBUG', None) is not None else lambda *args, **kwargs: None
from modules import shared, sd_models, devices, model_quant
from pipelines import generic
def load_auraflow(checkpoint_info, diffusers_load_config={}):
repo_id = sd_models.path_to_repo(checkpoint_info)
if 'torch_dtype' not in diffusers_load_config:
diffusers_load_config['torch_dtype'] = torch.float16
debug(f'Load model: type=AuraFlow repo="{repo_id}" config={diffusers_load_config}')
sd_models.hf_auth_check(checkpoint_info)
load_args, _quant_args = model_quant.get_dit_args(diffusers_load_config, allow_quant=False)
shared.log.debug(f'Load model: type=AuraFlow repo="{repo_id}" config={diffusers_load_config} offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype} args={load_args}')
transformer = generic.load_transformer(repo_id, cls_name=diffusers.AuraFlowTransformer2DModel, load_config=diffusers_load_config)
text_encoder = generic.load_text_encoder(repo_id, cls_name=transformers.T5EncoderModel, load_config=diffusers_load_config)
pipe = diffusers.AuraFlowPipeline.from_pretrained(
repo_id,
cache_dir = shared.opts.diffusers_dir,
**diffusers_load_config,
transformer=transformer,
text_encoder=text_encoder,
cache_dir=shared.opts.diffusers_dir,
**load_args,
)
del text_encoder
del transformer
devices.torch_gc(force=True, reason='load')
return pipe
+12 -59
View File
@@ -1,73 +1,26 @@
import os
import sys
import transformers
import diffusers
from modules import shared, devices, sd_models, model_quant, sd_hijack_te
def load_transformer(repo_id, diffusers_load_config={}):
load_args, quant_args = model_quant.get_dit_args(diffusers_load_config, module='Model', device_map=True)
fn = None
if shared.opts.sd_unet is not None and shared.opts.sd_unet != 'Default':
from modules import sd_unet
if shared.opts.sd_unet not in list(sd_unet.unet_dict):
shared.log.error(f'Load module: type=Transformer not found: {shared.opts.sd_unet}')
return None
fn = sd_unet.unet_dict[shared.opts.sd_unet] if os.path.exists(sd_unet.unet_dict[shared.opts.sd_unet]) else None
from pipelines.bria.transformer_bria import BriaTransformer2DModel
if fn is not None and 'gguf' in fn.lower():
shared.log.error('Load model: type=Bria format="gguf" unsupported')
transformer = None
elif fn is not None and 'safetensors' in fn.lower():
shared.log.debug(f'Load model: type=Bria transformer="{fn}" quant="{model_quant.get_quant(repo_id)}" args={load_args}')
transformer = BriaTransformer2DModel.from_single_file(
fn,
cache_dir=shared.opts.hfcache_dir,
**load_args,
)
else:
shared.log.debug(f'Load model: type=Bria transformer="{repo_id}" quant="{model_quant.get_quant_type(quant_args)}" args={load_args}')
transformer = BriaTransformer2DModel.from_pretrained(
repo_id,
subfolder="transformer",
cache_dir=shared.opts.hfcache_dir,
**load_args,
**quant_args,
)
if shared.opts.diffusers_offload_mode != 'none' and transformer is not None:
sd_models.move_model(transformer, devices.cpu)
return transformer
def load_text_encoder(repo_id, diffusers_load_config={}):
load_args, quant_args = model_quant.get_dit_args(diffusers_load_config, module='TE', device_map=True)
shared.log.debug(f'Load model: type=Bria te="{repo_id}" quant="{model_quant.get_quant_type(quant_args)}" args={load_args}')
text_encoder = transformers.T5EncoderModel.from_pretrained(
repo_id,
subfolder="text_encoder",
cache_dir=shared.opts.hfcache_dir,
**load_args,
**quant_args,
)
if shared.opts.diffusers_offload_mode != 'none' and text_encoder is not None:
sd_models.move_model(text_encoder, devices.cpu)
return text_encoder
from pipelines import generic
def load_bria(checkpoint_info, diffusers_load_config={}):
repo_id = sd_models.path_to_repo(checkpoint_info)
sd_models.hf_auth_check(checkpoint_info)
transformer = load_transformer(repo_id, diffusers_load_config)
text_encoder = load_text_encoder(repo_id, diffusers_load_config)
load_args, _quant_args = model_quant.get_dit_args(diffusers_load_config, module='Model')
shared.log.debug(f'Load model: type=Bria model="{checkpoint_info.name}" repo="{repo_id}" offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype} args={load_args}')
from pipelines.bria.bria_pipeline import BriaPipeline
sys.path.append(os.path.join(os.path.dirname(__file__), 'bria'))
from pipelines.bria.bria_pipeline import BriaPipeline
from pipelines.bria.transformer_bria import BriaTransformer2DModel
diffusers.BriaPipeline = BriaPipeline
diffusers.BriaTransformer2DModel = BriaTransformer2DModel
load_args, _quant_args = model_quant.get_dit_args(diffusers_load_config, allow_quant=False)
shared.log.debug(f'Load model: type=Bria repo="{repo_id}" config={diffusers_load_config} offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype} args={load_args}')
transformer = generic.load_transformer(repo_id, cls_name=BriaTransformer2DModel, load_config=diffusers_load_config)
text_encoder = generic.load_text_encoder(repo_id, cls_name=transformers.T5EncoderModel, load_config=diffusers_load_config)
pipe = BriaPipeline.from_pretrained(
repo_id,
+17 -264
View File
@@ -1,281 +1,34 @@
import os
import json
import torch
import diffusers
import transformers
from safetensors.torch import load_file
from huggingface_hub import hf_hub_download
from modules import shared, errors, devices, sd_models, sd_unet, model_te, model_quant, sd_hijack_te
from modules import shared, devices, sd_models, model_quant
from pipelines import generic
debug = shared.log.trace if os.environ.get('SD_LOAD_DEBUG', None) is not None else lambda *args, **kwargs: None
def load_chroma_quanto(checkpoint_info):
transformer, text_encoder = None, None
quanto = model_quant.load_quanto('Load model: type=Chroma')
if isinstance(checkpoint_info, str):
repo_path = checkpoint_info
else:
repo_path = checkpoint_info.path
try:
quantization_map = os.path.join(repo_path, "transformer", "quantization_map.json")
debug(f'Load model: type=Chroma quantization map="{quantization_map}" repo="{checkpoint_info.name}" component="transformer"')
if not os.path.exists(quantization_map):
repo_id = sd_models.path_to_repo(checkpoint_info)
quantization_map = hf_hub_download(repo_id, subfolder='transformer', filename='quantization_map.json', cache_dir=shared.opts.diffusers_dir)
with open(quantization_map, "r", encoding='utf8') as f:
quantization_map = json.load(f)
state_dict = load_file(os.path.join(repo_path, "transformer", "diffusion_pytorch_model.safetensors"))
dtype = state_dict['context_embedder.bias'].dtype
with torch.device("meta"):
transformer = diffusers.ChromaTransformer2DModel.from_config(os.path.join(repo_path, "transformer", "config.json")).to(dtype=dtype)
quanto.requantize(transformer, state_dict, quantization_map, device=torch.device("cpu"))
transformer_dtype = transformer.dtype
if transformer_dtype != devices.dtype:
try:
transformer = transformer.to(dtype=devices.dtype)
except Exception:
shared.log.error(f"Load model: type=Chroma Failed to cast transformer to {devices.dtype}, set dtype to {transformer_dtype}")
except Exception as e:
shared.log.error(f"Load model: type=Chroma failed to load Quanto transformer: {e}")
if debug:
errors.display(e, 'Chroma Quanto:')
try:
quantization_map = os.path.join(repo_path, "text_encoder", "quantization_map.json")
debug(f'Load model: type=Chroma quantization map="{quantization_map}" repo="{checkpoint_info.name}" component="text_encoder"')
if not os.path.exists(quantization_map):
repo_id = sd_models.path_to_repo(checkpoint_info)
quantization_map = hf_hub_download(repo_id, subfolder='text_encoder', filename='quantization_map.json', cache_dir=shared.opts.diffusers_dir)
with open(quantization_map, "r", encoding='utf8') as f:
quantization_map = json.load(f)
with open(os.path.join(repo_path, "text_encoder", "config.json"), encoding='utf8') as f:
t5_config = transformers.T5Config(**json.load(f))
state_dict = load_file(os.path.join(repo_path, "text_encoder", "model.safetensors"))
dtype = state_dict['encoder.block.0.layer.0.SelfAttention.relative_attention_bias.weight'].dtype
with torch.device("meta"):
text_encoder = transformers.T5EncoderModel(t5_config).to(dtype=dtype)
quanto.requantize(text_encoder, state_dict, quantization_map, device=torch.device("cpu"))
text_encoder_dtype = text_encoder.dtype
if text_encoder_dtype != devices.dtype:
try:
text_encoder = text_encoder.to(dtype=devices.dtype)
except Exception:
shared.log.error(f"Load model: type=Chroma Failed to cast text encoder to {devices.dtype}, set dtype to {text_encoder_dtype}")
except Exception as e:
shared.log.error(f"Load model: type=Chroma failed to load Quanto text encoder: {e}")
if debug:
errors.display(e, 'Chroma Quanto:')
return transformer, text_encoder
def load_chroma_bnb(checkpoint_info, diffusers_load_config): # pylint: disable=unused-argument
transformer, text_encoder = None, None
if isinstance(checkpoint_info, str):
repo_path = checkpoint_info
else:
repo_path = checkpoint_info.path
model_quant.load_bnb('Load model: type=Chroma')
quant = model_quant.get_quant(repo_path)
try:
# we ignore the distilled guidance layer because it degrades quality too much
# see: https://github.com/huggingface/diffusers/pull/11698#issuecomment-2969717180 for more details
if quant == 'fp8':
quantization_config = transformers.BitsAndBytesConfig(load_in_8bit=True, llm_int8_skip_modules=["distilled_guidance_layer"], bnb_4bit_compute_dtype=devices.dtype)
debug(f'Quantization: {quantization_config}')
transformer = diffusers.ChromaTransformer2DModel.from_single_file(repo_path, **diffusers_load_config, quantization_config=quantization_config)
elif quant == 'fp4':
quantization_config = transformers.BitsAndBytesConfig(load_in_4bit=True, llm_int8_skip_modules=["distilled_guidance_layer"], bnb_4bit_compute_dtype=devices.dtype, bnb_4bit_quant_type= 'fp4')
debug(f'Quantization: {quantization_config}')
transformer = diffusers.ChromaTransformer2DModel.from_single_file(repo_path, **diffusers_load_config, quantization_config=quantization_config)
elif quant == 'nf4':
quantization_config = transformers.BitsAndBytesConfig(load_in_4bit=True, llm_int8_skip_modules=["distilled_guidance_layer"], bnb_4bit_compute_dtype=devices.dtype, bnb_4bit_quant_type= 'nf4')
debug(f'Quantization: {quantization_config}')
transformer = diffusers.ChromaTransformer2DModel.from_single_file(repo_path, **diffusers_load_config, quantization_config=quantization_config)
else:
transformer = diffusers.ChromaTransformer2DModel.from_single_file(repo_path, **diffusers_load_config)
except Exception as e:
shared.log.error(f"Load model: type=Chroma failed to load BnB transformer: {e}")
transformer, text_encoder = None, None
if debug:
errors.display(e, 'Chroma:')
return transformer, text_encoder
def load_quants(kwargs, repo_id, cache_dir, allow_quant): # pylint: disable=unused-argument
try:
diffusers_load_config = {
"torch_dtype": devices.dtype,
"cache_dir": cache_dir,
}
if 'transformer' not in kwargs and model_quant.check_nunchaku('Model'):
shared.log.error(f'Load module: quant=Nunchaku module=transformer repo="{repo_id}" unsupported')
if 'transformer' not in kwargs and model_quant.check_quant('Model'):
load_args, quant_args = model_quant.get_dit_args(diffusers_load_config, module='Model', device_map=True, modules_to_not_convert=["distilled_guidance_layer"])
kwargs['transformer'] = diffusers.ChromaTransformer2DModel.from_pretrained(repo_id, subfolder="transformer", **load_args, **quant_args)
if 'text_encoder' not in kwargs and model_quant.check_nunchaku('TE'):
import nunchaku
nunchaku_precision = nunchaku.utils.get_precision()
nunchaku_repo = 'mit-han-lab/nunchaku-t5/awq-int4-flux.1-t5xxl.safetensors'
shared.log.debug(f'Load module: quant=Nunchaku module=t5 repo="{nunchaku_repo}" precision={nunchaku_precision}')
kwargs['text_encoder'] = nunchaku.NunchakuT5EncoderModel.from_pretrained(nunchaku_repo, torch_dtype=devices.dtype)
if 'text_encoder' not in kwargs and model_quant.check_quant('TE'):
load_args, quant_args = model_quant.get_dit_args(diffusers_load_config, module='TE', device_map=True)
kwargs['text_encoder'] = transformers.T5EncoderModel.from_pretrained(repo_id, subfolder="text_encoder", **load_args, **quant_args)
except Exception as e:
shared.log.error(f'Quantization: {e}')
errors.display(e, 'Quantization:')
return kwargs
def load_transformer(file_path): # triggered by opts.sd_unet change
if file_path is None or not os.path.exists(file_path):
return None
transformer = None
quant = model_quant.get_quant(file_path)
diffusers_load_config = {
"torch_dtype": devices.dtype,
"cache_dir": shared.opts.hfcache_dir,
}
if quant is not None and quant != 'none':
shared.log.info(f'Load module: type=UNet/Transformer file="{file_path}" offload={shared.opts.diffusers_offload_mode} prequant={quant} dtype={devices.dtype}')
if 'gguf' in file_path.lower():
from modules import ggml
_transformer = ggml.load_gguf(file_path, cls=diffusers.ChromaTransformer2DModel, compute_dtype=devices.dtype)
if _transformer is not None:
transformer = _transformer
elif quant in {'qint8', 'qint4'}:
_transformer, _text_encoder = load_chroma_quanto(file_path)
if _transformer is not None:
transformer = _transformer
elif quant in {'fp8', 'fp4', 'nf4'}:
_transformer, _text_encoder = load_chroma_bnb(file_path, diffusers_load_config)
if _transformer is not None:
transformer = _transformer
else:
load_args, quant_args = model_quant.get_dit_args(diffusers_load_config, module='Model', device_map=True, modules_to_not_convert=["distilled_guidance_layer"])
shared.log.debug(f'Load model: type=Chroma transformer file="{file_path}" offload={shared.opts.diffusers_offload_mode} args={load_args}')
transformer = diffusers.ChromaTransformer2DModel.from_single_file(file_path, **load_args, **quant_args)
if transformer is None:
shared.log.error('Failed to load UNet model')
shared.opts.sd_unet = 'Default'
return transformer
def load_chroma(checkpoint_info, diffusers_load_config): # triggered by opts.sd_checkpoint change
fn = checkpoint_info.path
def load_chroma(checkpoint_info, diffusers_load_config={}):
repo_id = sd_models.path_to_repo(checkpoint_info)
sd_models.hf_auth_check(checkpoint_info)
allow_post_quant = False
prequantized = model_quant.get_quant(checkpoint_info.path)
shared.log.debug(f'Load model: type=Chroma model="{checkpoint_info.name}" repo={repo_id or "none"} unet="{shared.opts.sd_unet}" te="{shared.opts.sd_text_encoder}" vae="{shared.opts.sd_vae}" quant={prequantized} offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype}')
debug(f'Load model: type=Chroma config={diffusers_load_config}')
load_args, _quant_args = model_quant.get_dit_args(diffusers_load_config, allow_quant=False)
shared.log.debug(f'Load model: type=Chroma repo="{repo_id}" config={diffusers_load_config} offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype} args={load_args}')
transformer = None
text_encoder = None
vae = None
transformer = generic.load_transformer(repo_id, cls_name=diffusers.ChromaTransformer2DModel, load_config=diffusers_load_config)
text_encoder = generic.load_text_encoder(repo_id, cls_name=transformers.T5EncoderModel, load_config=diffusers_load_config)
# unload current model
sd_models.unload_model_weights()
shared.sd_model = None
devices.torch_gc(force=True, reason='load')
pipe = diffusers.AuraFlowPipeline.from_pretrained(
repo_id,
transformer=transformer,
text_encoder=text_encoder,
cache_dir=shared.opts.diffusers_dir,
**load_args,
)
if shared.opts.teacache_enabled:
from modules import teacache
shared.log.debug(f'Transformers cache: type=teacache patch=forward cls={diffusers.ChromaTransformer2DModel.__name__}')
diffusers.ChromaTransformer2DModel.forward = teacache.teacache_chroma_forward # patch must be done before transformer is loaded
# load overrides if any
if shared.opts.sd_unet != 'Default':
try:
debug(f'Load model: type=Chroma unet="{shared.opts.sd_unet}"')
transformer = load_transformer(sd_unet.unet_dict[shared.opts.sd_unet])
if transformer is None:
shared.opts.sd_unet = 'Default'
sd_unet.failed_unet.append(shared.opts.sd_unet)
except Exception as e:
shared.log.error(f"Load model: type=Chroma failed to load UNet: {e}")
shared.opts.sd_unet = 'Default'
if debug:
errors.display(e, 'Chroma UNet:')
if shared.opts.sd_text_encoder != 'Default':
try:
debug(f'Load model: type=Chroma te="{shared.opts.sd_text_encoder}"')
from modules.model_te import load_t5
text_encoder = load_t5(name=shared.opts.sd_text_encoder, cache_dir=shared.opts.diffusers_dir)
except Exception as e:
shared.log.error(f"Load model: type=Chroma failed to load T5: {e}")
shared.opts.sd_text_encoder = 'Default'
if debug:
errors.display(e, 'Chroma T5:')
if shared.opts.sd_vae != 'Default' and shared.opts.sd_vae != 'Automatic':
try:
debug(f'Load model: type=Chroma vae="{shared.opts.sd_vae}"')
from modules import sd_vae
# vae = sd_vae.load_vae_diffusers(None, sd_vae.vae_dict[shared.opts.sd_vae], 'override')
vae_file = sd_vae.vae_dict[shared.opts.sd_vae]
if os.path.exists(vae_file):
vae_config = os.path.join('configs', 'chroma', 'vae', 'config.json')
vae = diffusers.AutoencoderKL.from_single_file(vae_file, config=vae_config, **diffusers_load_config)
except Exception as e:
shared.log.error(f"Load model: type=Chroma failed to load VAE: {e}")
shared.opts.sd_vae = 'Default'
if debug:
errors.display(e, 'Chroma VAE:')
# initialize pipeline with pre-loaded components
kwargs = {}
if transformer is not None:
kwargs['transformer'] = transformer
sd_unet.loaded_unet = shared.opts.sd_unet
if text_encoder is not None:
kwargs['text_encoder'] = text_encoder
model_te.loaded_te = shared.opts.sd_text_encoder
if vae is not None:
kwargs['vae'] = vae
cls = diffusers.ChromaPipeline
shared.log.debug(f'Load model: type=Chroma cls={cls.__name__} preloaded={list(kwargs)} revision={diffusers_load_config.get("revision", None)}')
for c in kwargs:
if getattr(kwargs[c], 'quantization_method', None) is not None or getattr(kwargs[c], 'gguf', None) is not None:
shared.log.debug(f'Load model: type=Chroma component={c} dtype={kwargs[c].dtype} quant={getattr(kwargs[c], "quantization_method", None) or getattr(kwargs[c], "gguf", None)}')
if kwargs[c].dtype == torch.float32 and devices.dtype != torch.float32:
try:
kwargs[c] = kwargs[c].to(dtype=devices.dtype)
shared.log.warning(f'Load model: type=Chroma component={c} dtype={kwargs[c].dtype} cast dtype={devices.dtype} recast')
except Exception:
pass
allow_quant = 'gguf' not in (sd_unet.loaded_unet or '') and (prequantized is None or prequantized == 'none')
if (fn is None) or (not os.path.exists(fn) or os.path.isdir(fn)):
kwargs = load_quants(kwargs, repo_id, cache_dir=shared.opts.diffusers_dir, allow_quant=allow_quant)
if fn.endswith('.safetensors') and os.path.isfile(fn):
pipe = cls.from_single_file(fn, cache_dir=shared.opts.diffusers_dir, **kwargs, **diffusers_load_config)
allow_post_quant = True
else:
pipe = cls.from_pretrained(repo_id, cache_dir=shared.opts.diffusers_dir, **kwargs, **diffusers_load_config)
if shared.opts.teacache_enabled and model_quant.check_nunchaku('Model'):
from nunchaku.caching.diffusers_adapters import apply_cache_on_pipe
apply_cache_on_pipe(pipe, residual_diff_threshold=0.12)
# register autopipline
diffusers.pipelines.auto_pipeline.AUTO_TEXT2IMAGE_PIPELINES_MAPPING["chroma"] = diffusers.ChromaPipeline
diffusers.pipelines.auto_pipeline.AUTO_IMAGE2IMAGE_PIPELINES_MAPPING["chroma"] = diffusers.ChromaImg2ImgPipeline
# TODO model load: add ChromaControlPipeline, ChromaInpaintPipeline
# Chroma will support inpainting *after* its training has finished: https://huggingface.co/lodestones/Chroma/discussions/28#6826dd2ed86f53ff983add5c
# release memory
transformer = None
text_encoder = None
vae = None
for k in kwargs.keys():
kwargs[k] = None
sd_hijack_te.init_hijack(pipe)
del text_encoder
del transformer
devices.torch_gc(force=True, reason='load')
return pipe, allow_post_quant
return pipe
+19 -47
View File
@@ -1,34 +1,19 @@
import transformers
import diffusers
from modules import shared, devices, sd_models, model_quant, modelloader
from modules import shared, devices, sd_models, model_quant
from pipelines import generic
def load_cogview3(checkpoint_info, diffusers_load_config={}):
modelloader.hf_login()
repo_id = sd_models.path_to_repo(checkpoint_info)
load_args, quant_args = model_quant.get_dit_args(diffusers_load_config, module='Model')
shared.log.debug(f'Load model: type=CogView3 transformer="{repo_id}" quant="{model_quant.get_quant_type(quant_args)}" args={load_args}')
transformer = diffusers.CogView3PlusTransformer2DModel.from_pretrained(
repo_id,
subfolder="transformer",
cache_dir=shared.opts.diffusers_dir,
**load_args,
**quant_args,
)
load_args, quant_args = model_quant.get_dit_args(diffusers_load_config, module='TE', device_map=True)
shared.log.debug(f'Load model: type=CogView3 te="{repo_id}" quant="{model_quant.get_quant_type(quant_args)}" args={load_args}')
text_encoder = transformers.T5EncoderModel.from_pretrained(
repo_id,
subfolder="text_encoder",
cache_dir=shared.opts.diffusers_dir,
**diffusers_load_config,
**quant_args,
)
sd_models.hf_auth_check(checkpoint_info)
load_args, _quant_args = model_quant.get_dit_args(diffusers_load_config, allow_quant=False)
shared.log.debug(f'Load model: type=CogView3 model="{checkpoint_info.name}" repo="{repo_id}" offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype} args={load_args}')
shared.log.debug(f'Load model: type=CogView3 repo="{repo_id}" config={diffusers_load_config} offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype} args={load_args}')
transformer = generic.load_transformer(repo_id, cls_name=diffusers.CogView3PlusTransformer2DModel, load_config=diffusers_load_config, subfolder="transformer")
text_encoder = generic.load_text_encoder(repo_id, cls_name=transformers.T5EncoderModel, load_config=diffusers_load_config, subfolder="text_encoder")
pipe = diffusers.CogView3PlusPipeline.from_pretrained(
repo_id,
text_encoder=text_encoder,
@@ -36,43 +21,30 @@ def load_cogview3(checkpoint_info, diffusers_load_config={}):
cache_dir=shared.opts.diffusers_dir,
**load_args,
)
del transformer
del text_encoder
devices.torch_gc()
return pipe
def load_cogview4(checkpoint_info, diffusers_load_config={}):
modelloader.hf_login()
repo_id = sd_models.path_to_repo(checkpoint_info)
load_args, quant_args = model_quant.get_dit_args(diffusers_load_config, module='Model')
shared.log.debug(f'Load model: type=CogView4 transformer="{repo_id}" quant="{model_quant.get_quant_type(quant_args)}" args={load_args}')
transformer = diffusers.CogView4Transformer2DModel.from_pretrained(
repo_id,
subfolder="transformer",
cache_dir=shared.opts.diffusers_dir,
**diffusers_load_config,
**quant_args,
)
load_args, quant_args = model_quant.get_dit_args(diffusers_load_config, module='TE', device_map=True)
shared.log.debug(f'Load model: type=CogView4 te="{repo_id}" quant="{model_quant.get_quant_type(quant_args)}" args={load_args}')
text_encoder = transformers.AutoModelForCausalLM.from_pretrained( # TODO model load: cogview4 balanced offload does not work for GlmModel
repo_id,
subfolder="text_encoder",
cache_dir=shared.opts.diffusers_dir,
**load_args,
# **quant_args,
)
sd_models.hf_auth_check(checkpoint_info)
load_args, _quant_args = model_quant.get_dit_args(diffusers_load_config, allow_quant=False)
shared.log.debug(f'Load model: type=CogView4 model="{checkpoint_info.name}" repo="{repo_id}" offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype} args={load_args}')
pipe = diffusers.CogView4Pipeline.from_pretrained(
shared.log.debug(f'Load model: type=CogView4 repo="{repo_id}" config={diffusers_load_config} offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype} args={load_args}')
transformer = generic.load_transformer(repo_id, cls_name=diffusers.CogView4Transformer2DModel, load_config=diffusers_load_config, subfolder="transformer")
text_encoder = generic.load_text_encoder(repo_id, cls_name=transformers.T5EncoderModel, load_config=diffusers_load_config, subfolder="text_encoder")
pipe = diffusers.CogView3PlusPipeline.from_pretrained(
repo_id,
text_encoder=text_encoder,
transformer=transformer,
cache_dir=shared.opts.diffusers_dir,
**load_args,
)
pipe.enable_model_cpu_offload()
del transformer
del text_encoder
devices.torch_gc()
return pipe
+7 -54
View File
@@ -1,68 +1,21 @@
import os
import transformers
import diffusers
from modules import shared, devices, sd_models, model_quant, sd_hijack_te
def load_transformer(repo_id, diffusers_load_config={}):
load_args, quant_args = model_quant.get_dit_args(diffusers_load_config, module='Model', device_map=True)
fn = None
if shared.opts.sd_unet is not None and shared.opts.sd_unet != 'Default':
from modules import sd_unet
if shared.opts.sd_unet not in list(sd_unet.unet_dict):
shared.log.error(f'Load module: type=Transformer not found: {shared.opts.sd_unet}')
return None
fn = sd_unet.unet_dict[shared.opts.sd_unet] if os.path.exists(sd_unet.unet_dict[shared.opts.sd_unet]) else None
if fn is not None and 'gguf' in fn.lower():
shared.log.error('Load model: type=Cosmos format="gguf" unsupported')
transformer = None
elif fn is not None and 'safetensors' in fn.lower():
shared.log.debug(f'Load model: type=Cosmos transformer="{fn}" quant="{model_quant.get_quant(repo_id)}" args={load_args}')
transformer = diffusers.CosmosTransformer3DModel.from_single_file(fn, cache_dir=shared.opts.hfcache_dir, **load_args)
else:
shared.log.debug(f'Load model: type=Cosmos transformer="{repo_id}" quant="{model_quant.get_quant_type(quant_args)}" args={load_args}')
transformer = diffusers.CosmosTransformer3DModel.from_pretrained(
repo_id,
subfolder="transformer",
cache_dir=shared.opts.hfcache_dir,
**load_args,
**quant_args,
)
if shared.opts.diffusers_offload_mode != 'none' and transformer is not None:
sd_models.move_model(transformer, devices.cpu)
return transformer
def load_text_encoder(repo_id, diffusers_load_config={}):
load_args, quant_args = model_quant.get_dit_args(diffusers_load_config, module='TE', device_map=True)
shared.log.debug(f'Load model: type=Cosmos te="{repo_id}" quant="{model_quant.get_quant_type(quant_args)}" args={load_args}')
text_encoder = transformers.T5EncoderModel.from_pretrained(
repo_id,
subfolder="text_encoder",
cache_dir=shared.opts.hfcache_dir,
**load_args,
**quant_args,
)
if shared.opts.diffusers_offload_mode != 'none' and text_encoder is not None:
sd_models.move_model(text_encoder, devices.cpu)
return text_encoder
from pipelines import generic
def load_cosmos_t2i(checkpoint_info, diffusers_load_config={}):
repo_id = sd_models.path_to_repo(checkpoint_info)
sd_models.hf_auth_check(checkpoint_info)
transformer = load_transformer(repo_id, diffusers_load_config)
text_encoder = load_text_encoder(repo_id, diffusers_load_config)
load_args, _quant_args = model_quant.get_dit_args(diffusers_load_config, allow_quant=False)
shared.log.debug(f'Load model: type=Cosmos repo="{repo_id}" config={diffusers_load_config} offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype} args={load_args}')
transformer = generic.load_transformer(repo_id, cls_name=diffusers.CosmosTransformer3DModel, load_config=diffusers_load_config, subfolder="transformer")
text_encoder = generic.load_text_encoder(repo_id, cls_name=transformers.T5EncoderModel, load_config=diffusers_load_config, subfolder="text_encoder")
safety_checker = Fake_safety_checker()
load_args, _quant_args = model_quant.get_dit_args(diffusers_load_config, module='Model')
shared.log.debug(f'Load model: type=Cosmos model="{checkpoint_info.name}" repo="{repo_id}" offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype} args={load_args}')
cls = diffusers.Cosmos2TextToImagePipeline
pipe = cls.from_pretrained(
pipe = diffusers.Cosmos2TextToImagePipeline.from_pretrained(
repo_id,
transformer=transformer,
text_encoder=text_encoder,
+6 -55
View File
@@ -1,81 +1,32 @@
import os
import transformers
import diffusers
from modules import shared, devices, sd_models, model_quant, sd_hijack_te
def load_transformer(repo_id, diffusers_load_config={}):
load_args, quant_args = model_quant.get_dit_args(diffusers_load_config, module='Model', device_map=True)
fn = None
if shared.opts.sd_unet is not None and shared.opts.sd_unet != 'Default':
from modules import sd_unet
if shared.opts.sd_unet not in list(sd_unet.unet_dict):
shared.log.error(f'Load module: type=Transformer not found: {shared.opts.sd_unet}')
return None
fn = sd_unet.unet_dict[shared.opts.sd_unet] if os.path.exists(sd_unet.unet_dict[shared.opts.sd_unet]) else None
if fn is not None and 'gguf' in fn.lower():
shared.log.error('Load model: type=HiDream format="gguf" unsupported')
transformer = None
from modules import ggml
transformer = ggml.load_gguf(fn, cls=diffusers.HiDreamImageTransformer2DModel, compute_dtype=devices.dtype)
elif fn is not None and 'safetensors' in fn.lower():
shared.log.debug(f'Load model: type=FLEX transformer="{repo_id}" quant="{model_quant.get_quant(repo_id)}" args={load_args}')
transformer = diffusers.FluxTransformer2DModel.from_single_file(fn, cache_dir=shared.opts.hfcache_dir, **load_args)
else:
shared.log.debug(f'Load model: type=FLEX transformer="{repo_id}" quant="{model_quant.get_quant_type(quant_args)}" args={load_args}')
transformer = diffusers.FluxTransformer2DModel.from_pretrained(
repo_id,
subfolder="transformer",
cache_dir=shared.opts.hfcache_dir,
**load_args,
**quant_args,
)
if shared.opts.diffusers_offload_mode != 'none' and transformer is not None:
sd_models.move_model(transformer, devices.cpu)
return transformer
def load_text_encoders(repo_id, diffusers_load_config={}):
load_args, quant_args = model_quant.get_dit_args(diffusers_load_config, module='TE', device_map=True)
shared.log.debug(f'Load model: type=FLEX t5="{repo_id}" quant="{model_quant.get_quant_type(quant_args)}" args={load_args}')
text_encoder_2 = transformers.T5EncoderModel.from_pretrained(
repo_id,
subfolder="text_encoder_2",
cache_dir=shared.opts.hfcache_dir,
**load_args,
**quant_args,
)
if shared.opts.diffusers_offload_mode != 'none' and text_encoder_2 is not None:
sd_models.move_model(text_encoder_2, devices.cpu)
return text_encoder_2
from pipelines import generic
def load_flex(checkpoint_info, diffusers_load_config={}):
repo_id = sd_models.path_to_repo(checkpoint_info)
sd_models.hf_auth_check(checkpoint_info)
transformer = load_transformer(repo_id, diffusers_load_config)
text_encoder_2 = load_text_encoders(repo_id, diffusers_load_config)
load_args, _quant_args = model_quant.get_dit_args(diffusers_load_config, allow_quant=False)
shared.log.debug(f'Load model: type=Flex repo="{repo_id}" config={diffusers_load_config} offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype} args={load_args}')
load_args, _quant_args = model_quant.get_dit_args(diffusers_load_config, module='Model')
shared.log.debug(f'Load model: type=FLEX model="{checkpoint_info.name}" repo="{repo_id}" offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype} args={load_args}')
transformer = generic.load_transformer(repo_id, cls_name=diffusers.FluxTransformer2DModel, load_config=diffusers_load_config)
text_encoder_2 = generic.load_text_encoder(repo_id, cls_name=transformers.T5EncoderModel, load_config=diffusers_load_config, subfolder="text_encoder_2")
from pipelines.flex2 import Flex2Pipeline
pipe = Flex2Pipeline.from_pretrained(
repo_id,
# custom_pipeline=repo_id,
transformer=transformer,
text_encoder_2=text_encoder_2,
cache_dir=shared.opts.diffusers_dir,
**load_args,
)
sd_hijack_te.init_hijack(pipe)
diffusers.pipelines.auto_pipeline.AUTO_TEXT2IMAGE_PIPELINES_MAPPING["flex2"] = Flex2Pipeline
diffusers.pipelines.auto_pipeline.AUTO_IMAGE2IMAGE_PIPELINES_MAPPING["flex2"] = Flex2Pipeline
diffusers.pipelines.auto_pipeline.AUTO_INPAINT_PIPELINES_MAPPING["flex2"] = Flex2Pipeline
sd_hijack_te.init_hijack(pipe)
del text_encoder_2
del transformer
+13 -39
View File
@@ -1,52 +1,26 @@
import sys
import diffusers
import transformers
from modules import shared, devices, sd_models, model_quant, sd_hijack_te
def load_dit(repo_id, diffusers_load_config={}):
load_args, quant_args = model_quant.get_dit_args(diffusers_load_config, module='Model', device_map=True)
shared.log.debug(f'Load model: type=FLite dit="{repo_id}" quant="{model_quant.get_quant_type(quant_args)}" args={load_args}')
import pipelines.f_lite
sys.modules['f_lite'] = pipelines.f_lite
transformer = pipelines.f_lite.DiT.from_pretrained(
repo_id,
subfolder="dit_model",
cache_dir=shared.opts.hfcache_dir,
**load_args,
**quant_args,
)
if shared.opts.diffusers_offload_mode != 'none' and transformer is not None:
sd_models.move_model(transformer, devices.cpu)
return transformer
def load_text_encoder(repo_id, diffusers_load_config={}):
load_args, quant_args = model_quant.get_dit_args(diffusers_load_config, module='TE', device_map=True)
shared.log.debug(f'Load model: type=FLite te="{repo_id}" quant="{model_quant.get_quant_type(quant_args)}" args={load_args}')
text_encoder = transformers.T5EncoderModel.from_pretrained(
repo_id,
subfolder="text_encoder",
cache_dir=shared.opts.hfcache_dir,
**load_args,
**quant_args,
)
if shared.opts.diffusers_offload_mode != 'none' and text_encoder is not None:
sd_models.move_model(text_encoder, devices.cpu)
return text_encoder
from pipelines import generic
def load_flite(checkpoint_info, diffusers_load_config={}):
repo_id = sd_models.path_to_repo(checkpoint_info)
sd_models.hf_auth_check(checkpoint_info)
from pipelines.f_lite import FLitePipeline
dit_model = load_dit(repo_id, diffusers_load_config)
text_encoder = load_text_encoder(repo_id, diffusers_load_config)
load_args, _quant_args = model_quant.get_dit_args(diffusers_load_config, allow_quant=False)
shared.log.debug(f'Load model: type=FLite repo="{repo_id}" config={diffusers_load_config} offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype} args={load_args}')
load_args, _quant_args = model_quant.get_dit_args(diffusers_load_config, module='Model')
shared.log.debug(f'Load model: type=FLite model="{checkpoint_info.name}" repo="{repo_id}" offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype} args={load_args}')
pipe = FLitePipeline.from_pretrained(
repo_id,
from pipelines import f_lite
diffusers.FLitePipeline = f_lite.FLitePipeline
sys.modules['f_lite'] = f_lite
dit_model = generic.load_transformer(repo_id, cls_name=f_lite.DiT, load_config=diffusers_load_config, subfolder="dit_model")
text_encoder = generic.load_text_encoder(repo_id, cls_name=transformers.T5EncoderModel, load_config=diffusers_load_config, subfolder="text_encoder")
pipe = f_lite.FLitePipeline.from_pretrained(
"Freepik/F-Lite", # pr only exists on main repo
revision="refs/pr/8",
dit_model=dit_model,
text_encoder=text_encoder,
+1 -1
View File
@@ -181,7 +181,7 @@ def load_transformer(file_path): # triggered by opts.sd_unet change
_transformer, _text_encoder_2 = load_flux_bnb(file_path, diffusers_load_config)
if _transformer is not None:
transformer = _transformer
elif 'nf4' in quant: # TODO flux: loader for civitai nf4 models
elif 'nf4' in quant:
from pipelines.model_flux_nf4 import load_flux_nf4
_transformer, _text_encoder_2 = load_flux_nf4(file_path, prequantized=True)
if _transformer is not None:
+8 -57
View File
@@ -4,62 +4,12 @@ import diffusers
from modules import shared, devices, sd_models, model_quant, sd_hijack_te
def load_transformer(repo_id, diffusers_load_config={}):
load_args, quant_args = model_quant.get_dit_args(diffusers_load_config, module='Model', device_map=True)
fn = None
if shared.opts.sd_unet is not None and shared.opts.sd_unet != 'Default':
from modules import sd_unet
if shared.opts.sd_unet not in list(sd_unet.unet_dict):
shared.log.error(f'Load module: type=Transformer not found: {shared.opts.sd_unet}')
return None
fn = sd_unet.unet_dict[shared.opts.sd_unet] if os.path.exists(sd_unet.unet_dict[shared.opts.sd_unet]) else None
if fn is not None and 'gguf' in fn.lower():
shared.log.error('Load model: type=HiDream format="gguf" unsupported')
transformer = None
# from modules import ggml
# transformer = ggml.load_gguf(fn, cls=diffusers.HiDreamImageTransformer2DModel, compute_dtype=devices.dtype)
elif fn is not None and 'safetensors' in fn.lower():
shared.log.debug(f'Load model: type=HiDream transformer="{repo_id}" offload={shared.opts.diffusers_offload_mode} quant="{model_quant.get_quant(repo_id)}" args={load_args}')
transformer = diffusers.HiDreamImageTransformer2DModel.from_single_file(fn, cache_dir=shared.opts.hfcache_dir, **load_args, **quant_args)
# elif model_quant.check_nunchaku('Model'):
# shared.log.error(f'Load model: type=HiDream transformer="{repo_id}" quant="Nunchaku" unsupported')
# transformer = None
else:
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,
subfolder="transformer",
cache_dir=shared.opts.hfcache_dir,
**load_args,
**quant_args,
)
if shared.opts.diffusers_offload_mode != 'none' and transformer is not None:
sd_models.move_model(transformer, devices.cpu)
return transformer
def load_text_encoders(repo_id, diffusers_load_config={}):
if repo_id == 'HiDream-ai/HiDream-E1-Full':
repo_id = 'HiDream-ai/HiDream-I1-Full' # use I1 for t5 and llm
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,
subfolder="text_encoder_3",
cache_dir=shared.opts.hfcache_dir,
**load_args,
**quant_args,
)
if shared.opts.diffusers_offload_mode != 'none' and text_encoder_3 is not None:
sd_models.move_model(text_encoder_3, devices.cpu)
def load_llama(repo_id, diffusers_load_config={}):
load_args, quant_args = model_quant.get_dit_args(diffusers_load_config, module='TE', device_map=True)
llama_repo = shared.opts.model_h1_llama_repo if shared.opts.model_h1_llama_repo != 'Default' else 'meta-llama/Meta-Llama-3.1-8B-Instruct'
shared.log.debug(f'Load model: type=HiDream te4="{llama_repo}" quant="{model_quant.get_quant_type(quant_args)}" args={load_args}')
sd_models.hf_auth_check(llama_repo)
text_encoder_4 = transformers.LlamaForCausalLM.from_pretrained(
llama_repo,
output_hidden_states=True,
@@ -75,18 +25,19 @@ def load_text_encoders(repo_id, diffusers_load_config={}):
)
if shared.opts.diffusers_offload_mode != 'none' and text_encoder_4 is not None:
sd_models.move_model(text_encoder_4, devices.cpu)
return text_encoder_3, text_encoder_4, tokenizer_4
return text_encoder_4, tokenizer_4
def load_hidream(checkpoint_info, diffusers_load_config={}):
repo_id = sd_models.path_to_repo(checkpoint_info)
sd_models.hf_auth_check(checkpoint_info)
transformer = load_transformer(repo_id, diffusers_load_config)
text_encoder_3, text_encoder_4, tokenizer_4 = load_text_encoders(repo_id, diffusers_load_config)
load_args, _quant_args = model_quant.get_dit_args(diffusers_load_config, allow_quant=False)
shared.log.debug(f'Load model: type=HiDream repo="{repo_id}" config={diffusers_load_config} offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype} args={load_args}')
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}')
transformer = generic.load_transformer(repo_id, cls_name=diffusers.HiDreamImageTransformer2DModel, load_config=diffusers_load_config, subfolder="transformer")
text_encoder_3 = generic.load_text_encoder(repo_id, cls_name=transformers.T5EncoderModel, load_config=diffusers_load_config, subfolder="text_encoder_3")
text_encoder_4, tokenizer_4 = load_text_encoders(repo_id, diffusers_load_config)
if shared.opts.teacache_enabled:
from modules import teacache
+13 -20
View File
@@ -1,12 +1,14 @@
import transformers
import diffusers
from huggingface_hub import file_exists
from modules import shared, devices, modelloader, sd_models, model_quant
from pipelines import generic
def load_pixart(checkpoint_info, diffusers_load_config={}):
from modules import shared, devices, modelloader, sd_models, model_quant
modelloader.hf_login()
repo_id = sd_models.path_to_repo(checkpoint_info)
sd_models.hf_auth_check(checkpoint_info)
repo_id_tenc = repo_id
repo_id_pipe = repo_id
@@ -15,30 +17,21 @@ def load_pixart(checkpoint_info, diffusers_load_config={}):
if not file_exists(repo_id_pipe, "model_index.json"):
repo_id_pipe = "PixArt-alpha/PixArt-Sigma-XL-2-1024-MS"
load_args, quant_args = model_quant.get_dit_args(diffusers_load_config, module='Model')
transformer = diffusers.PixArtTransformer2DModel.from_pretrained(
repo_id,
subfolder='transformer',
cache_dir=shared.opts.hfcache_dir,
**load_args,
**quant_args,
)
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_tenc,
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)
shared.log.debug(f'Load model: type=AuraFlow repo="{repo_id}" config={diffusers_load_config} offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype} args={load_args}')
transformer = generic.load_transformer(repo_id, cls_name=diffusers.PixArtTransformer2DModel, load_config=diffusers_load_config)
text_encoder = generic.load_text_encoder(repo_id_tenc, cls_name=transformers.T5EncoderModel, load_config=diffusers_load_config)
pipe = diffusers.PixArtSigmaPipeline.from_pretrained(
repo_id_pipe,
cache_dir=shared.opts.diffusers_dir,
transformer=transformer,
text_encoder=text_encoder,
cache_dir=shared.opts.diffusers_dir,
**load_args,
)
del text_encoder
del transformer
devices.torch_gc(force=True, reason='load')
return pipe
+4 -33
View File
@@ -1,48 +1,19 @@
import transformers
import diffusers
from modules import shared, devices, sd_models, model_quant, sd_hijack_te
def load_transformer(repo_id, diffusers_load_config={}):
load_args, quant_args = model_quant.get_dit_args(diffusers_load_config, module='Model', device_map=True)
shared.log.debug(f'Load model: type=Qwen transformer="{repo_id}" quant="{model_quant.get_quant_type(quant_args)}" args={load_args}')
transformer = diffusers.QwenImageTransformer2DModel.from_pretrained(
repo_id,
subfolder="transformer",
cache_dir=shared.opts.hfcache_dir,
**load_args,
**quant_args,
)
if shared.opts.diffusers_offload_mode != 'none' and transformer is not None:
sd_models.move_model(transformer, devices.cpu)
return transformer
def load_text_encoder(repo_id, diffusers_load_config={}):
load_args, quant_args = model_quant.get_dit_args(diffusers_load_config, module='TE', device_map=True)
shared.log.debug(f'Load model: type=Qwen te="{repo_id}" quant="{model_quant.get_quant_type(quant_args)}" args={load_args}')
text_encoder = transformers.Qwen2_5_VLForConditionalGeneration.from_pretrained(
repo_id,
subfolder="text_encoder",
cache_dir=shared.opts.hfcache_dir,
**load_args,
**quant_args,
)
if shared.opts.diffusers_offload_mode != 'none' and text_encoder is not None:
sd_models.move_model(text_encoder, devices.cpu)
return text_encoder
from pipelines import generic
def load_qwen(checkpoint_info, diffusers_load_config={}):
repo_id = sd_models.path_to_repo(checkpoint_info)
sd_models.hf_auth_check(checkpoint_info)
transformer = load_transformer(repo_id, diffusers_load_config)
text_encoder = load_text_encoder(repo_id, diffusers_load_config)
load_args, _quant_args = model_quant.get_dit_args(diffusers_load_config, module='Model')
shared.log.debug(f'Load model: type=Qwen model="{checkpoint_info.name}" repo="{repo_id}" offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype} args={load_args}')
transformer = generic.load_transformer(repo_id, cls_name=diffusers.QwenImageTransformer2DModel, load_config=diffusers_load_config)
text_encoder = generic.load_text_encoder(repo_id, cls_name=transformers.Qwen2_5_VLForConditionalGeneration, load_config=diffusers_load_config)
cls = diffusers.QwenImagePipeline
pipe = cls.from_pretrained(
repo_id,
+21 -113
View File
@@ -1,128 +1,36 @@
import os
import diffusers
import transformers
from modules import shared, devices, errors, sd_models, sd_unet, model_quant, model_tools
from modules import shared, devices, sd_models, model_quant
from pipelines import generic
def load_overrides(kwargs, cache_dir):
if shared.opts.sd_unet != 'Default':
try:
fn = sd_unet.unet_dict[shared.opts.sd_unet]
if fn.endswith('.safetensors'):
kwargs['transformer'] = diffusers.SD3Transformer2DModel.from_single_file(fn, cache_dir=cache_dir, torch_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=safetensors')
elif fn.endswith('.gguf'):
from modules import ggml
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')
except Exception as e:
shared.log.error(f"Load model: type=SD3 failed to load UNet: {e}")
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
if 'vit-l' in shared.opts.sd_text_encoder.lower():
kwargs['text_encoder'] = load_vit_l()
shared.log.debug(f'Load model: type=SD3 variant="vit-l" te="{shared.opts.sd_text_encoder}"')
elif 'vit-g' in shared.opts.sd_text_encoder.lower():
kwargs['text_encoder_2'] = load_vit_g()
shared.log.debug(f'Load model: type=SD3 variant="vit-g" te="{shared.opts.sd_text_encoder}"')
else:
kwargs['text_encoder_3'] = load_t5(name=shared.opts.sd_text_encoder, cache_dir=shared.opts.diffusers_dir)
shared.log.debug(f'Load model: type=SD3 variant="t5" te="{shared.opts.sd_text_encoder}"')
except Exception as e:
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
vae_file = sd_vae.vae_dict[shared.opts.sd_vae]
if os.path.exists(vae_file):
vae_config = os.path.join('configs', 'sd3', 'vae', 'config.json')
kwargs['vae'] = diffusers.AutoencoderKL.from_single_file(vae_file, config=vae_config, cache_dir=cache_dir, torch_dtype=devices.dtype)
shared.log.debug(f'Load model: type=SD3 vae="{shared.opts.sd_vae}"')
except Exception as e:
shared.log.error(f"Load model: type=SD3 failed to load VAE: {e}")
errors.display(e, 'VAE')
shared.opts.sd_vae = 'Default'
return kwargs
def load_quants(kwargs, repo_id, cache_dir):
quant_args = model_quant.create_config(module='Model')
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)
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
def load_missing(kwargs, fn, cache_dir):
keys = model_tools.get_safetensor_keys(fn)
size = os.stat(fn).st_size // 1024 // 1024
if size > 15000:
repo_id = 'stabilityai/stable-diffusion-3.5-large'
else:
repo_id = 'stabilityai/stable-diffusion-3-medium-diffusers'
if 'text_encoder' not in kwargs and 'text_encoder' not in keys:
kwargs['text_encoder'] = transformers.CLIPTextModelWithProjection.from_pretrained(repo_id, subfolder='text_encoder', cache_dir=cache_dir, torch_dtype=devices.dtype)
shared.log.debug(f'Load model: type=SD3 missing=te1 repo="{repo_id}"')
if 'text_encoder_2' not in kwargs and 'text_encoder_2' not in keys:
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:
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}"')
return kwargs
def load_sd3(checkpoint_info, cache_dir=None, config=None):
def load_sd3(checkpoint_info, diffusers_load_config={}):
repo_id = sd_models.path_to_repo(checkpoint_info)
sd_models.hf_auth_check(checkpoint_info)
fn = checkpoint_info.path
kwargs = {}
kwargs = load_overrides(kwargs, cache_dir)
if (fn is None) or (not os.path.exists(fn) or os.path.isdir(fn)):
kwargs = load_quants(kwargs, repo_id, cache_dir)
loader = diffusers.StableDiffusion3Pipeline.from_pretrained
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
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_missing(kwargs, fn, cache_dir)
kwargs['variant'] = 'fp16'
else:
kwargs['variant'] = 'fp16'
shared.log.debug(f'Load model: type=SD3 kwargs={list(kwargs)} repo="{repo_id}"')
load_args, _quant_args = model_quant.get_dit_args(diffusers_load_config, allow_quant=False)
shared.log.debug(f'Load model: type=SD3 repo="{repo_id}" config={diffusers_load_config} offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype} args={load_args}')
transformer = generic.load_transformer(repo_id, cls_name=diffusers.SD3Transformer2DModel, load_config=diffusers_load_config)
# text_encoder = generic.load_text_encoder(repo_id, cls_name=transformers.CLIPTextModelWithProjection, load_config=diffusers_load_config, subfolder="text_encoder")
# text_encoder_2 = generic.load_text_encoder(repo_id, cls_name=transformers.CLIPTextModelWithProjection, load_config=diffusers_load_config, subfolder="text_encoder_2")
if shared.opts.model_sd3_disable_te5:
shared.log.debug('Load model: type=SD3 option="disable-te5"')
kwargs['text_encoder_3'] = None
text_encoder_3 = None
else:
text_encoder_3 = generic.load_text_encoder(repo_id, cls_name=transformers.T5EncoderModel, load_config=diffusers_load_config, subfolder="text_encoder_3")
pipe = loader(
pipe = diffusers.StableDiffusion3Pipeline.from_pretrained(
repo_id,
torch_dtype=devices.dtype,
cache_dir=cache_dir,
config=config,
**kwargs,
transformer=transformer,
# text_encoder=text_encoder,
# text_encoder_2=text_encoder_2,
text_encoder_3=text_encoder_3,
cache_dir=shared.opts.diffusers_dir,
**load_args,
)
del text_encoder_3
del transformer
devices.torch_gc(force=True, reason='load')
return pipe
+1 -1
View File
@@ -106,7 +106,7 @@ def image_guard(image, policy:str=None) -> str:
attn_implementation='flash_attention_2',
torch_dtype=devices.dtype,
device_map="auto",
cache_dir='/mnt/models/huggingface',
cache_dir=shared.opts.hfcache_dir,
)
processor = transformers.AutoProcessor.from_pretrained(repo_id, cache_dir=shared.opts.hfcache_dir)
shared.log.info(f'NudeNet load: model="{repo_id}"')
+1 -1
View File
@@ -61,7 +61,7 @@ class NudeDetector:
self.model_path = model or hf.hf_hub_download(
repo_id='vladmandic/nudenet',
filename='nudenet.onnx',
cache_dir=shared.opts.diffusers_dir,
cache_dir=shared.opts.hfcache_dir,
)
if session is None:
log.info(f'NudeNet load: model="{self.model_path}" providers={providers}')