From 98fba8ddc3a7af53780dfd8ae8573dbde1f38ff9 Mon Sep 17 00:00:00 2001 From: CalamitousFelicitousness Date: Mon, 25 May 2026 05:47:56 +0100 Subject: [PATCH] feat(pipelines): wire native_spec dispatch in generic.load_transformer New native_spec=None kwarg. When set and the UNET dropdown points at a .safetensors, dispatches to native_transformer.load (threading allow_quant/dtype/modules_to_not_convert/modules_dtype_dict). Pipelines without a spec stay on cls.from_single_file unchanged. --- pipelines/generic.py | 22 +++++++++++++++++++++- pipelines/native_transformer.py | 24 ++++++++++++++++++++++-- 2 files changed, 43 insertions(+), 3 deletions(-) diff --git a/pipelines/generic.py b/pipelines/generic.py index b6aa3e92b..768077ff7 100644 --- a/pipelines/generic.py +++ b/pipelines/generic.py @@ -19,7 +19,17 @@ def _loader(component): return 'runai' if shared.opts.runai_streamer_transformers else 'default' -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, **kwargs): +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, native_spec=None, **kwargs): + """Load a DiT transformer from the base repo, or from a user-selected + single file when the UNET dropdown (``shared.opts.sd_unet``) is set. + + When ``native_spec`` is supplied and a .safetensors override is selected, + dispatches to :func:`pipelines.native_transformer.load` so the per-arch + spec (multi-prefix detection, optional converter, optional sibling + partitioning, forbidden markers) drives the load. Pipelines without a + spec continue to use the legacy ``from_single_file`` path; this preserves + behavior for Mode D arches until they explicitly opt in. + """ if shared.state.interrupted: return None transformer = None @@ -55,6 +65,16 @@ def load_transformer(repo_id, cls_name, load_config=None, subfolder="transformer **load_args, ) transformer = model_quant.do_post_load_quant(transformer, allow=quant_type is not None) + elif local_file is not None and local_file.lower().endswith('.safetensors') 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_transformer args={load_args}') + 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, + ) elif local_file is not None and local_file.lower().endswith('.safetensors'): log.debug(f'Load model: transformer="{local_file}" cls={cls_name.__name__} quant="{quant_type}" loader={_loader("diffusers")} args={load_args}') if dtype is not None: diff --git a/pipelines/native_transformer.py b/pipelines/native_transformer.py index 5f46de9fa..57d2e3ed4 100644 --- a/pipelines/native_transformer.py +++ b/pipelines/native_transformer.py @@ -171,6 +171,11 @@ def load( spec: TransformerSpec, diffusers_cfg: dict | None = None, sibling_classes: dict[str, type] | None = None, + *, + allow_quant: bool = True, + dtype=None, + modules_to_not_convert: list | None = None, + modules_dtype_dict: dict | None = None, ) -> tuple[object, dict[str, object]]: """Load the transformer (and any bundled siblings) from ``local_file``. @@ -180,6 +185,11 @@ def load( Missing sibling classes raise ``ValueError`` if the corresponding sibling keys are present in the bundled file. + Keyword-only arguments ``allow_quant``, ``dtype``, ``modules_to_not_convert``, + and ``modules_dtype_dict`` mirror the corresponding kwargs of + :func:`pipelines.generic.load_transformer` so the dispatch from there can + plumb the caller's intent through unchanged. + Returns ``(transformer, siblings_dict)``. ``siblings_dict`` is keyed by sibling name and is empty for non-sibling specs, or for sibling specs whose keys are absent from the bundled file. @@ -197,7 +207,10 @@ def load( ) _, quant_args = model_quant.get_dit_args( - diffusers_cfg, module="Model", device_map=True, allow_quant=True, + diffusers_cfg, 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) @@ -213,6 +226,7 @@ def load( f"transformer_keys={len(transformer_sd)} siblings={sibling_counts or '{}'}" ) + effective_dtype = dtype if dtype is not None else devices.dtype transformer_cfg = fetch_component_config(repo_id, spec.subfolder) transformer = build_component( component_name="transformer", @@ -223,6 +237,7 @@ def load( acceptable_missing=spec.acceptable_missing, quant_args=quant_args, quant_type=quant_type, + dtype=effective_dtype, ) del transformer_sd devices.torch_gc() @@ -248,6 +263,7 @@ def load( acceptable_missing=sibling_spec.acceptable_missing, quant_args={}, quant_type=None, + dtype=effective_dtype, ) sd_models.allow_post_quant = False @@ -362,9 +378,13 @@ def build_component( acceptable_missing: tuple[str, ...], quant_args: dict, quant_type: str | None, + dtype=None, ) -> object: """Convert (if needed), instantiate, load weights, dtype-cast, quantize, and offload-place a single component. Raises on any hard failure. + + ``dtype`` overrides ``devices.dtype`` when supplied; otherwise the global + default is used. """ try: sd = converter(state_dict) if converter is not None else state_dict @@ -373,7 +393,7 @@ def build_component( validate_state_dict_load(component_name, missing, unexpected, acceptable_missing) del sd devices.torch_gc() - component = component.to(dtype=devices.dtype) + component = component.to(dtype=dtype if dtype is not None else devices.dtype) except Exception as e: log.error(f"Load model: native_transformer {component_name} load failed: {e}") errors.display(e, "Load")