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=...).
This commit is contained in:
CalamitousFelicitousness
2026-05-25 05:47:56 +01:00
parent 98fba8ddc3
commit b681cccccb
2 changed files with 31 additions and 0 deletions
+29
View File
@@ -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)
+2
View File
@@ -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,