from enum import Enum import sys import time import copy import inspect import logging import os import os.path import diffusers import diffusers.loaders.single_file_utils import torch import huggingface_hub as hf from modules.logger import log from modules import timer, paths, shared, modelloader, devices, script_callbacks, sd_vae, sd_unet, errors, sd_models_compile, sd_detect, model_quant, sd_hijack_te, sd_hijack_vae, sd_hijack_accelerate, sd_hijack_safetensors, sd_hijack_transformers, sd_hijack_hfhub, attention from modules.memstats import memory_stats, gpu_stats from modules.shared_helpers import walk_files from modules.modeldata import model_data from modules.sd_checkpoint import CheckpointInfo, select_checkpoint, list_models, checkpoint_titles, get_closest_checkpoint_match, update_model_hashes, write_metadata, checkpoints_list # pylint: disable=unused-import from modules.sd_offload import get_module_names, disable_offload, set_diffuser_offload, apply_balanced_offload, set_accelerate, offload_ondemand, reapply_offload # pylint: disable=unused-import from modules.sd_models_utils import NoWatermark, get_signature, get_call, path_to_repo, apply_function_to_model, read_state_dict, get_state_dict_from_checkpoint # pylint: disable=unused-import model_dir = "Stable-diffusion" model_path = os.path.abspath(os.path.join(paths.models_path, model_dir)) sd_metadata = None sd_metadata_pending = 0 sd_metadata_timer = 0 loaded_te = None # tracks the text-encoder selection currently loaded, to detect sd_text_encoder changes debug_move = log.trace if os.environ.get('SD_MOVE_DEBUG', None) is not None else lambda *args, **kwargs: None debug_load = os.environ.get('SD_LOAD_DEBUG', None) debug_process = log.trace if os.environ.get('SD_PROCESS_DEBUG', None) is not None else lambda *args, **kwargs: None diffusers_version = int(diffusers.__version__.split('.')[1]) get_closet_checkpoint_match = get_closest_checkpoint_match # legacy compatibility checkpoint_tiles = checkpoint_titles # legacy compatibility allow_post_quant = None pipe_switch_task_exclude = [ 'AnimateDiffPipeline', 'AnimateDiffSDXLPipeline', 'FluxControlPipeline', 'FluxFillPipeline', 'InstantIRPipeline', 'LTXConditionPipeline', 'OmniGenPipeline', 'OmniGen2Pipeline', 'PhotoMakerStableDiffusionXLPipeline', 'PixelSmithXLPipeline', 'StableDiffusion3ControlNetPipeline', 'StableDiffusionAdapterPipeline', 'StableDiffusionXLAdapterPipeline', 'StableDiffusionControlNetXSPipeline', 'StableDiffusionXLControlNetXSPipeline', 'StableDiffusionReferencePipeline', 'StableDiffusionXLInstantIDPipeline', 'XOmniPipeline', 'HunyuanImagePipeline', 'NucleusMoEImagePipeline', 'AuraFlowPipeline', 'ChronoEditPipeline', 'Kandinsky5I2IPipeline', 'GoogleNanoBananaPipeline', 'Step1XEditPipeline', 'BooguImagePipeline', 'BooguImageTurboPipeline', ] i2i_pipes = [ 'LEditsPPPipelineStableDiffusion', 'LEditsPPPipelineStableDiffusionXL', 'OmniGenPipeline', 'OmniGen2Pipeline', 'StableDiffusionAdapterPipeline', 'StableDiffusionXLAdapterPipeline', 'StableDiffusionControlNetXSPipeline', 'StableDiffusionXLControlNetXSPipeline', 'Step1XEditPipeline', ] def set_huggingface_options(quiet=False): if shared.opts.diffusers_to_gpu: # and model_type.startswith('Stable Diffusion'): sd_hijack_accelerate.hijack_accelerate() else: sd_hijack_accelerate.restore_accelerate() if (shared.opts.runai_streamer_diffusers or shared.opts.runai_streamer_transformers) and (sys.platform == 'linux'): if not quiet: log.debug(f'Loader: runai enabled chunk={os.environ.get("RUNAI_STREAMER_CHUNK_BYTESIZE", "N/A")} limit={os.environ.get("RUNAI_STREAMER_MEMORY_LIMIT", "N/A")}') sd_hijack_safetensors.hijack_safetensors(shared.opts.runai_streamer_diffusers, shared.opts.runai_streamer_transformers) else: sd_hijack_safetensors.restore_safetensors() sd_hijack_hfhub.init_hijack() sd_hijack_transformers.hijack_transformers() def set_caption_load_options(): if shared.opts.caption_to_gpu: sd_hijack_accelerate.hijack_accelerate() else: sd_hijack_accelerate.restore_accelerate() if (shared.opts.runai_streamer_diffusers or shared.opts.runai_streamer_transformers) and (sys.platform == 'linux'): log.debug(f'LLM loader: gpu={shared.opts.caption_to_gpu} runai=True chunk={os.environ.get("RUNAI_STREAMER_CHUNK_BYTESIZE", "N/A")} limit={os.environ.get("RUNAI_STREAMER_MEMORY_LIMIT", "N/A")}') sd_hijack_safetensors.hijack_safetensors(shared.opts.runai_streamer_diffusers, shared.opts.runai_streamer_transformers) else: if shared.opts.caption_to_gpu: log.debug(f'LLM loader: gpu={shared.opts.caption_to_gpu}') sd_hijack_safetensors.restore_safetensors() sd_hijack_hfhub.init_hijack() def set_vae_options(sd_model, vae=None, op:str='model', quiet:bool=False): ops = {} if hasattr(sd_model, "vae"): if vae is not None: sd_model.vae = vae ops['name'] = f"{sd_vae.loaded_vae_file}" if shared.opts.diffusers_vae_upcast != 'default': sd_model.vae.config.force_upcast = True if shared.opts.diffusers_vae_upcast == 'true' else False ops['upcast'] = sd_model.vae.config.force_upcast if shared.opts.no_half_vae and op not in {'decode', 'encode'}: devices.dtype_vae = torch.float32 sd_model.vae.to(devices.dtype_vae) ops['no-half'] = True if hasattr(sd_model, 'vae') and hasattr(sd_model.vae, 'enable_slicing') and hasattr(sd_model.vae, 'disable_slicing'): ops['slicing'] = shared.opts.diffusers_vae_slicing try: if shared.opts.diffusers_vae_slicing: sd_model.vae.enable_slicing() else: sd_model.vae.disable_slicing() except Exception: pass if hasattr(sd_model, 'vae') and hasattr(sd_model.vae, 'enable_tiling') and hasattr(sd_model.vae, 'disable_tiling'): ops['tiling'] = shared.opts.diffusers_vae_tiling try: if shared.opts.diffusers_vae_tiling: if hasattr(sd_model, 'vae') and hasattr(sd_model.vae, 'config') and hasattr(sd_model.vae.config, 'sample_size') and isinstance(sd_model.vae.config.sample_size, int): if getattr(sd_model.vae, "tile_sample_min_size_backup", None) is None: sd_model.vae.tile_sample_min_size_backup = sd_model.vae.tile_sample_min_size sd_model.vae.tile_latent_min_size_backup = sd_model.vae.tile_latent_min_size sd_model.vae.tile_overlap_factor_backup = sd_model.vae.tile_overlap_factor if shared.opts.diffusers_vae_tile_size > 0: sd_model.vae.tile_sample_min_size = int(shared.opts.diffusers_vae_tile_size) sd_model.vae.tile_latent_min_size = int(shared.opts.diffusers_vae_tile_size / (2 ** (len(sd_model.vae.config.block_out_channels) - 1))) else: sd_model.vae.tile_sample_min_size = getattr(sd_model.vae, "tile_sample_min_size_backup", sd_model.vae.tile_sample_min_size) sd_model.vae.tile_latent_min_size = getattr(sd_model.vae, "tile_latent_min_size_backup", sd_model.vae.tile_latent_min_size) if shared.opts.diffusers_vae_tile_overlap != 0.25: sd_model.vae.tile_overlap_factor = float(shared.opts.diffusers_vae_tile_overlap) else: sd_model.vae.tile_overlap_factor = getattr(sd_model.vae, "tile_overlap_factor_backup", sd_model.vae.tile_overlap_factor) ops['tile'] = sd_model.vae.tile_sample_min_size ops['overlap'] = sd_model.vae.tile_overlap_factor sd_model.vae.enable_tiling() else: sd_model.vae.disable_tiling() except Exception: pass if hasattr(sd_model, "vqvae"): ops['upcast'] = True sd_model.vqvae.to(torch.float32) # vqvae is producing nans in fp16 if not quiet and len(ops) > 0: fn = f'{sys._getframe(2).f_code.co_name}:{sys._getframe(1).f_code.co_name}' # pylint: disable=protected-access log.quiet(quiet, f'Setting {op}: component=vae {ops} fn={fn}') def set_diffuser_options(sd_model, vae=None, op:str='model', offload:bool=True, quiet:bool=False): if sd_model is None: log.warning(f'{op} is not loaded') return if hasattr(sd_model, "watermark"): sd_model.watermark = NoWatermark() if not (hasattr(sd_model, "has_accelerate") and sd_model.has_accelerate): sd_model.has_accelerate = False clear_caches() set_vae_options(sd_model, vae, op, quiet) attention.set_diffusers_attention(sd_model, quiet) if shared.opts.diffusers_fuse_projections and hasattr(sd_model, 'fuse_qkv_projections'): try: sd_model.fuse_qkv_projections() log.quiet(quiet, f'Setting {op}: fused-qkv=True') except Exception as e: log.error(f'Setting {op}: fused-qkv=True {e}') if shared.opts.diffusers_fuse_projections and hasattr(sd_model, 'transformer') and hasattr(sd_model.transformer, 'fuse_qkv_projections'): try: sd_model.transformer.fuse_qkv_projections() log.quiet(quiet, f'Setting {op}: fused-qkv=True') except Exception as e: log.error(f'Setting {op}: fused-qkv=True {e}') if shared.opts.diffusers_eval: log.debug(f'Setting {op}: eval=True') def eval_model(model, op=None, sd_model=None): # pylint: disable=unused-argument if hasattr(model, "requires_grad_"): model.requires_grad_(False) model.eval() return model sd_model = apply_function_to_model(sd_model, eval_model, ["Model", "VAE", "TE"], op="eval") if shared.opts.opt_channelslast and hasattr(sd_model, 'unet'): log.quiet(quiet, f'Setting {op}: channels-last=True') sd_model.unet.to(memory_format=torch.channels_last) for module_name in get_module_names(sd_model): module = getattr(sd_model, module_name, None) if hasattr(module, "quantization_config") and getattr(module.quantization_config, "quant_method", None) == "sdnq": if module_name.startswith("text_encoder"): if shared.opts.sdnq_quantize_matmul_mode_te == "Same as model": sdnq_use_quantized_matmul = shared.opts.sdnq_quantize_matmul_mode != "disabled" else: sdnq_use_quantized_matmul = shared.opts.sdnq_quantize_matmul_mode_te != "disabled" else: sdnq_use_quantized_matmul = shared.opts.sdnq_quantize_matmul_mode != "disabled" if module.quantization_config.use_quantized_matmul != sdnq_use_quantized_matmul: from sdnq.loader import apply_sdnq_options_to_model # log.debug(f'Setting {op} {module_name}: sdnq_use_quantized_matmul={sdnq_use_quantized_matmul}') module = apply_sdnq_options_to_model(module, use_quantized_matmul=sdnq_use_quantized_matmul) setattr(sd_model, module_name, module) if offload: set_diffuser_offload(sd_model, op, quiet) def move_model(model, device=None, force=False): def set_execution_device(module, device): if device == torch.device('cpu'): return if hasattr(module, "_hf_hook") and hasattr(module._hf_hook, "execution_device"): # pylint: disable=protected-access try: """ for k, v in module.named_parameters(recurse=True): if v.device == torch.device('meta'): from accelerate.utils import set_module_tensor_to_device set_module_tensor_to_device(module, k, device, tied_params_map=module._hf_hook.tied_params_map) """ module._hf_hook.execution_device = device # pylint: disable=protected-access # module._hf_hook.offload = True except Exception as e: if os.environ.get('SD_MOVE_DEBUG', None): log.error(f'Model move execution device: device={device} {e}') if model is None or device is None: return if getattr(model, 'sdnext_ondemand', False) and device == devices.device: # on-demand components onload at their entry points instead of pre-moves return if hasattr(model, 'pipe'): move_model(model.pipe, device, force) fn = f'{sys._getframe(2).f_code.co_name}:{sys._getframe(1).f_code.co_name}' # pylint: disable=protected-access if getattr(model, 'vae', None) is not None and get_diffusers_task(model) != DiffusersTaskType.TEXT_2_IMAGE: if (device == devices.device) and (model.vae.device.type != "meta") and not getattr(model.vae, 'sdnext_ondemand', False): # force vae back to gpu if not in txt2img mode; on-demand vaes onload at their entry point instead model.vae.to(device) if hasattr(model.vae, '_hf_hook'): debug_move(f'Model move: to={device} class={model.vae.__class__} fn={fn}') # pylint: disable=protected-access model.vae._hf_hook.execution_device = device # pylint: disable=protected-access if hasattr(model, "components"): # accelerate patch for name, m in model.components.items(): if not hasattr(m, "_hf_hook"): # not accelerate hook break if not isinstance(m, torch.nn.Module) or name in getattr(model, '_exclude_from_cpu_offload', []): # modular pipelines lack the attr continue for module in m.modules(): set_execution_device(module, device) # set_execution_device(model, device) if getattr(model, 'has_accelerate', False) and not force: return if hasattr(model, "device") and devices.normalize_device(model.device) == devices.normalize_device(device) and not force: return try: t0 = time.time() try: if hasattr(model, 'device') and model.device == torch.device('meta'): set_execution_device(model, device) elif hasattr(model, 'to'): if device == devices.device and getattr(model, 'sdnext_ondemand_modules', None): pass # the group engine already placed every component; a pipe-level move would only drag on-demand components to the accelerator for the trailing eviction to undo else: model.to(device) if hasattr(model, "prior_pipe"): model.prior_pipe.to(device) if device == devices.device: offload_ondemand(model) # a bulk move must not strand on-demand components on the accelerator; their entry points onload them when needed except Exception as e0: if 'Cannot copy out of meta tensor' in str(e0) or 'must be Tensor, not NoneType' in str(e0): if hasattr(model, "components"): for _name, component in model.components.items(): if hasattr(component, 'modules'): for module in component.modules(): try: if hasattr(module, 'to'): module.to(device) except Exception as e2: if 'Cannot copy out of meta tensor' in str(e2): if os.environ.get('SD_MOVE_DEBUG', None): log.warning(f'Model move meta: module={module.__class__}') module.to_empty(device=device) elif 'enable_sequential_cpu_offload' in str(e0): pass # ignore model move if sequential offload is enabled elif 'Params4bit' in str(e0) or 'Params8bit' in str(e0): 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() except Exception as e1: t1 = time.time() log.warning(f'Model move: device={device} {e1}') if 'move' not in timer.process.records: timer.process.records['move'] = 0 timer.process.records['move'] += t1 - t0 if os.environ.get('SD_MOVE_DEBUG', None) is not None or (t1-t0) > 2: log.debug(f'Model move: device={device} class={model.__class__.__name__} accelerate={getattr(model, "has_accelerate", False)} fn={fn} time={t1-t0:.2f}') # pylint: disable=protected-access devices.torch_gc() def move_base(model, device): if hasattr(model, 'transformer'): key = 'transformer' elif hasattr(model, 'unet'): key = 'unet' else: log.warning(f'Model move: model={model.__class__} device={device} key=unknown') return None log.debug(f'Model move: module={key} device={device}') model = getattr(model, key) R = model.device move_model(model, device) return R def load_diffuser_initial(diffusers_load_config: dict, op='model'): sd_model = None checkpoint_info = None ckpt_basename = os.path.basename(shared.cmd_opts.ckpt) model_name = modelloader.find_diffuser(ckpt_basename) if model_name is not None: log.info(f'Load model {op}: path="{model_name}"') model_file = modelloader.download_diffusers_model(hub_id=model_name, variant=diffusers_load_config.get('variant', None)) try: log.debug(f'Load {op}: config={diffusers_load_config}') diffusers_load_config.pop('cache_dir', None) sd_model = diffusers.DiffusionPipeline.from_pretrained(model_file, cache_dir=shared.opts.diffusers_dir, **diffusers_load_config) except Exception as e: log.error(f'Failed loading model: {model_file} {e}') errors.display(e, f'Load {op}: path="{model_file}"') return None, None list_models() # rescan for downloaded model checkpoint_info = CheckpointInfo(model_name) return sd_model, checkpoint_info def hf_prefetch_configs(checkpoint_info: CheckpointInfo | str, diffusers_load_config: dict, op='model'): # diffusers pipeline downloads build subfolder config allow-patterns with os.path.join, and huggingface_hub>=1.22 # matches patterns with fnmatchcase which does not normalize separators (huggingface/huggingface_hub#4435), # so on windows component config.json files are never downloaded and the incomplete snapshot # still passes the diffusers cache-completeness check so it is never repaired on load # prefetching configs with forward-slash patterns both avoids and repairs such snapshots if os.name != 'nt' or shared.opts.offline_mode: return repo_id = path_to_repo(checkpoint_info) if repo_id is None or '/' not in repo_id or os.path.exists(repo_id): # local folder or single-file model return try: t0 = time.time() hf.snapshot_download(repo_id, cache_dir=shared.opts.diffusers_dir, revision=diffusers_load_config.get('revision', None), allow_patterns=['*/config.json']) log.debug(f'Load {op}: repo="{repo_id}" prefetch=configs time={time.time()-t0:.2f}') except Exception as e: if debug_load: log.debug(f'Load {op}: repo="{repo_id}" prefetch=configs {e}') def load_diffuser_force(detected_model_type: str, checkpoint_info: CheckpointInfo, diffusers_load_config: dict, op='model'): import sdnq # pylint: disable=unused-import sd_model = None global allow_post_quant # pylint: disable=global-statement unload_model_weights(op=op) shared.sd_model = None model_type = detected_model_type.removesuffix('SDNQ').strip() try: if model_type in ['Stable Cascade']: from pipelines.model_stablecascade import load_cascade_combined sd_model = load_cascade_combined(checkpoint_info, diffusers_load_config) allow_post_quant = True elif model_type in ['InstaFlow']: from pipelines.model_instaflow import load_instaflow sd_model = load_instaflow(checkpoint_info, diffusers_load_config) allow_post_quant = True elif model_type in ['SegMoE']: from pipelines.model_segmoe import load_segmoe sd_model = load_segmoe(checkpoint_info, diffusers_load_config) allow_post_quant = True elif model_type in ['PixArtAlpha']: from pipelines.model_pixart import load_pixart_alpha sd_model = load_pixart_alpha(checkpoint_info, diffusers_load_config) allow_post_quant = False elif model_type in ['PixArtSigma']: from pipelines.model_pixart import load_pixart_sigma sd_model = load_pixart_sigma(checkpoint_info, diffusers_load_config) allow_post_quant = False elif model_type in ['Sana']: from pipelines.model_sana import load_sana sd_model = load_sana(checkpoint_info, diffusers_load_config) allow_post_quant = False elif model_type in ['LuminaNext']: from pipelines.model_lumina import load_lumina sd_model = load_lumina(checkpoint_info, diffusers_load_config) allow_post_quant = True elif model_type in ['LuminaDiMOO']: from pipelines.model_lumina import load_lumina_dimoo sd_model = load_lumina_dimoo(checkpoint_info, diffusers_load_config) allow_post_quant = False elif model_type in ['Kolors']: from pipelines.model_kolors import load_kolors sd_model = load_kolors(checkpoint_info, diffusers_load_config) allow_post_quant = True elif model_type in ['AuraFlow']: from pipelines.model_auraflow import load_auraflow sd_model = load_auraflow(checkpoint_info, diffusers_load_config) allow_post_quant = False elif model_type in ['UltraFlux']: from pipelines.model_ultraflux import load_ultraflux sd_model = load_ultraflux(checkpoint_info, diffusers_load_config) allow_post_quant = False elif model_type in ['FLUX']: from pipelines.model_flux import load_flux sd_model = load_flux(checkpoint_info, diffusers_load_config) allow_post_quant = False elif model_type in ['FLUX2']: from pipelines.model_flux2 import load_flux2 sd_model = load_flux2(checkpoint_info, diffusers_load_config) allow_post_quant = False elif model_type in ['FLUX2Klein']: from pipelines.model_flux2_klein import load_flux2_klein sd_model = load_flux2_klein(checkpoint_info, diffusers_load_config) allow_post_quant = False elif model_type in ['FLEX']: from pipelines.model_flex import load_flex sd_model = load_flex(checkpoint_info, diffusers_load_config) allow_post_quant = False elif model_type in ['ZetaChroma']: from pipelines.model_zetachroma import load_zetachroma sd_model = load_zetachroma(checkpoint_info, diffusers_load_config) allow_post_quant = False elif model_type in ['Chroma']: from pipelines.model_chroma import load_chroma sd_model = load_chroma(checkpoint_info, diffusers_load_config) allow_post_quant = False elif model_type in ['Lumina2']: from pipelines.model_lumina import load_lumina2 sd_model = load_lumina2(checkpoint_info, diffusers_load_config) allow_post_quant = False elif model_type in ['Stable Diffusion 3']: from pipelines.model_sd3 import load_sd3 sd_model = load_sd3(checkpoint_info, diffusers_load_config) allow_post_quant = False elif model_type in ['CogView3']: from pipelines.model_cogview import load_cogview3 sd_model = load_cogview3(checkpoint_info, diffusers_load_config) allow_post_quant = False elif model_type in ['CogView4']: from pipelines.model_cogview import load_cogview4 sd_model = load_cogview4(checkpoint_info, diffusers_load_config) allow_post_quant = False elif model_type in ['Meissonic']: from pipelines.model_meissonic import load_meissonic sd_model = load_meissonic(checkpoint_info, diffusers_load_config) allow_post_quant = True elif model_type in ['OmniGen2']: from pipelines.model_omnigen import load_omnigen2 sd_model = load_omnigen2(checkpoint_info, diffusers_load_config) allow_post_quant = False elif model_type in ['OmniGen']: from pipelines.model_omnigen import load_omnigen sd_model = load_omnigen(checkpoint_info, diffusers_load_config) allow_post_quant = False elif model_type in ['HiDreamO1']: from pipelines.model_hidream import load_hidream_o1 sd_model = load_hidream_o1(checkpoint_info, diffusers_load_config) allow_post_quant = False elif model_type in ['HiDream']: from pipelines.model_hidream import load_hidream sd_model = load_hidream(checkpoint_info, diffusers_load_config) allow_post_quant = False elif model_type in ['Cosmos']: from pipelines.model_cosmos import load_cosmos_t2i sd_model = load_cosmos_t2i(checkpoint_info, diffusers_load_config) allow_post_quant = False elif model_type in ['Anima']: from pipelines.model_anima import load_anima sd_model = load_anima(checkpoint_info, diffusers_load_config) allow_post_quant = False elif model_type in ['FLite']: from pipelines.model_flite import load_flite sd_model = load_flite(checkpoint_info, diffusers_load_config) allow_post_quant = False elif model_type in ['WanAI']: from pipelines.model_wanai import load_wan sd_model = load_wan(checkpoint_info, diffusers_load_config) allow_post_quant = False elif model_type in ['MiniMaxH3']: from pipelines.model_minimax import load_minimax sd_model = load_minimax(checkpoint_info, diffusers_load_config) allow_post_quant = False elif model_type in ['ChronoEdit']: from pipelines.model_chrono import load_chrono sd_model = load_chrono(checkpoint_info, diffusers_load_config) allow_post_quant = False elif model_type in ['Bria']: from pipelines.model_bria import load_bria sd_model = load_bria(checkpoint_info, diffusers_load_config) allow_post_quant = False elif model_type in ['Step1XEdit']: from pipelines.model_step1x_edit import load_step1x_edit sd_model = load_step1x_edit(checkpoint_info, diffusers_load_config) allow_post_quant = False elif model_type in ['VIBE']: from pipelines.model_vibe import load_vibe sd_model = load_vibe(checkpoint_info, diffusers_load_config) allow_post_quant = False elif model_type in ['JoyEdit']: from pipelines.model_joy import load_joyedit sd_model = load_joyedit(checkpoint_info, diffusers_load_config) allow_post_quant = False elif model_type in ['Qwen']: from pipelines.model_qwen import load_qwen sd_model = load_qwen(checkpoint_info, diffusers_load_config) allow_post_quant = False elif model_type in ['Boogu']: from pipelines.model_boogu import load_boogu sd_model = load_boogu(checkpoint_info, diffusers_load_config) allow_post_quant = False elif model_type in ['HunyuanDiT']: from pipelines.model_hunyuandit import load_hunyuandit sd_model = load_hunyuandit(checkpoint_info, diffusers_load_config) allow_post_quant = False elif model_type in ['Lens']: from pipelines.model_lens import load_lens sd_model = load_lens(checkpoint_info, diffusers_load_config) allow_post_quant = False elif model_type in ['Kandinsky21']: from pipelines.model_kandinsky import load_kandinsky21 sd_model = load_kandinsky21(checkpoint_info, diffusers_load_config) allow_post_quant = True elif model_type in ['Kandinsky22']: from pipelines.model_kandinsky import load_kandinsky22 sd_model = load_kandinsky22(checkpoint_info, diffusers_load_config) allow_post_quant = True elif model_type in ['Kandinsky30']: from pipelines.model_kandinsky import load_kandinsky3 sd_model = load_kandinsky3(checkpoint_info, diffusers_load_config) allow_post_quant = False elif model_type in ['Kandinsky50']: from pipelines.model_kandinsky import load_kandinsky5 sd_model = load_kandinsky5(checkpoint_info, diffusers_load_config) allow_post_quant = False elif model_type in ['NextStep']: from pipelines.model_nextstep import load_nextstep sd_model = load_nextstep(checkpoint_info, diffusers_load_config) # pylint: disable=assignment-from-none allow_post_quant = False elif model_type in ['HunyuanImage']: from pipelines.model_hyimage import load_hyimage sd_model = load_hyimage(checkpoint_info, diffusers_load_config) # pylint: disable=assignment-from-none allow_post_quant = False elif model_type in ['HunyuanImage3']: from pipelines.model_hyimage import load_hyimage3 sd_model = load_hyimage3(checkpoint_info, diffusers_load_config) # pylint: disable=assignment-from-none allow_post_quant = False elif model_type in ['XOmni']: from pipelines.model_xomni import load_xomni sd_model = load_xomni(checkpoint_info, diffusers_load_config) # pylint: disable=assignment-from-none allow_post_quant = False elif model_type in ['NanoBanana']: from pipelines.model_google import load_nanobanana sd_model = load_nanobanana(checkpoint_info, diffusers_load_config) allow_post_quant = False elif model_type == 'PRXPixel': from pipelines.model_prx_pixel import load_prx_pixel sd_model = load_prx_pixel(checkpoint_info, diffusers_load_config) allow_post_quant = False elif model_type == 'PRX': from pipelines.model_prx import load_prx sd_model = load_prx(checkpoint_info, diffusers_load_config) allow_post_quant = False elif model_type in ['ERNIEImage']: from pipelines.model_ernie import load_ernie_image sd_model = load_ernie_image(checkpoint_info, diffusers_load_config) allow_post_quant = False elif model_type in ['NucleusImage']: from pipelines.model_nucleus import load_nucleus sd_model = load_nucleus(checkpoint_info, diffusers_load_config) allow_post_quant = False elif model_type in ['ZImage']: from pipelines.model_z_image import load_z_image sd_model = load_z_image(checkpoint_info, diffusers_load_config) allow_post_quant = False elif model_type in ['Krea2']: from pipelines.model_krea2 import load_krea2 sd_model = load_krea2(checkpoint_info, diffusers_load_config) allow_post_quant = False elif model_type in ['Ideogram4']: from pipelines.model_ideogram4 import load_ideogram4 sd_model = load_ideogram4(checkpoint_info, diffusers_load_config) allow_post_quant = False elif model_type in ['LongCat']: from pipelines.model_longcat import load_longcat sd_model = load_longcat(checkpoint_info, diffusers_load_config) allow_post_quant = False elif model_type in ['OvisImage']: from pipelines.model_ovis import load_ovis sd_model = load_ovis(checkpoint_info, diffusers_load_config) allow_post_quant = False elif model_type in ['GLMImage']: from pipelines.model_glm import load_glm_image sd_model = load_glm_image(checkpoint_info, diffusers_load_config) allow_post_quant = False elif model_type in ['SDXS']: from pipelines.model_sdxs import load_sdxs sd_model = load_sdxs(checkpoint_info, diffusers_load_config) allow_post_quant = False elif model_type in ['SeFi']: from pipelines.model_sefi import load_sefi sd_model = load_sefi(checkpoint_info, diffusers_load_config) allow_post_quant = False elif model_type in ['MageFlow']: from pipelines.model_mageflow import load_mageflow sd_model = load_mageflow(checkpoint_info, diffusers_load_config) allow_post_quant = True except Exception as e: log.error(f'Load {op}: path="{checkpoint_info.path}" {e}') errors.display(e, 'Load') return None, True if sd_model is not None: return sd_model, True else: return sd_model, False def load_diffuser_folder(model_type: str, pipeline, checkpoint_info: CheckpointInfo, diffusers_load_config: dict, op='model'): sd_model = None files = walk_files(checkpoint_info.path, ['.safetensors', '.bin', '.ckpt']) if 'variant' not in diffusers_load_config and any('diffusion_pytorch_model.fp16' in f for f in files): # deal with diffusers lack of variant fallback when loading diffusers_load_config['variant'] = 'fp16' err0, err1, err2, err3 = None, None, None, None if os.path.exists(checkpoint_info.path) and os.path.isdir(checkpoint_info.path): if os.path.exists(os.path.join(checkpoint_info.path, 'unet', 'diffusion_pytorch_model.bin')): log.debug(f'Load {op}: type=pickle') diffusers_load_config['use_safetensors'] = False if debug_load: log.debug(f'Load {op}: args={diffusers_load_config}') try: #0 - using detected model type and pipeline if (model_type is not None) and (pipeline is not None): if ('sdnq' in model_type.lower()) or ('sdnq' in checkpoint_info.path.lower()): global allow_post_quant # pylint: disable=global-statement allow_post_quant = False sd_model = pipeline.from_pretrained(checkpoint_info.path, cache_dir=shared.opts.diffusers_dir, **diffusers_load_config) sd_model.model_type = sd_model.__class__.__name__ except Exception as e: err0 = e if debug_load: errors.display(e, 'Load Detected') try: # 1 - autopipeline, best choice but not all pipelines are available try: if err0 is not None: sd_model = diffusers.AutoPipelineForText2Image.from_pretrained(checkpoint_info.path, cache_dir=shared.opts.diffusers_dir, **diffusers_load_config) sd_model.model_type = sd_model.__class__.__name__ except ValueError as e: if 'no variant default' in str(e): log.warning(f'Load {op}: variant={diffusers_load_config["variant"]} model="{checkpoint_info.path}" using default variant') diffusers_load_config.pop('variant', None) sd_model = diffusers.AutoPipelineForText2Image.from_pretrained(checkpoint_info.path, cache_dir=shared.opts.diffusers_dir, **diffusers_load_config) sd_model.model_type = sd_model.__class__.__name__ elif 'safetensors found in directory' in str(err1): log.warning(f'Load {op}: type=pickle') diffusers_load_config['use_safetensors'] = False sd_model = diffusers.AutoPipelineForText2Image.from_pretrained(checkpoint_info.path, cache_dir=shared.opts.diffusers_dir, **diffusers_load_config) sd_model.model_type = sd_model.__class__.__name__ else: raise ValueError from e # reraise except Exception as e: err1 = e if debug_load: errors.display(e, 'Load AutoPipeline') try: # 2 - diffusion pipeline, works for most non-linked pipelines if err1 is not None: sd_model = diffusers.DiffusionPipeline.from_pretrained(checkpoint_info.path, cache_dir=shared.opts.diffusers_dir, **diffusers_load_config) sd_model.model_type = sd_model.__class__.__name__ except Exception as e: err2 = e if debug_load: errors.display(e, "Load DiffusionPipeline") try: # 3 - try basic pipeline just in case if err2 is not None: sd_model = diffusers.StableDiffusionXLPipeline.from_pretrained(checkpoint_info.path, cache_dir=shared.opts.diffusers_dir, **diffusers_load_config) sd_model.model_type = sd_model.__class__.__name__ except Exception as e: err3 = e if debug_load: errors.display(e, "Load StableDiffusionPipeline") if sd_model is None: log.error(f'Load {op}: path="{checkpoint_info.path}" pipeline={pipeline.__class__.__name__} type="{model_type}"') log.error(f'Load {op}: attempt detected: {err0}') log.error(f'Load {op}: attempt auto: {err1}') log.error(f'Load {op}: attempt diffusion: {err2}') log.error(f'Load {op}: attempt base: {err3}') return None return sd_model def load_diffuser_file(model_type: str, pipeline, checkpoint_info: CheckpointInfo, diffusers_load_config: dict, op='model'): sd_model = None diffusers_load_config["extract_ema"] = shared.opts.diffusers_extract_ema if pipeline is None: log.error(f'Load {op}: pipeline={shared.opts.diffusers_pipeline} not initialized') return None try: if model_type.startswith('Stable Diffusion'): if shared.opts.diffusers_force_zeros: diffusers_load_config['force_zeros_for_empty_prompt '] = shared.opts.diffusers_force_zeros else: model_config = sd_detect.get_load_config(checkpoint_info.path, model_type, config_type='json') if model_config is not None: if debug_load: log.debug(f'Load {op}: config="{model_config}"') 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, 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['safety_checker'] = None # sd15 specific but we cant know ahead of time diffusers_load_config['requires_safety_checker'] = False # sd15 specific but we cant know ahead of time diffusers_load_config['use_safetensors'] = True diffusers_load_config.pop('cache_dir', None) if shared.opts.stream_load: diffusers_load_config['disable_mmap'] = True if shared.opts.disable_accelerate: from diffusers.utils import import_utils import_utils._accelerate_available = False # pylint: disable=protected-access sd_model = pipeline.from_single_file(checkpoint_info.path, cache_dir=shared.opts.diffusers_dir, **diffusers_load_config) # sd_model = patch_diffuser_config(sd_model, checkpoint_info.path) elif hasattr(pipeline, 'from_ckpt'): diffusers_load_config['cache_dir'] = shared.opts.hfcache_dir sd_model = pipeline.from_ckpt(checkpoint_info.path, **diffusers_load_config) else: log.error(f'Load {op}: file="{checkpoint_info.path}" {shared.opts.diffusers_pipeline} cannot load safetensor model') return None if shared.opts.diffusers_vae_upcast != 'default' and model_type in ['Stable Diffusion', 'Stable Diffusion XL']: diffusers_load_config['force_upcast'] = True if shared.opts.diffusers_vae_upcast == 'true' else False if sd_model is not None: diffusers_load_config.pop('vae', None) diffusers_load_config.pop('safety_checker', None) diffusers_load_config.pop('requires_safety_checker', None) diffusers_load_config.pop('config_files', None) diffusers_load_config.pop('local_files_only', None) log.debug(f'Setting {op}: pipeline={sd_model.__class__.__name__} config={diffusers_load_config}') # pylint: disable=protected-access except Exception as e: log.error(f'Load {op}: file="{checkpoint_info.path}" pipeline={shared.opts.diffusers_pipeline} config={diffusers_load_config} {e}') if 'Weights for this component appear to be missing in the checkpoint' in str(e): log.error(f'Load {op}: file="{checkpoint_info.path}" is not a complete model') else: errors.display(e, 'Load') return None return sd_model def load_sdnq_module(fn: str, module_name: str, load_method: str): t0 = time.time() quantization_config = None quantization_config_path = os.path.join(fn, module_name, 'quantization_config.json') model_config_path = os.path.join(fn, module_name, 'config.json') root_quantization_config_path = os.path.join(fn, 'quantization_config.json') if os.path.exists(quantization_config_path): quantization_config = shared.readfile(quantization_config_path, silent=True, as_type="dict") elif os.path.exists(model_config_path): quantization_config = shared.readfile(model_config_path, silent=True, as_type="dict").get("quantization_config", None) elif os.path.exists(root_quantization_config_path): quantization_config = shared.readfile(root_quantization_config_path, silent=True, as_type="dict") if debug_load: log.debug(f'SDNQ: load_sdnq_module fn={fn} module={module_name} quant_path={quantization_config_path} model_config={model_config_path} root_quant={root_quantization_config_path} found={quantization_config is not None}') if isinstance(quantization_config, dict): log.debug(f'SDNQ: quantization_config keys={list(quantization_config.keys())}') if quantization_config is None: return None, module_name, 0 model_name = os.path.join(fn, module_name) try: import sdnq module = sdnq.load_sdnq_model( model_path=model_name, quantization_config=quantization_config, device=devices.device if shared.opts.diffusers_to_gpu else devices.cpu, dtype=devices.dtype, load_method=load_method, ) t1 = time.time() return module, module_name, t1 - t0 except Exception as e: log.error(f'Load sdnq: model="{fn}" module="{module_name}" {e}') errors.display(e, 'Load') return None, module_name, 0 def load_sdnq_model(checkpoint_info: CheckpointInfo, pipeline, diffusers_load_config: dict, op: str): modules = {} global allow_post_quant # pylint: disable=global-statement allow_post_quant = False t0 = time.time() if shared.opts.runai_streamer_diffusers and (sys.platform == 'linux'): load_method = 'streamer' from installer import install install('runai_model_streamer>=0.15.1') elif shared.opts.sd_parallel_load: load_method = 'threaded' else: load_method = 'safetensors' if debug_load: log.debug(f'SDNQ: load_sdnq_model path={checkpoint_info.path} method={load_method}') for module_name in os.listdir(checkpoint_info.path): module, name, t = load_sdnq_module(checkpoint_info.path, module_name, load_method=load_method) if module is not None: modules[name] = module if debug_load: log.debug(f'SDNQ: loaded module={name} path={os.path.join(checkpoint_info.path, module_name)} time={t:.2f}') log.debug(f'Load {op}: module="{checkpoint_info.name}" module="{name}" gpu={shared.opts.diffusers_to_gpu} prequant=sdnq method={load_method} time={t:.2f}') t1 = time.time() log.debug(f'Load {op}: model="{checkpoint_info.name}" modules={list(modules.keys())} prequant=sdnq time={t1-t0:.2f}') sd_model = pipeline.from_pretrained( checkpoint_info.path, cache_dir=shared.opts.diffusers_dir, **modules, **diffusers_load_config, ) return sd_model def set_overrides(sd_model, checkpoint_info: CheckpointInfo, model_type: str): checkpoint_info_name = checkpoint_info.name.lower() if "Kandinsky" in sd_model.__class__.__name__: sd_model.scheduler.name = 'DDIM' elif ( checkpoint_info.path.lower().endswith('.safetensors') and model_type.startswith("Stable Diffusion") and model_type != "Stable Diffusion 3" ): # SDXL and SD 1.5 scheduler_config = sd_model.scheduler.config # scheduler_config['beta_schedule'] = 'scaled_linear' # scheduler_config['timestep_spacing'] = 'trailing' sd_model.scheduler = diffusers.EulerAncestralDiscreteScheduler.from_config(scheduler_config) if 'bigaspv25' in checkpoint_info_name or 'noobai-rf' in checkpoint_info_name or ('flow' in checkpoint_info_name and 'flower' not in checkpoint_info_name): scheduler_config = sd_model.scheduler.config scheduler_config['prediction_type'] = 'flow_prediction' scheduler_config['beta_schedule'] = 'linear' scheduler_config['use_flow_sigmas'] = True scheduler_config["flow_shift"] = 2.5 sd_model.scheduler = diffusers.UniPCMultistepScheduler.from_config(scheduler_config) log.info(f'Setting override: model="{checkpoint_info.name}" component=scheduler prediction="flow-prediction"') elif 'vpred' in checkpoint_info_name or 'v-pred' in checkpoint_info_name or 'v_pred' in checkpoint_info_name: scheduler_config = sd_model.scheduler.config scheduler_config['prediction_type'] = 'v_prediction' scheduler_config['beta_schedule'] = 'scaled_linear' scheduler_config['rescale_betas_zero_snr'] = True sd_model.scheduler = diffusers.EulerAncestralDiscreteScheduler.from_config(scheduler_config) log.info(f'Setting override: model="{checkpoint_info.name}" component=scheduler prediction="v-prediction" rescale=True') else: try: from safetensors import safe_open with safe_open(checkpoint_info.path, framework='pt') as f: keys = f.keys() if 'v_pred' in keys: # NoobAI VPred models added empty v_pred and ztsnr keys scheduler_config = sd_model.scheduler.config scheduler_config['prediction_type'] = 'v_prediction' scheduler_config['beta_schedule'] = 'scaled_linear' if 'ztsnr' in keys: scheduler_config['rescale_betas_zero_snr'] = True sd_model.scheduler = diffusers.EulerAncestralDiscreteScheduler.from_config(scheduler_config) log.info(f'Setting override: model="{checkpoint_info.name}" component=scheduler prediction="v-prediction" rescale={scheduler_config.get("rescale_betas_zero_snr", False)}') except Exception as e: log.debug(f'Setting override from keys failed: {e}') def set_defaults(sd_model, checkpoint_info: CheckpointInfo): sd_model.sd_model_hash = checkpoint_info.calculate_shorthash() # pylint: disable=attribute-defined-outside-init sd_model.sd_checkpoint_info = checkpoint_info # pylint: disable=attribute-defined-outside-init sd_model.sd_model_checkpoint = checkpoint_info.filename # pylint: disable=attribute-defined-outside-init if hasattr(sd_model, "prior_pipe"): sd_model.default_scheduler = copy.deepcopy(sd_model.prior_pipe.scheduler) if hasattr(sd_model.prior_pipe, "scheduler") else None else: sd_model.default_scheduler = copy.deepcopy(sd_model.scheduler) if hasattr(sd_model, "scheduler") else None sd_model.is_sdxl = False # a1111 compatibility item sd_model.is_sd2 = hasattr(sd_model, 'cond_stage_model') and hasattr(sd_model.cond_stage_model, 'model') # a1111 compatibility item sd_model.is_sd1 = not sd_model.is_sd2 # a1111 compatibility item sd_model.logvar = sd_model.logvar.to(devices.device) if hasattr(sd_model, 'logvar') else None # fix for training shared.opts.data["sd_checkpoint_hash"] = checkpoint_info.sha256 if hasattr(sd_model, "set_progress_bar_config"): sd_model.set_progress_bar_config(bar_format='Progress {rate_fmt}{postfix} {bar:15} {percentage:3.0f}% {n_fmt}/{total_fmt} {elapsed} {remaining}', ncols=120, colour='#327fba') def load_diffuser(checkpoint_info: CheckpointInfo | None = None, op='model', revision=None): # pylint: disable=unused-argument global allow_post_quant # pylint: disable=global-statement allow_post_quant = True # assume default logging.getLogger("diffusers").setLevel(logging.ERROR) timer.load.record("diffusers") diffusers_load_config = { "low_cpu_mem_usage": True, "torch_dtype": devices.dtype, "load_connected_pipeline": True, } if shared.opts.stream_load: diffusers_load_config['disable_mmap'] = True if shared.opts.offload_state_dict: diffusers_load_config['offload_state_dict'] = True diffusers_load_config['offload_folder'] = paths.temp_dir if revision is not None: diffusers_load_config['revision'] = revision if shared.opts.diffusers_model_load_variant != 'default': diffusers_load_config['variant'] = shared.opts.diffusers_model_load_variant if shared.opts.diffusers_pipeline == 'Custom Diffusers Pipeline' and len(shared.opts.custom_diffusers_pipeline) > 0: log.debug(f'Model pipeline: pipeline="{shared.opts.custom_diffusers_pipeline}"') diffusers_load_config['custom_pipeline'] = shared.opts.custom_diffusers_pipeline if shared.opts.data.get('sd_model_checkpoint', '') == 'model.safetensors' or shared.opts.data.get('sd_model_checkpoint', '') == '': shared.opts.data['sd_model_checkpoint'] = "stabilityai/stable-diffusion-xl-base-1.0" if (op == 'model' or op == 'dict'): if (model_data.sd_model is not None) and (checkpoint_info is not None) and (getattr(model_data.sd_model, 'sd_checkpoint_info', None) is not None) and (checkpoint_info.hash == model_data.sd_model.sd_checkpoint_info.hash): # trying to load the same model return else: if (model_data.sd_refiner is not None) and (checkpoint_info is not None) and (getattr(model_data.sd_refiner, 'sd_checkpoint_info', None) is not None) and (checkpoint_info.hash == model_data.sd_refiner.sd_checkpoint_info.hash): # trying to load the same model return sd_model = None handled = False try: # initial load only if sd_model is None: if shared.cmd_opts.ckpt is not None and os.path.isdir(shared.cmd_opts.ckpt) and model_data.initial: sd_model, checkpoint_info = load_diffuser_initial(diffusers_load_config, op) # unload current model checkpoint_info = checkpoint_info or select_checkpoint(op=op) if checkpoint_info is None: unload_model_weights(op=op) return # handle offline mode if shared.opts.offline_mode: log.info(f'Load {op}: offline=True') diffusers_load_config["local_files_only"] = True os.environ['HF_HUB_OFFLINE'] = '1' else: os.environ.pop('HF_HUB_OFFLINE', None) os.unsetenv('HF_HUB_OFFLINE') # detect pipeline pipeline, model_type = sd_detect.detect_pipeline(checkpoint_info, op) set_huggingface_options() # preload vae so it can be used as param vae = None sd_vae.loaded_vae_file = None if model_type is None: log.error(f'Load {op}: pipeline={shared.opts.diffusers_pipeline} not detected') return hf_prefetch_configs(checkpoint_info, diffusers_load_config, op) vae_file = None if model_type.startswith('Stable Diffusion') and (op == 'model' or op == 'refiner'): # preload vae for sd models vae_file, vae_source = sd_vae.resolve_vae(checkpoint_info.filename) vae = sd_vae.load_vae(checkpoint_info.path, vae_file, vae_source) if vae is not None: diffusers_load_config["vae"] = vae timer.load.record("vae") # load with custom loader if sd_model is None and not handled: sd_model, handled = load_diffuser_force(model_type, checkpoint_info, diffusers_load_config, op) if sd_model is not None and not sd_model: log.error(f'Load {op}: type="{model_type}" pipeline="{pipeline}" not loaded') return # load sdnq-prequantized model if sd_model is None and not handled: if model_type.endswith('SDNQ'): sd_model = load_sdnq_model(checkpoint_info, pipeline, diffusers_load_config, op) model_type = model_type.replace(' SDNQ', '') # load from single-file if sd_model is None and not handled: if os.path.isfile(checkpoint_info.path) and checkpoint_info.path.lower().endswith('.safetensors'): sd_model = load_diffuser_file(model_type, pipeline, checkpoint_info, diffusers_load_config, op) # load from hf folder-style if sd_model is None and not handled: if os.path.isdir(checkpoint_info.path) or (checkpoint_info.type == 'huggingface') or (checkpoint_info.type == 'transformer') or (checkpoint_info.type == 'reference'): sd_model = load_diffuser_folder(model_type, pipeline, checkpoint_info, diffusers_load_config, op) if sd_model is None: log.error(f'Load {op}: name="{checkpoint_info.name if checkpoint_info is not None else None}" not loaded') return # a family loader that returns None falls through to the generic folder loader, and for a modular pipeline that # yields an object holding nothing but its from_config helpers, since a modular load registers the fetched # components empty until load_components runs. it reports as a loaded model and then fails on the first # component something reaches, far from the cause specs = getattr(sd_model, '_component_specs', None) # pylint: disable=protected-access if isinstance(specs, dict): fetched = [name for name, spec in specs.items() if getattr(spec, 'default_creation_method', None) == 'from_pretrained'] if len(fetched) > 0 and all(getattr(sd_model, name, None) is None for name in fetched): log.error(f'Load {op}: name="{checkpoint_info.name if checkpoint_info is not None else None}" cls={sd_model.__class__.__name__} no components loaded') return set_overrides(sd_model, checkpoint_info, model_type) set_defaults(sd_model, checkpoint_info) if hasattr(sd_model, "unet") and model_type not in ['Stable Cascade']: # others calls load_diffuser again sd_unet.load_unet(sd_model, checkpoint_info.path) add_noise_pred_to_diffusers_callback(sd_model) timer.load.record("load") if op == 'refiner': model_data.sd_refiner = sd_model else: model_data.sd_model = sd_model reload_text_encoder(initial=True) # must be before embeddings timer.load.record("te") if debug_load: log.trace(f'Model components: {list(get_signature(sd_model).values())}') from modules import textual_inversion sd_model.embedding_db = textual_inversion.EmbeddingDatabase() sd_model.embedding_db.add_embedding_dir(shared.opts.embeddings_dir) sd_model.embedding_db.load_textual_inversion_embeddings(force_reload=True) timer.load.record("embeddings") from modules import prompt_parser_diffusers prompt_parser_diffusers.insert_parser_highjack(sd_model.__class__.__name__) prompt_parser_diffusers.cache.clear() set_diffuser_options(sd_model, vae, op, offload=False) attention.set_attention_dispatcher(sd_model) sd_model = model_quant.do_post_load_quant(sd_model, allow=allow_post_quant) # run this before move model so it can be compressed in CPU timer.load.record("options") set_diffuser_offload(sd_model, op) if op == 'model' and not (os.path.isdir(checkpoint_info.path) or checkpoint_info.type == 'huggingface'): if getattr(shared.sd_model, 'sd_checkpoint_info', None) is not None and vae_file is not None: sd_vae.apply_vae_config(shared.sd_model.sd_checkpoint_info.filename, vae_file, sd_model) # move_model(sd_model, devices.device) timer.load.record("move") except Exception as e: log.error(f"Load {op}: {e}") errors.display(e, "Model") try: if shared.opts.ipex_optimize: sd_model = sd_models_compile.ipex_optimize(sd_model) if (shared.opts.cuda_compile_backend != 'none') and len(shared.opts.cuda_compile) > 0: if 'components' in shared.opts.cuda_compile_options: sd_model = sd_models_compile.compile_diffusers(sd_model, apply_to_components=True) else: if 'Model' in shared.opts.cuda_compile: if hasattr(sd_model, "unet"): sd_model.unet = sd_models_compile.compile_diffusers(sd_model.unet, apply_to_components=False) if hasattr(sd_model, "transformer"): sd_model.transformer = sd_models_compile.compile_diffusers(sd_model.transformer, apply_to_components=False) if 'TE' in shared.opts.cuda_compile: if hasattr(sd_model, "text_encoder"): sd_model.text_encoder = sd_models_compile.compile_diffusers(sd_model.text_encoder, apply_to_components=False) if hasattr(sd_model, "text_encoder_2"): sd_model.text_encoder_2 = sd_models_compile.compile_diffusers(sd_model.text_encoder_2, apply_to_components=False) if hasattr(sd_model, "text_encoder_3"): sd_model.text_encoder_3 = sd_models_compile.compile_diffusers(sd_model.text_encoder_3, apply_to_components=False) if 'VAE' in shared.opts.cuda_compile: if hasattr(sd_model, "vae"): sd_model.vae = sd_models_compile.compile_diffusers(sd_model.vae, apply_to_components=False) timer.load.record("compile") except Exception as e: log.error(f"Compile {op}: {e}") errors.display(e, "Compile") if shared.opts.diffusers_offload_mode != 'balanced': devices.torch_gc(force=True, reason='load') if sd_model is not None: script_callbacks.model_loaded_callback(sd_model) if debug_load: from modules import modelstats modelstats.analyze() log.info(f"Load {op}: type={shared.sd_model_type} time={timer.load.dct()} native={get_native(sd_model)} memory={memory_stats()}") from modules.platform import cleanup cleanup() shared.opts.save(silent=True) class DiffusersTaskType(Enum): TEXT_2_IMAGE = 1 IMAGE_2_IMAGE = 2 INPAINTING = 3 INSTRUCT = 4 MODULAR = 5 def get_diffusers_task(pipe: diffusers.DiffusionPipeline) -> DiffusersTaskType: cls = pipe.__class__.__name__ if cls in i2i_pipes: # special case return DiffusersTaskType.IMAGE_2_IMAGE elif 'ImageToVideo' in cls or cls in ['LTXConditionPipeline', 'StableVideoDiffusionPipeline']: # i2v pipelines return DiffusersTaskType.IMAGE_2_IMAGE elif 'Instruct' in cls: return DiffusersTaskType.INSTRUCT elif 'Modular' in cls: return DiffusersTaskType.MODULAR elif pipe.__class__ in diffusers.pipelines.auto_pipeline.AUTO_IMAGE2IMAGE_PIPELINES_MAPPING.values(): return DiffusersTaskType.IMAGE_2_IMAGE elif pipe.__class__ in diffusers.pipelines.auto_pipeline.AUTO_INPAINT_PIPELINES_MAPPING.values(): return DiffusersTaskType.INPAINTING else: return DiffusersTaskType.TEXT_2_IMAGE def switch_pipe(cls: type[diffusers.DiffusionPipeline] | str, pipeline: diffusers.DiffusionPipeline | None = None, force = False, args: dict | None = None): """ args: - cls: can be pipeline class or a string from custom pipelines for example: diffusers.StableDiffusionPipeline or 'mixture_tiling' - pipeline: source model to be used, if not provided currently loaded model is used - args: any additional components to load into the pipeline for example: { 'vae': None } """ try: if args is None: args = {} if isinstance(cls, str): log.debug(f'Pipeline switch: custom={cls}') cls_object = diffusers.utils.get_class_from_dynamic_module(cls, module_file='pipeline.py') if not cls_object: log.error(f"Pipeline switch: Failed to get class for '{cls}'") if shared.sd_model is not None: return shared.sd_model raise RuntimeError("Pipeline switch: No existing pipeline to fall back to") else: cls_object = cls if pipeline is None: if shared.sd_model is None: raise RuntimeError("Pipeline switch: No existing pipeline to use as default") pipeline = shared.sd_model new_pipe = None signature = get_signature(cls_object) possible = signature.keys() if not force and isinstance(pipeline, cls_object) and args == {}: return pipeline pipe_dict = {} components_used = [] components_skipped = [] components_missing = [] switch_mode = 'none' if hasattr(pipeline, '_internal_dict'): for item in pipeline._internal_dict.keys(): # pylint: disable=protected-access if item in possible: pipe_dict[item] = getattr(pipeline, item, None) components_used.append(item) else: components_skipped.append(item) for item in possible: if item in ['self', 'args', 'kwargs']: # skip continue if signature[item].default != inspect._empty: # has default value so we dont have to worry about it # pylint: disable=protected-access continue if item not in components_used: log.warning(f'Pipeling switch: missing component={item} type={signature[item].annotation}') pipe_dict[item] = None # try but not likely to work components_missing.append(item) new_pipe = cls_object(**pipe_dict) switch_mode = 'auto' elif 'tokenizer_2' in possible and hasattr(pipeline, 'tokenizer_2'): new_pipe = cls_object( vae=pipeline.vae, text_encoder=pipeline.text_encoder, text_encoder_2=pipeline.text_encoder_2, tokenizer=pipeline.tokenizer, tokenizer_2=pipeline.tokenizer_2, unet=pipeline.unet, scheduler=pipeline.scheduler, feature_extractor=getattr(pipeline, 'feature_extractor', None), ) move_model(new_pipe, pipeline.device) switch_mode = 'sdxl' elif 'tokenizer' in possible and hasattr(pipeline, 'tokenizer'): new_pipe = cls_object( vae=pipeline.vae, text_encoder=pipeline.text_encoder, tokenizer=pipeline.tokenizer, unet=pipeline.unet, scheduler=pipeline.scheduler, feature_extractor=getattr(pipeline, 'feature_extractor', None), requires_safety_checker=False, safety_checker=None, ) move_model(new_pipe, pipeline.device) switch_mode = 'sd' else: log.error(f'Pipeline switch error: {pipeline.__class__.__name__} unrecognized') return pipeline if new_pipe is not None: for k, v in args.items(): if k in possible: setattr(new_pipe, k, v) components_used.append(k) else: log.warning(f'Pipeline switch skipping unknown: component={k}') components_skipped.append(k) if new_pipe is not None: copy_diffuser_options(new_pipe, pipeline) sd_hijack_te.init_hijack(new_pipe) if hasattr(new_pipe, "watermark"): new_pipe.watermark = NoWatermark() if switch_mode == 'auto': log.debug(f'Pipeline switch: from={pipeline.__class__.__name__} to={new_pipe.__class__.__name__} components={components_used} skipped={components_skipped} missing={components_missing}') else: log.debug(f'Pipeline switch: from={pipeline.__class__.__name__} to={new_pipe.__class__.__name__} mode={switch_mode}') return new_pipe else: log.error(f'Pipeline switch error: from={pipeline.__class__.__name__} to={cls_object.__name__} empty pipeline') except Exception as e: log.error(f'Pipeline switch error: from={pipeline.__class__.__name__} to={cls if isinstance(cls, str) else cls.__name__} {e}') errors.display(e, 'Pipeline switch') return pipeline def clean_diffuser_pipe(pipe): if pipe is not None and shared.sd_model_type == 'sdxl' and hasattr(pipe, 'config') and 'requires_aesthetics_score' in pipe.config and hasattr(pipe, '_internal_dict'): debug_process(f'Pipeline clean: {pipe.__class__.__name__}') # diffusers adds requires_aesthetics_score with img2img and complains if requires_aesthetics_score exist in txt2img internal_dict = dict(pipe._internal_dict) # pylint: disable=protected-access internal_dict.pop('requires_aesthetics_score', None) del pipe._internal_dict pipe.register_to_config(**internal_dict) def copy_diffuser_options(new_pipe, orig_pipe): new_pipe.sd_checkpoint_info = getattr(orig_pipe, 'sd_checkpoint_info', None) new_pipe.sd_model_checkpoint = getattr(orig_pipe, 'sd_model_checkpoint', None) new_pipe.embedding_db = getattr(orig_pipe, 'embedding_db', None) new_pipe.loaded_loras = getattr(orig_pipe, 'loaded_loras', {}) new_pipe.sd_model_hash = getattr(orig_pipe, 'sd_model_hash', None) new_pipe.has_accelerate = getattr(orig_pipe, 'has_accelerate', False) new_pipe.current_attn_name = getattr(orig_pipe, 'current_attn_name', None) new_pipe.current_attn_overrides = getattr(orig_pipe, 'current_attn_overrides', None) new_pipe.default_scheduler = getattr(orig_pipe, 'default_scheduler', None) new_pipe.image_encoder = getattr(orig_pipe, 'image_encoder', None) new_pipe.feature_extractor = getattr(orig_pipe, 'feature_extractor', None) new_pipe.mask_processor = getattr(orig_pipe, 'mask_processor', None) new_pipe.restore_pipeline = getattr(orig_pipe, 'restore_pipeline', None) new_pipe.is_sdxl = getattr(orig_pipe, 'is_sdxl', False) # a1111 compatibility item new_pipe.is_sd2 = getattr(orig_pipe, 'is_sd2', False) new_pipe.is_sd1 = getattr(orig_pipe, 'is_sd1', True) add_noise_pred_to_diffusers_callback(new_pipe) if getattr(new_pipe, 'task_args', None) is None: new_pipe.task_args = {} new_pipe.task_args.update(getattr(orig_pipe, 'task_args', {})) if new_pipe.has_accelerate: set_accelerate(new_pipe) def backup_pipe_components(pipe): if pipe is None: return {} return { 'sd_checkpoint_info': getattr(pipe, "sd_checkpoint_info", None), 'sd_model_checkpoint': getattr(pipe, "sd_model_checkpoint", None), 'embedding_db': getattr(pipe, "embedding_db", None), 'loaded_loras': getattr(pipe, "loaded_loras", {}), 'sd_model_hash': getattr(pipe, "sd_model_hash", None), 'has_accelerate': getattr(pipe, "has_accelerate", None), 'current_attn_name': getattr(pipe, "current_attn_name", None), 'current_attn_overrides': getattr(pipe, "current_attn_overrides", None), 'default_scheduler': getattr(pipe, "default_scheduler", None), 'image_encoder': getattr(pipe, "image_encoder", None), 'feature_extractor': getattr(pipe, "feature_extractor", None), 'mask_processor': getattr(pipe, "mask_processor", None), 'restore_pipeline': getattr(pipe, "restore_pipeline", None), 'task_args': getattr(pipe, "task_args", None), 'hijack_prompt': hasattr(pipe, "orig_encode_prompt"), 'hijack_vae': hasattr(pipe, "vae") and hasattr(pipe.vae, "orig_decode") } def restore_pipe_components(pipe, components): if pipe is None or components is None: return pipe.sd_checkpoint_info = components['sd_checkpoint_info'] pipe.sd_model_checkpoint = components['sd_model_checkpoint'] pipe.embedding_db = components['embedding_db'] pipe.loaded_loras = components['loaded_loras'] if components['loaded_loras'] is not None else {} pipe.sd_model_hash = components['sd_model_hash'] pipe.has_accelerate = components['has_accelerate'] pipe.current_attn_name = components['current_attn_name'] pipe.current_attn_overrides = components.get('current_attn_overrides') pipe.default_scheduler = components['default_scheduler'] if components['image_encoder'] is not None: pipe.image_encoder = components['image_encoder'] if components['feature_extractor'] is not None: pipe.feature_extractor = components['feature_extractor'] if components['mask_processor'] is not None: pipe.mask_processor = components['mask_processor'] if components['restore_pipeline'] is not None: pipe.restore_pipeline = components['restore_pipeline'] if components['task_args'] is not None: pipe.task_args = components['task_args'] if components['hijack_prompt']: sd_hijack_te.init_hijack(pipe) if components['hijack_vae']: sd_hijack_vae.init_hijack(pipe) if pipe.__class__.__name__ in ['FluxPipeline', 'StableDiffusion3Pipeline']: pipe.register_modules(image_encoder = components['image_encoder']) pipe.register_modules(feature_extractor = components['feature_extractor']) def set_diffuser_pipe(pipe, new_pipe_type): has_errors = False if new_pipe_type == DiffusersTaskType.TEXT_2_IMAGE: clean_diffuser_pipe(pipe) if hasattr(pipe, 'no_task_switch'): del pipe.no_task_switch return pipe if get_diffusers_task(pipe) == new_pipe_type: return pipe if get_diffusers_task(pipe) == DiffusersTaskType.MODULAR: return pipe # skip specific pipelines cls = pipe.__class__.__name__ if cls in pipe_switch_task_exclude: return pipe if 'Video' in cls: return pipe if 'Onnx' in cls: return pipe # in some cases we want to reset the pipeline to parent as they dont have their own variants if (new_pipe_type == DiffusersTaskType.IMAGE_2_IMAGE) or (new_pipe_type == DiffusersTaskType.INPAINTING): if cls == 'StableDiffusionPAGPipeline': pipe = switch_pipe(diffusers.StableDiffusionPipeline, pipe) if cls == 'StableDiffusionXLPAGPipeline': pipe = switch_pipe(diffusers.StableDiffusionXLPipeline, pipe) new_pipe = None components_backup = backup_pipe_components(pipe) if hasattr(pipe, 'config'): # real pipeline which can be auto-switched try: if new_pipe_type == DiffusersTaskType.TEXT_2_IMAGE: new_pipe = diffusers.AutoPipelineForText2Image.from_pipe(pipe) elif new_pipe_type == DiffusersTaskType.IMAGE_2_IMAGE: new_pipe = diffusers.AutoPipelineForImage2Image.from_pipe(pipe) elif new_pipe_type == DiffusersTaskType.INPAINTING: new_pipe = diffusers.AutoPipelineForInpainting.from_pipe(pipe) else: log.warning(f'Pipeline class change failed: type={new_pipe_type} pipeline={cls}') return pipe except Exception as e: # pylint: disable=unused-variable fn = f'{sys._getframe(2).f_code.co_name}:{sys._getframe(1).f_code.co_name}' # pylint: disable=protected-access log.trace(f"Pipeline class change requested: target={new_pipe_type} fn={fn}") # pylint: disable=protected-access log.warning(f'Pipeline class change failed: type={new_pipe_type} pipeline={cls} {e}') if debug_load: errors.display(e, 'Pipeline switch') has_errors = True if not hasattr(pipe, 'config') or has_errors: try: # maybe a wrapper pipeline so just change the class if new_pipe_type == DiffusersTaskType.TEXT_2_IMAGE: pipe.__class__ = diffusers.pipelines.auto_pipeline._get_task_class(diffusers.pipelines.auto_pipeline.AUTO_TEXT2IMAGE_PIPELINES_MAPPING, cls) # pylint: disable=protected-access new_pipe = pipe elif new_pipe_type == DiffusersTaskType.IMAGE_2_IMAGE: pipe.__class__ = diffusers.pipelines.auto_pipeline._get_task_class(diffusers.pipelines.auto_pipeline.AUTO_IMAGE2IMAGE_PIPELINES_MAPPING, cls) # pylint: disable=protected-access new_pipe = pipe elif new_pipe_type == DiffusersTaskType.INPAINTING: pipe.__class__ = diffusers.pipelines.auto_pipeline._get_task_class(diffusers.pipelines.auto_pipeline.AUTO_INPAINT_PIPELINES_MAPPING, cls) # pylint: disable=protected-access new_pipe = pipe else: log.error(f'Pipeline class set failed: type={new_pipe_type} pipeline={cls}') return pipe except Exception as e: # pylint: disable=unused-variable log.warning(f'Pipeline class set failed: type={new_pipe_type} pipeline={cls} {e}') if debug_load: errors.display(e, 'Pipeline switch') has_errors = True return pipe if new_pipe is None: return pipe restore_pipe_components(new_pipe, components_backup) components_backup = None # free memory new_pipe.is_sdxl = getattr(pipe, 'is_sdxl', False) # a1111 compatibility item new_pipe.is_sd2 = getattr(pipe, 'is_sd2', False) new_pipe.is_sd1 = getattr(pipe, 'is_sd1', True) if hasattr(new_pipe, 'watermark'): new_pipe.watermark = NoWatermark() add_noise_pred_to_diffusers_callback(new_pipe) if hasattr(new_pipe, 'pipe'): # also handle nested pipelines new_pipe.pipe = set_diffuser_pipe(new_pipe.pipe, new_pipe_type) add_noise_pred_to_diffusers_callback(new_pipe.pipe) fn = f'{sys._getframe(2).f_code.co_name}:{sys._getframe(1).f_code.co_name}' # pylint: disable=protected-access log.debug(f"Pipeline class change: source={cls} target={new_pipe.__class__.__name__} fn={fn}") # pylint: disable=protected-access if shared.opts.diffusers_offload_mode == 'none' and hasattr(pipe, 'device'): move_model(new_pipe, pipe.device) else: set_diffuser_offload(new_pipe, op='model') pipe = new_pipe return pipe def add_noise_pred_to_diffusers_callback(pipe): if not hasattr(pipe, "_callback_tensor_inputs"): return pipe if pipe.__class__.__name__.startswith("Anima"): return pipe if pipe.__class__.__name__.startswith("ErnieImage"): return pipe if pipe.__class__.__name__.startswith("StableCascade") and ("predicted_image_embedding" not in pipe._callback_tensor_inputs): # pylint: disable=protected-access pipe.prior_pipe._callback_tensor_inputs.append("predicted_image_embedding") # pylint: disable=protected-access elif "noise_pred" not in pipe._callback_tensor_inputs: # pylint: disable=protected-access if pipe.__class__.__name__.startswith("StableDiffusion"): pipe._callback_tensor_inputs.append("noise_pred") # pylint: disable=protected-access elif hasattr(pipe, "scheduler") and "flow" in pipe.scheduler.__class__.__name__.lower(): pipe._callback_tensor_inputs.append("noise_pred") # pylint: disable=protected-access elif hasattr(pipe, "scheduler") and hasattr(pipe.scheduler, "config") and (getattr(pipe.scheduler.config, "prediction_type", "none") == "flow_prediction"): pipe._callback_tensor_inputs.append("noise_pred") # pylint: disable=protected-access elif hasattr(pipe, "default_scheduler") and ("flow" in pipe.default_scheduler.__class__.__name__.lower()): pipe._callback_tensor_inputs.append("noise_pred") # pylint: disable=protected-access elif hasattr(pipe, "default_scheduler") and hasattr(pipe.default_scheduler, "config") and (getattr(pipe.default_scheduler.config, "prediction_type", "none") == "flow_prediction"): pipe._callback_tensor_inputs.append("noise_pred") # pylint: disable=protected-access return pipe def get_native(pipe: diffusers.DiffusionPipeline): if hasattr(pipe, "vae") and hasattr(pipe.vae, "config") and hasattr(pipe.vae.config, "sample_size"): size = pipe.vae.config.sample_size # Stable Diffusion elif hasattr(pipe, "movq") and hasattr(pipe.movq.config, "sample_size"): size = pipe.movq.config.sample_size # Kandinsky elif hasattr(pipe, "unet") and hasattr(pipe.unet.config, "sample_size"): size = pipe.unet.config.sample_size else: size = 0 return size def reload_text_encoder(initial=False): global loaded_te # pylint: disable=global-statement te = shared.opts.sd_text_encoder if initial and (te is None or te == 'Default'): loaded_te = te return # dont unload if not initial and te == loaded_te: return # selection unchanged since it was loaded if shared.sd_model is None: loaded_te = te return signature = get_signature(shared.sd_model) t5 = [k for k, v in signature.items() if 'T5EncoderModel' in str(v)] if hasattr(shared.sd_model, 'text_encoder') and te is not None and 'vit' in te.lower(): from modules.model_te import set_clip set_clip(pipe=shared.sd_model) elif len(t5) > 0: from modules.model_te import set_t5 log.debug(f'Load module: type=t5 path="{te}" module="{t5[0]}"') set_t5(pipe=shared.sd_model, module=t5[0], t5=te, cache_dir=shared.opts.hfcache_dir) elif hasattr(shared.sd_model, 'text_encoder_3'): from modules.model_te import set_t5 log.debug(f'Load module: type=t5 path="{te}" module="text_encoder_3"') set_t5(pipe=shared.sd_model, module='text_encoder_3', t5=te, cache_dir=shared.opts.hfcache_dir) elif not initial: # generic text encoder with no in-place swap path (e.g. Qwen3-VL): reload the model so the # newly selected encoder is read at load time. loaded_te is set first so the reload's own # initial=True call is a no-op rather than recursing. log.info(f'Load module: type=te name="{te}" reloading model to apply') loaded_te = te reload_model_weights(force=True) return loaded_te = te clear_caches(full=True) apply_balanced_offload(shared.sd_model) def reload_model_weights(sd_model=None, info: CheckpointInfo | None = None, op='model', force=False, revision=None): global loaded_te # pylint: disable=global-statement checkpoint_info = info or select_checkpoint(op=op) # are we selecting model or dictionary if checkpoint_info is None: unload_model_weights(op=op) return None jobid = shared.state.begin('Load model') if sd_model is None: sd_model = model_data.sd_model if op == 'model' or op == 'dict' else model_data.sd_refiner loaded_ckpt: CheckpointInfo | None = getattr(sd_model, 'sd_checkpoint_info', None) if sd_model is not None else None changed_checkpoint = loaded_ckpt is None or checkpoint_info is None or loaded_ckpt.filename != checkpoint_info.filename reset_unet = shared.opts.sd_unet not in (None, 'Default', 'None') reset_unet_secondary = shared.opts.sd_unet_secondary not in (None, 'Default', 'None') reset_te = shared.opts.sd_text_encoder not in (None, 'Default', 'None') if op == 'model' and sd_model is not None and changed_checkpoint and (reset_unet or reset_unet_secondary or reset_te): # compare detected model type, not pipeline class: custom-loader arches (e.g. Krea2) load as a # concrete class but detect as generic DiffusionPipeline, so a class compare would falsely reset # across same-arch checkpoints (Base vs Turbo). detect both sides so the comparison is symmetric. try: _, new_type = sd_detect.detect_pipeline(checkpoint_info, op) _, old_type = sd_detect.detect_pipeline(loaded_ckpt, op) if loaded_ckpt is not None else (None, None) except Exception: new_type = old_type = None if new_type is not None and old_type is not None and new_type != old_type: # architecture changed: custom components no longer fit if reset_unet: log.info(f'Load model: type="{old_type}" changed="{new_type}" unet="{shared.opts.sd_unet}" set to default') shared.opts.data["sd_unet"] = 'Default' sd_unet.loaded_unet = None if reset_unet_secondary: log.info(f'Load model: type="{old_type}" changed="{new_type}" unet_secondary="{shared.opts.sd_unet_secondary}" set to default') shared.opts.data["sd_unet_secondary"] = 'Default' sd_unet.loaded_unet_secondary = None if reset_te: log.info(f'Load model: type="{old_type}" changed="{new_type}" te="{shared.opts.sd_text_encoder}" set to default') shared.opts.data["sd_text_encoder"] = 'Default' loaded_te = None if sd_model is None: # previous model load failed current_checkpoint_info = None else: current_checkpoint_info = getattr(sd_model, 'sd_checkpoint_info', None) if current_checkpoint_info is not None and checkpoint_info is not None and current_checkpoint_info.filename == checkpoint_info.filename and not force: shared.state.end(jobid) return None else: move_model(sd_model, devices.cpu) unload_model_weights(op=op) sd_model = None timer.load = timer.Timer() # TODO model load: implement model in-memory caching timer.load.record("config") if sd_model is None or force: sd_model = None load_diffuser(checkpoint_info, op=op, revision=revision) shared.state.end(jobid) if op == 'model': shared.opts.data["sd_model_checkpoint"] = checkpoint_info.title return model_data.sd_model else: shared.opts.data["sd_model_refiner"] = checkpoint_info.title return model_data.sd_refiner shared.state.end(jobid) return None # should not be here def clear_caches(full: bool = False): from modules import prompt_parser_diffusers, memstats, sd_offload from modules.lora import lora_common, lora_load prompt_parser_diffusers.cache.clear() memstats.reset_stats() lora_common.loaded_networks.clear() lora_common.previously_loaded_networks.clear() lora_load.lora_cache.clear() fn = f'{sys._getframe(2).f_code.co_name}:{sys._getframe(1).f_code.co_name}' # pylint: disable=protected-access log.debug(f'Cache clear: full={full} fn={fn}') if full: sd_offload.offload_hook_instance = None def unload_model_weights(op='model'): fn = f'{sys._getframe(2).f_code.co_name}:{sys._getframe(1).f_code.co_name}' # pylint: disable=protected-access if model_data.sd_model or model_data.sd_refiner: clear_caches(full=True) devices.torch_reset() if shared.compiled_model_state is not None: shared.compiled_model_state.compiled_cache.clear() shared.compiled_model_state.req_cache.clear() shared.compiled_model_state.partitioned_modules.clear() if (op == 'model' or op == 'dict') and model_data.sd_model: log.debug(f'Current {op}: {memory_stats()}') if not ('Model' in shared.opts.cuda_compile and (shared.opts.cuda_compile_backend == "openvino_fx" or shared.opts.cuda_compile_backend == "openvino")): disable_offload(model_data.sd_model) move_model(model_data.sd_model, 'meta') model_data.sd_model = None from sdnq.common import reset_compile_caches reset_compile_caches() # dead compiled-dequant graphs and their lifetime recompile counters otherwise accumulate across switches devices.torch_gc(force=True, reason='unload') log.debug(f'Unload {op}: {memory_stats()} fn={fn}') elif (op == 'refiner') and model_data.sd_refiner: log.debug(f'Current {op}: {memory_stats()}') disable_offload(model_data.sd_refiner) move_model(model_data.sd_refiner, 'meta') model_data.sd_refiner = None devices.torch_gc(force=True, reason='unload') log.debug(f'Unload {op}: {memory_stats()} fn={fn}') def hf_auth_check(checkpoint_info: CheckpointInfo | str, force:bool=False): if shared.opts.offline_mode: log.info('Offline mode: skipping auth check') return False login = None if not force: try: fn = checkpoint_info.path if isinstance(checkpoint_info, CheckpointInfo) else checkpoint_info if (fn.endswith('.safetensors') and os.path.isfile(fn)): # skip check for single-file safetensors models return True if os.path.exists(fn) and os.path.isdir(fn) and any(os.path.isfile(os.path.join(fn, f)) for f in ('model_index.json', 'modular_model_index.json')): # skip check for local diffusers folders return True except Exception: pass repo_id = path_to_repo(checkpoint_info) # already handles str or CheckpointInfo if repo_id is None or '/' not in repo_id: # log.warning(f'Auth: repo="{repo_id}" invalid repo id') return False auth_ok = False try: login = modelloader.hf_login() token = os.environ.get('HF_TOKEN', None) hf.auth_check(repo_id, write=False, token=token) auth_ok = True except Exception as e: log.error(f'Auth: repo="{repo_id}" login={login} auth={auth_ok} {e}') return auth_ok def save_model(name: str, path: str | None = None, shard: str = "5GB", overwrite = False): if (name is None) or len(name.strip()) == 0: log.error('Save model: invalid model name') return 'Invalid model name' if not shared.sd_loaded: log.error('Save model: model not loaded') return 'Model not loaded' from sdnq import save_sdnq_model if path is None: path = shared.opts.diffusers_dir model_name = os.path.join(path.strip(), name.strip()) if os.path.exists(model_name) and not overwrite: log.error(f'Save model: path="{model_name}" exists') return f'Path exists: {model_name}' if not shard.strip(): shard = "5GB" # Guard against empty input try: torch.cuda.synchronize() except Exception: pass jobid = shared.state.begin('Save model') try: t0 = time.time() log.info(f'Save model: path="{model_name}" cls={shared.sd_model.__class__.__name__} start') if hasattr(shared.sd_model, '_component_specs'): # modular pipeline: the saved index must reference the destination folder, not the source repos; save_sdnq_model lives in the sdnq submodule and does not pass this flag import functools shared.sd_model.save_pretrained = functools.partial(shared.sd_model.save_pretrained, overwrite_modular_index=True) save_sdnq_model( model=shared.sd_model, model_path=model_name, max_shard_size=shard, is_pipeline=True, ) t1 = time.time() log.info(f'Save model: path="{model_name}" cls={shared.sd_model.__class__.__name__} time={t1 - t0:.2f}') return f'Saved: {model_name}' except Exception as e: log.error(f'Save model: path="{model_name}" {e}') errors.display(e, 'Save model') return f'Error: {e}' finally: if 'save_pretrained' in vars(shared.sd_model): del shared.sd_model.save_pretrained # drop the instance shadow, restoring the class method shared.state.end(jobid) def list_hfcache(): checkpoints = [] for f in os.scandir(shared.opts.hfcache_dir): if not os.path.isdir(f) or not f.name.startswith('models--'): continue checkpoint = CheckpointInfo(filename=f.path, name=path_to_repo(f.name), model_type='hfcache') checkpoints.append(checkpoint) return checkpoints def warn_group_offload(min_vram: int = 0): vram = gpu_stats() vram = round(vram['total'] if "total" in vram else 0) if (0 < vram < min_vram) and (shared.opts.diffusers_offload_mode in ['none', 'balanced', 'model']): log.warning(f'Load model: vram={vram} min={min_vram} offload={shared.opts.diffusers_offload_mode} recommended=group reason="insufficient vram"')