Files
automatic/pipelines/generic_transformer.py
Vladimir Mandic 05eef5e919 update nunchaku-lite
Signed-off-by: Vladimir Mandic <mandic00@live.com>
2026-08-18 18:00:34 +02:00

209 lines
9.7 KiB
Python

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 = {}
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
return cls_name.from_pretrained(
repo_id,
cache_dir=shared.opts.hfcache_dir,
**load_args,
**quant_args,
**kwargs,
)
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}')
transformer = loader(
local_file,
cache_dir=shared.opts.hfcache_dir,
**load_args,
**quant_args,
**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