From b681cccccbe7f3df2c3816245457c7a35292be7b Mon Sep 17 00:00:00 2001 From: CalamitousFelicitousness Date: Mon, 25 May 2026 05:47:56 +0100 Subject: [PATCH] feat(ernie): register ERNIE_SPEC for native dispatch ErnieImageTransformer2DModel has no from_single_file; selecting an Ernie finetune in the UNET dropdown previously crashed with "is not a valid JSON file" from from_pretrained. Probe of jibMixErnie_v20.safetensors: 409/409 keys overlap with the model state dict after stripping model.diffusion_model., zero missing or unexpected. Spec is the minimum TransformerSpec(cls=...). --- pipelines/ernie/__init__.py | 29 +++++++++++++++++++++++++++++ pipelines/model_ernie.py | 2 ++ 2 files changed, 31 insertions(+) diff --git a/pipelines/ernie/__init__.py b/pipelines/ernie/__init__.py index e69de29bb..dfed6ff84 100644 --- a/pipelines/ernie/__init__.py +++ b/pipelines/ernie/__init__.py @@ -0,0 +1,29 @@ +"""ERNIE-Image pipeline package. + +Exports :data:`ERNIE_SPEC` for use by :mod:`pipelines.model_ernie` together +with :mod:`pipelines.native_transformer`. The spec captures the +ERNIE-Image-specific knobs that differ from the native-loader defaults: + +- No converter is needed: community Ernie trainer dumps use BFL-style keys + (``model.diffusion_model.``-prefixed) whose names match diffusers' + ``ErnieImageTransformer2DModel.state_dict()`` verbatim after prefix strip. + Probed against ``jibMixErnie_v20.safetensors`` (the upstream community + finetune): 409/409 keys overlap with zero missing or unexpected. +- No siblings; no forbidden markers; default prefixes + (``model.diffusion_model.``, ``diffusion_model.``, ``net.``) cover every + exporter seen in the wild. + +Before this spec, selecting an Ernie finetune via the UNET dropdown crashed +in :func:`diffusers.loaders.ModelMixin.from_pretrained` with a misleading +``OSError: ... is not a valid JSON file`` because +``ErnieImageTransformer2DModel`` lacks ``from_single_file`` support and the +fallback path treats the safetensors as a directory looking for +``config.json``. +""" + +import diffusers + +from pipelines.native_transformer import TransformerSpec + + +ERNIE_SPEC = TransformerSpec(cls=diffusers.ErnieImageTransformer2DModel) diff --git a/pipelines/model_ernie.py b/pipelines/model_ernie.py index 382b1fe80..471c081be 100644 --- a/pipelines/model_ernie.py +++ b/pipelines/model_ernie.py @@ -14,10 +14,12 @@ def load_ernie_image(checkpoint_info, diffusers_load_config=None): load_args, _quant_args = model_quant.get_dit_args(diffusers_load_config, allow_quant=False) log.debug(f'Load model: type=ERNIE-Image repo="{repo_id}" offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype} args={load_args} pe={shared.opts.model_ernie_enable_pe}') + from pipelines.ernie import ERNIE_SPEC transformer = generic.load_transformer( repo_id, cls_name=diffusers.ErnieImageTransformer2DModel, load_config=diffusers_load_config, + native_spec=ERNIE_SPEC, ) text_encoder = generic.load_text_encoder( repo_id,