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}')