diff --git a/CHANGELOG.md b/CHANGELOG.md
index 04770ea8c..9259d02ee 100644
--- a/CHANGELOG.md
+++ b/CHANGELOG.md
@@ -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
diff --git a/cli/test-all-models.py b/cli/test-all-models.py
index baf59f1c8..7115bc8f2 100755
--- a/cli/test-all-models.py
+++ b/cli/test-all-models.py
@@ -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')
diff --git a/modules/api/loras.py b/modules/api/loras.py
index c387bdaa0..8acc8f0cf 100644
--- a/modules/api/loras.py
+++ b/modules/api/loras.py
@@ -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():
diff --git a/modules/api/middleware.py b/modules/api/middleware.py
index ed433170f..d276bd46b 100644
--- a/modules/api/middleware.py
+++ b/modules/api/middleware.py
@@ -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,
diff --git a/modules/memstats.py b/modules/memstats.py
index 71dc2cb63..1e2be27c8 100644
--- a/modules/memstats.py
+++ b/modules/memstats.py
@@ -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
diff --git a/modules/sd_models.py b/modules/sd_models.py
index d0c3423b1..553da62eb 100644
--- a/modules/sd_models.py
+++ b/modules/sd_models.py
@@ -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
diff --git a/modules/shared.py b/modules/shared.py
index 628b83911..55fc32010 100644
--- a/modules/shared.py
+++ b/modules/shared.py
@@ -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("
Optional
", "", 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"),
diff --git a/pipelines/generic.py b/pipelines/generic.py
new file mode 100644
index 000000000..08fe00275
--- /dev/null
+++ b/pipelines/generic.py
@@ -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
diff --git a/pipelines/model_auraflow.py b/pipelines/model_auraflow.py
index 175ba12bf..c6f2ade77 100644
--- a/pipelines/model_auraflow.py
+++ b/pipelines/model_auraflow.py
@@ -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
diff --git a/pipelines/model_bria.py b/pipelines/model_bria.py
index bcd458dc4..900c58a2a 100644
--- a/pipelines/model_bria.py
+++ b/pipelines/model_bria.py
@@ -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,
diff --git a/pipelines/model_chroma.py b/pipelines/model_chroma.py
index a1ca25ffc..102eabb08 100644
--- a/pipelines/model_chroma.py
+++ b/pipelines/model_chroma.py
@@ -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
diff --git a/pipelines/model_cogview.py b/pipelines/model_cogview.py
index 400038dc3..bb3b8eb1a 100644
--- a/pipelines/model_cogview.py
+++ b/pipelines/model_cogview.py
@@ -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
diff --git a/pipelines/model_cosmos.py b/pipelines/model_cosmos.py
index 419dc3f65..1c1c9cee9 100644
--- a/pipelines/model_cosmos.py
+++ b/pipelines/model_cosmos.py
@@ -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,
diff --git a/pipelines/model_flex.py b/pipelines/model_flex.py
index 4a11152ff..f7a285348 100644
--- a/pipelines/model_flex.py
+++ b/pipelines/model_flex.py
@@ -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
diff --git a/pipelines/model_flite.py b/pipelines/model_flite.py
index 9c1426fb4..ad883a564 100644
--- a/pipelines/model_flite.py
+++ b/pipelines/model_flite.py
@@ -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,
diff --git a/pipelines/model_flux.py b/pipelines/model_flux.py
index 5c1ba745b..da4bf70e9 100644
--- a/pipelines/model_flux.py
+++ b/pipelines/model_flux.py
@@ -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:
diff --git a/pipelines/model_hidream.py b/pipelines/model_hidream.py
index 6df9092e2..a5c18d3bc 100644
--- a/pipelines/model_hidream.py
+++ b/pipelines/model_hidream.py
@@ -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
diff --git a/pipelines/model_pixart.py b/pipelines/model_pixart.py
index 254326abc..f8d950ad7 100644
--- a/pipelines/model_pixart.py
+++ b/pipelines/model_pixart.py
@@ -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
diff --git a/pipelines/model_qwen.py b/pipelines/model_qwen.py
index 622b35bc9..4ac91e3cd 100644
--- a/pipelines/model_qwen.py
+++ b/pipelines/model_qwen.py
@@ -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,
diff --git a/pipelines/model_sd3.py b/pipelines/model_sd3.py
index 6130ad81c..177655e54 100644
--- a/pipelines/model_sd3.py
+++ b/pipelines/model_sd3.py
@@ -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
diff --git a/scripts/nudenet/imageguard.py b/scripts/nudenet/imageguard.py
index 55bcfaf86..edb2b7e8a 100644
--- a/scripts/nudenet/imageguard.py
+++ b/scripts/nudenet/imageguard.py
@@ -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}"')
diff --git a/scripts/nudenet/nudenet.py b/scripts/nudenet/nudenet.py
index 67b29b6a6..06ab7db25 100755
--- a/scripts/nudenet/nudenet.py
+++ b/scripts/nudenet/nudenet.py
@@ -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}')