import os import torch from modules import shared, devices, errors, sd_models, sd_offload, model_quant from modules.logger import log from pipelines.generic_util import get_loader from pipelines.generic_map import transformers_map debug = os.environ.get('SD_LOAD_DEBUG', None) is not None def load_transformer( repo_id, cls_name, load_config=None, subfolder="transformer", allow_quant=True, variant=None, dtype=None, modules_to_not_convert=None, modules_dtype_dict=None, use_safetensors=True, native_spec=None, override_slot='primary', trust_remote_code=False, **kwargs): """Load a DiT transformer from the base repo, or from a user-selected single file when the slot's UNET override dropdown is set. ``override_slot`` selects which dropdown this call consumes: ``'primary'`` reads ``shared.opts.sd_unet``, ``'secondary'`` reads ``shared.opts.sd_unet_secondary`` (dual-transformer arches give each transformer its own slot). With ``native_spec`` set and a .safetensors override selected, dispatches to :func:`pipelines.native_transformer.load`. Without a spec, a single-file override falls back to ``from_single_file``. """ if repo_id is None or repo_id.lower() == 'none': return None if shared.state.interrupted: return None transformer = None if load_config is None: load_config = {} if modules_to_not_convert is None: modules_to_not_convert = [] if modules_dtype_dict is None: modules_dtype_dict = {} if cls_name is None: from diffusers import AutoModel cls_name = AutoModel offline_args = {'local_files_only': True} if shared.opts.offline_mode else {} jobid = shared.state.begin('Load DiT') try: load_args, quant_args = model_quant.get_dit_args(load_config, module='Model', device_map=True, allow_quant=allow_quant, modules_to_not_convert=modules_to_not_convert, modules_dtype_dict=modules_dtype_dict) quant_type = model_quant.get_quant_type(quant_args) dtype = dtype or devices.dtype if 'nunchaku-lite' in repo_id.lower(): from modules.attention import hijack_kernels hijack_kernels() def load_from_repo(): nonlocal quant_args log.debug(f'Load model: transformer="{repo_id}" cls={cls_name.__name__} subfolder={subfolder} loader={get_loader("diffusers")} args={load_args}') if 'sdnq-' in repo_id.lower(): quant_args = {} if dtype is not None: load_args['torch_dtype'] = dtype if subfolder is not None: load_args['subfolder'] = subfolder if variant is not None: load_args['variant'] = variant if use_safetensors: load_args['use_safetensors'] = True if trust_remote_code: load_args['trust_remote_code'] = True load_kwargs = {**load_args, **quant_args, **offline_args, **kwargs} module = cls_name.from_pretrained( repo_id, cache_dir=shared.opts.hfcache_dir, **load_kwargs, ) if cls_name.__name__ == 'AutoModel': log.debug(f'Load model: transformer="{repo_id}" cls={module.__class__.__name__}') return module local_file = None override_name = None fallback = True from modules import sd_unet if override_slot == 'primary': override_opt, tracker_attr = 'sd_unet', 'loaded_unet' elif override_slot == 'secondary': override_opt, tracker_attr = 'sd_unet_secondary', 'loaded_unet_secondary' else: raise ValueError(f'load_transformer: unknown override_slot={override_slot}') selected = getattr(shared.opts, override_opt, None) if selected is not None and selected != 'Default': if selected not in list(sd_unet.unet_dict): log.error(f'Load module: type=transformer slot={override_slot} file="{selected}" not found') elif os.path.exists(sd_unet.unet_dict[selected]): local_file = sd_unet.unet_dict[selected] override_name = selected if repo_id.startswith(shared.opts.ckpt_dir) and os.path.exists(repo_id): log.error(f'Load model: transformer="{repo_id}" is incorrectly placed in the checkpoints folder') local_file = repo_id if shared.opts.allow_incomplete_model: log.warning(f'Load model: transformer="{repo_id}" is a local path, attempting to map to a HuggingFace for config fetch') repo_id = transformers_map.get(cls_name.__name__, repo_id) log.warning(f'Load model: transformer="{repo_id}" repo="{repo_id}" attempting to load...') else: return None fallback = False # 1. load gguf if local_file is not None and local_file.lower().endswith('.gguf'): log.debug(f'Load model: transformer="{local_file}" cls={cls_name.__name__} quant="{quant_type}" loader={get_loader("diffusers")} args={load_args}') from modules import ggml transformer = ggml.load_gguf_diffusers(local_file, cls=cls_name, compute_dtype=dtype, config=repo_id, subfolder=subfolder, variant=variant) # transformer = model_quant.do_post_load_quant(transformer, allow=quant_type is not None) # 2. load safetensors with native loader if spec is available elif local_file is not None and local_file.lower().endswith('.safetensors') and native_spec is not None: from pipelines import native_transformer log.debug(f'Load model: transformer="{local_file}" cls={cls_name.__name__} quant="{quant_type}" loader=native args={load_args}') try: transformer, _ = native_transformer.load( local_file, repo_id, native_spec, load_config, allow_quant=allow_quant, dtype=dtype, modules_to_not_convert=modules_to_not_convert, modules_dtype_dict=modules_dtype_dict, quant_args=quant_args, quant_type=quant_type, **kwargs, ) except native_transformer.OverrideArchMismatch as e: log.warning(f'Load model: transformer="{local_file}" override incompatible with cls={cls_name.__name__} ({e})') if fallback: log.warning(f'Load model: transformer="{local_file}" ignoring override and loading base transformer') shared.opts.data[override_opt] = 'Default' setattr(sd_unet, tracker_attr, None) transformer = load_from_repo() # 3. load safetensors with diffusers loader elif local_file is not None and local_file.lower().endswith('.safetensors'): if dtype is not None: load_args['torch_dtype'] = dtype load_args.pop('device_map', None) # single-file uses different syntax loader = cls_name.from_single_file if hasattr(cls_name, 'from_single_file') else cls_name.from_pretrained log.debug(f'Load model: transformer="{local_file}" cls={cls_name.__name__} quant="{quant_type}" loader={get_loader("diffusers")} method={loader.__name__} args={load_args}') load_kwargs = {**load_args, **quant_args, **offline_args, **kwargs} transformer = loader( local_file, cache_dir=shared.opts.hfcache_dir, **load_kwargs, ) # 4. default loading from diffusers repo (also the fallback when an incompatible override is dropped above) else: transformer = load_from_repo() # mark the dropdown selection as loaded so the slot's onchange callback, does not force a redundant full reload for an already-consumed override if (transformer is not None) and (override_name is not None) and getattr(shared.opts, override_opt, None) == override_name: setattr(sd_unet, tracker_attr, override_name) sd_models.allow_post_quant = False # we already handled it if shared.opts.diffusers_offload_mode != 'none' and transformer is not None: sd_models.move_model(transformer, devices.cpu) if (transformer is not None) and not hasattr(transformer, 'quantization_config'): # attach quantization_config if hasattr(transformer, 'config') and hasattr(transformer.config, 'quantization_config'): transformer.quantization_config = transformer.config.quantization_config elif (quant_type is not None) and (quant_args.get('quantization_config', None) is not None): transformer.quantization_config = quant_args.get('quantization_config', None) except Exception as e: log.error(f'Load model: transformer="{repo_id}" cls={cls_name.__name__} {e}') errors.display(e, 'Load') raise devices.torch_gc() shared.state.end(jobid) if transformer is not None: module_size, param_num = sd_offload.get_module_size(transformer) module_memory = sd_offload.get_module_memory(transformer) log.debug(f'Load model: transformer="{repo_id}" quant="{quant_type}" size={module_size:.3f} params={param_num:.3f} memory={module_memory}') try: # quantized models legitimately report the storage dtype (e.g. fp8 comfy_quant adopted via SDNQ); the compute dtype lives in the dequantizers, not the params if getattr(transformer, 'quantization_config', None) is None: actual_dtype = transformer.dtype if isinstance(actual_dtype, torch.dtype) and isinstance(dtype, torch.dtype) and actual_dtype != dtype: force = shared.opts.force_dtype log.warning(f'Load model: transformer="{repo_id}" dtype desired={dtype} actual={actual_dtype} force={force}') if force: transformer = transformer.to(dtype) except Exception: pass return transformer