mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 17:24:32 +02:00
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:
@@ -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
@@ -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')
|
||||
|
||||
@@ -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():
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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"),
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
|
||||
|
||||
@@ -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
@@ -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
@@ -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,
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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
|
||||
|
||||
@@ -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}"')
|
||||
|
||||
@@ -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}')
|
||||
|
||||
Reference in New Issue
Block a user