fix(chroma): route through native_transformer to avoid SDNQ pre-mode crash

Loading a Chroma transformer override via UNET dropdown crashed at
inference with a shape mismatch in unpack_uint4: from_single_file does
not integrate quantization_config the way from_pretrained does, so the
dequantizer state was set up but the .weight tensor was never packed.

CHROMA_SPEC routes through native_transformer with the diffusers
Chroma converter, which loads bare and quantizes explicitly via
sdnq_quantize_model.
This commit is contained in:
CalamitousFelicitousness
2026-05-25 23:00:56 +01:00
parent b29178743c
commit cda4822ca0
2 changed files with 20 additions and 1 deletions
+18
View File
@@ -0,0 +1,18 @@
"""Chroma pipeline package.
Exports :data:`CHROMA_SPEC`. Chroma community files use BFL-style
``model.diffusion_model.``-prefixed keys that need renaming into the
diffusers naming convention, so the spec plugs in
:func:`convert_chroma_transformer_checkpoint_to_diffusers` explicitly.
"""
import diffusers
from diffusers.loaders.single_file_utils import convert_chroma_transformer_checkpoint_to_diffusers
from pipelines.native_transformer import TransformerSpec
CHROMA_SPEC = TransformerSpec(
cls=diffusers.ChromaTransformer2DModel,
converter=convert_chroma_transformer_checkpoint_to_diffusers,
)
+2 -1
View File
@@ -14,7 +14,8 @@ def load_chroma(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=Chroma repo="{repo_id}" config={diffusers_load_config} offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype} args={load_args}')
transformer = generic.load_transformer(repo_id, cls_name=diffusers.ChromaTransformer2DModel, load_config=diffusers_load_config, modules_to_not_convert=["distilled_guidance_layer"])
from pipelines.chroma import CHROMA_SPEC
transformer = generic.load_transformer(repo_id, cls_name=diffusers.ChromaTransformer2DModel, load_config=diffusers_load_config, modules_to_not_convert=["distilled_guidance_layer"], native_spec=CHROMA_SPEC)
text_encoder = generic.load_text_encoder(repo_id, cls_name=transformers.T5EncoderModel, load_config=diffusers_load_config)
pipe = diffusers.ChromaPipeline.from_pretrained(