refactor(pipelines): move per-arch transformer specs into packages

Specs for chrono, cogview, flux2_klein, glm, hunyuandit, hyimage, joy,
kandinsky, longcat, nucleus, ovis, pixart and prx lived module-level in
model_<arch>.py. Move each into pipelines/<arch>/__init__.py to match
the layout used by anima, bria, ernie, f_lite, lens, qwen, step1x and
vibe.

Each model_<arch>.py now imports its spec lazily inside the load
function, so the package only gets pulled in when that arch is actually
loaded (kandinsky 2.x never goes through native dispatch and stays
untouched).

Also drop the "without this spec..." crash paragraph from the new
package docstrings plus ernie/__init__.py and bria/__init__.py.
This commit is contained in:
CalamitousFelicitousness
2026-05-25 17:12:33 +01:00
parent 9326804d69
commit bd60372cc5
28 changed files with 232 additions and 78 deletions
+2 -4
View File
@@ -7,12 +7,10 @@ Bria ships two transformer variants:
- The original Bria family uses a custom :class:`BriaTransformer2DModel`
imported from :mod:`pipelines.bria.transformer_bria`. The custom class
has no entry in diffusers' ``SINGLE_FILE_LOADABLE_CLASSES`` table, so a
community single-file selection previously crashed in
``ModelMixin.from_pretrained``.
has no entry in diffusers' ``SINGLE_FILE_LOADABLE_CLASSES`` table.
- Bria FIBO (and FIBO Edit) use the upstream
:class:`diffusers.BriaFiboTransformer2DModel`, which also lacks a
``SINGLE_FILE_LOADABLE_CLASSES`` entry. Same crash.
``SINGLE_FILE_LOADABLE_CLASSES`` entry.
Both specs use the default 3-prefix detection
(``model.diffusion_model.``, ``diffusion_model.``, ``net.``), no
+15
View File
@@ -0,0 +1,15 @@
"""ChronoEdit pipeline package.
Exports :data:`CHRONOEDIT_SPEC`. The minimum
``TransformerSpec(cls=ChronoEditTransformer3DModel)`` works because
ChronoEdit community files use BFL-style ``model.diffusion_model.``
prefixed keys whose names match the diffusers state_dict verbatim after
prefix strip. No siblings, no converter, no forbidden markers.
"""
import diffusers
from pipelines.native_transformer import TransformerSpec
CHRONOEDIT_SPEC = TransformerSpec(cls=diffusers.ChronoEditTransformer3DModel)
+15
View File
@@ -0,0 +1,15 @@
"""CogView pipeline package.
Exports :data:`COGVIEW3_SPEC` and :data:`COGVIEW4_SPEC`. Both default
specs work because CogView community files use BFL-style
``model.diffusion_model.``-prefixed keys whose names match the diffusers
state_dict verbatim after prefix strip.
"""
import diffusers
from pipelines.native_transformer import TransformerSpec
COGVIEW3_SPEC = TransformerSpec(cls=diffusers.CogView3PlusTransformer2DModel)
COGVIEW4_SPEC = TransformerSpec(cls=diffusers.CogView4Transformer2DModel)
-7
View File
@@ -6,13 +6,6 @@ trainer dumps use BFL-style ``model.diffusion_model.``-prefixed keys
whose names match the diffusers state_dict verbatim after prefix strip
(probed against a community finetune: full key overlap with zero
missing or unexpected). No siblings, no converter, no forbidden markers.
Without this spec, selecting an Ernie file via the UNET dropdown crashes
in :func:`diffusers.loaders.ModelMixin.from_pretrained` with a
misleading ``OSError: ... is not a valid JSON file``, because
``ErnieImageTransformer2DModel`` lacks ``from_single_file`` and the
fallback treats the safetensors as a directory looking for
``config.json``.
"""
import diffusers
+26
View File
@@ -0,0 +1,26 @@
"""Flux 2 Klein pipeline package.
Exports :data:`FLUX2_KLEIN_SPEC`. Klein shares
:class:`Flux2Transformer2DModel` with full Flux 2 but uses a smaller
config (hidden_size and friends). diffusers' ``from_single_file`` picks
the class default (= Flux 2 full), so loading a Klein-shaped community
file crashes at ``load_model_dict_into_meta`` with a shape mismatch like
``expected (36864, 6144), got (24576, 4096)``.
Routing through :mod:`pipelines.native_transformer` pulls the Klein
``transformer/config.json`` from the base repo first and instantiates
``Flux2Transformer2DModel`` at the right size, then runs the diffusers
Flux 2 converter to split fused QKV blocks and rename BFL keys into the
diffusers-expected names.
"""
import diffusers
from diffusers.loaders.single_file_utils import convert_flux2_transformer_checkpoint_to_diffusers
from pipelines.native_transformer import TransformerSpec
FLUX2_KLEIN_SPEC = TransformerSpec(
cls=diffusers.Flux2Transformer2DModel,
converter=convert_flux2_transformer_checkpoint_to_diffusers,
)
+15
View File
@@ -0,0 +1,15 @@
"""GLM-Image pipeline package.
Exports :data:`GLM_IMAGE_SPEC`. The minimum
``TransformerSpec(cls=GlmImageTransformer2DModel)`` works because
GLM-Image community files use BFL-style ``model.diffusion_model.``
prefixed keys whose names match the diffusers state_dict verbatim after
prefix strip.
"""
import diffusers
from pipelines.native_transformer import TransformerSpec
GLM_IMAGE_SPEC = TransformerSpec(cls=diffusers.GlmImageTransformer2DModel)
+14
View File
@@ -0,0 +1,14 @@
"""HunyuanDiT pipeline package.
Exports :data:`HUNYUANDIT_SPEC`. The minimum
``TransformerSpec(cls=HunyuanDiT2DModel)`` works because HunyuanDiT
community files use BFL-style ``model.diffusion_model.``-prefixed keys
whose names match the diffusers state_dict verbatim after prefix strip.
"""
import diffusers
from pipelines.native_transformer import TransformerSpec
HUNYUANDIT_SPEC = TransformerSpec(cls=diffusers.HunyuanDiT2DModel)
+19
View File
@@ -0,0 +1,19 @@
"""HunyuanImage pipeline package.
Exports :data:`HUNYUANIMAGE_SPEC`. The minimum
``TransformerSpec(cls=HunyuanImageTransformer2DModel)`` works because
HunyuanImage 2.1 community files use BFL-style
``model.diffusion_model.``-prefixed keys whose names match the diffusers
state_dict verbatim after prefix strip.
The HunyuanImage 3 path is a transformers ``AutoModelForCausalLM`` and
does not go through the native transformer loader, so it does not need a
spec here.
"""
import diffusers
from pipelines.native_transformer import TransformerSpec
HUNYUANIMAGE_SPEC = TransformerSpec(cls=diffusers.HunyuanImageTransformer2DModel)
+15
View File
@@ -0,0 +1,15 @@
"""Joy-Image-Edit pipeline package.
Exports :data:`JOY_SPEC`. The minimum
``TransformerSpec(cls=JoyImageEditTransformer3DModel)`` works because
Joy community files use BFL-style ``model.diffusion_model.``-prefixed
keys whose names match the diffusers state_dict verbatim after prefix
strip.
"""
import diffusers
from pipelines.native_transformer import TransformerSpec
JOY_SPEC = TransformerSpec(cls=diffusers.JoyImageEditTransformer3DModel)
+23
View File
@@ -0,0 +1,23 @@
"""Kandinsky pipeline package.
Exports :data:`KANDINSKY3_UNET_SPEC` and :data:`KANDINSKY5_SPEC`.
Kandinsky ships several generations under one family:
- Kandinsky 2.1 / 2.2 are unet-based and load through diffusers'
combined pipelines without going through the native loader, so no
spec is needed.
- Kandinsky 3 uses :class:`Kandinsky3UNet` in the ``unet`` subfolder of
the repo, hence ``subfolder='unet'`` instead of the default
``'transformer'``.
- Kandinsky 5 uses the new :class:`Kandinsky5Transformer3DModel` and
follows the default layout.
"""
import diffusers
from pipelines.native_transformer import TransformerSpec
KANDINSKY3_UNET_SPEC = TransformerSpec(cls=diffusers.Kandinsky3UNet, subfolder='unet')
KANDINSKY5_SPEC = TransformerSpec(cls=diffusers.Kandinsky5Transformer3DModel)
+15
View File
@@ -0,0 +1,15 @@
"""LongCat pipeline package.
Exports :data:`LONGCAT_SPEC`. The minimum
``TransformerSpec(cls=LongCatImageTransformer2DModel)`` works because
LongCat community files use BFL-style ``model.diffusion_model.``
prefixed keys whose names match the diffusers state_dict verbatim after
prefix strip.
"""
import diffusers
from pipelines.native_transformer import TransformerSpec
LONGCAT_SPEC = TransformerSpec(cls=diffusers.LongCatImageTransformer2DModel)
+1 -4
View File
@@ -3,10 +3,6 @@ import transformers
from modules import shared, devices, sd_models, model_quant, sd_hijack_te, sd_hijack_vae
from modules.logger import log
from pipelines import generic
from pipelines.native_transformer import TransformerSpec
CHRONOEDIT_SPEC = TransformerSpec(cls=diffusers.ChronoEditTransformer3DModel)
def postprocess(p, result): # pylint: disable=unused-argument
@@ -25,6 +21,7 @@ def load_chrono(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=ChronoEdit repo="{repo_id}" config={diffusers_load_config} offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype} args={load_args}')
from pipelines.chrono import CHRONOEDIT_SPEC
transformer = generic.load_transformer(repo_id, cls_name=diffusers.ChronoEditTransformer3DModel, load_config=diffusers_load_config, subfolder="transformer", native_spec=CHRONOEDIT_SPEC)
text_encoder = generic.load_text_encoder(repo_id, cls_name=transformers.UMT5EncoderModel, load_config=diffusers_load_config, subfolder="text_encoder")
+2 -5
View File
@@ -3,11 +3,6 @@ import diffusers
from modules import shared, devices, sd_models, model_quant, sd_hijack_te
from modules.logger import log
from pipelines import generic
from pipelines.native_transformer import TransformerSpec
COGVIEW3_SPEC = TransformerSpec(cls=diffusers.CogView3PlusTransformer2DModel)
COGVIEW4_SPEC = TransformerSpec(cls=diffusers.CogView4Transformer2DModel)
def load_cogview3(checkpoint_info, diffusers_load_config=None):
@@ -19,6 +14,7 @@ def load_cogview3(checkpoint_info, diffusers_load_config=None):
load_args, _quant_args = model_quant.get_dit_args(diffusers_load_config)
log.debug(f'Load model: type=CogView3 repo="{repo_id}" config={diffusers_load_config} offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype} args={load_args}')
from pipelines.cogview import COGVIEW3_SPEC
transformer = generic.load_transformer(repo_id, cls_name=diffusers.CogView3PlusTransformer2DModel, load_config=diffusers_load_config, subfolder="transformer", native_spec=COGVIEW3_SPEC)
text_encoder = generic.load_text_encoder(repo_id, cls_name=transformers.T5EncoderModel, load_config=diffusers_load_config, subfolder="text_encoder")
@@ -45,6 +41,7 @@ def load_cogview4(checkpoint_info, diffusers_load_config=None):
load_args, _quant_args = model_quant.get_dit_args(diffusers_load_config)
log.debug(f'Load model: type=CogView4 repo="{repo_id}" config={diffusers_load_config} offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype} args={load_args}')
from pipelines.cogview import COGVIEW4_SPEC
transformer = generic.load_transformer(repo_id, cls_name=diffusers.CogView4Transformer2DModel, load_config=diffusers_load_config, subfolder="transformer", native_spec=COGVIEW4_SPEC)
text_encoder = generic.load_text_encoder(repo_id, cls_name=transformers.GlmModel, load_config=diffusers_load_config, subfolder="text_encoder", allow_quant=True)
+1 -17
View File
@@ -1,25 +1,8 @@
import transformers
import diffusers
from diffusers.loaders.single_file_utils import convert_flux2_transformer_checkpoint_to_diffusers
from modules import shared, devices, sd_models, model_quant, sd_hijack_te, sd_hijack_vae
from modules.logger import log
from pipelines import generic
from pipelines.native_transformer import TransformerSpec
# Klein shares Flux2Transformer2DModel with full Flux 2, but uses a smaller
# config (hidden_size and friends). diffusers' from_single_file picks the
# class default (= Flux 2 full), so loading a Klein-shaped community file
# crashes at load_model_dict_into_meta with a shape mismatch like
# "expected (36864, 6144), got (24576, 4096)". Routing through
# native_transformer pulls the Klein transformer/config.json from the base
# repo first and instantiates Flux2Transformer2DModel at the right size,
# then runs the diffusers Flux 2 converter to split fused QKV blocks and
# rename BFL keys into the diffusers-expected names.
FLUX2_KLEIN_SPEC = TransformerSpec(
cls=diffusers.Flux2Transformer2DModel,
converter=convert_flux2_transformer_checkpoint_to_diffusers,
)
def load_flux2_klein(checkpoint_info, diffusers_load_config=None):
@@ -32,6 +15,7 @@ def load_flux2_klein(checkpoint_info, diffusers_load_config=None):
log.debug(f'Load model: type=Flux2Klein repo="{repo_id}" config={diffusers_load_config} offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype} args={load_args}')
# Load transformer - Klein uses Flux2Transformer2DModel (same class as Flux2, different size)
from pipelines.flux2_klein import FLUX2_KLEIN_SPEC
transformer = generic.load_transformer(repo_id, cls_name=diffusers.Flux2Transformer2DModel, load_config=diffusers_load_config, native_spec=FLUX2_KLEIN_SPEC)
# Load text encoder - Klein uses Qwen3 (4B for Klein-4B, 8B for Klein-9B)
+1 -4
View File
@@ -5,10 +5,6 @@ import diffusers
from modules import shared, devices, sd_models, model_quant, sd_hijack_te
from modules.logger import log, console
from pipelines import generic
from pipelines.native_transformer import TransformerSpec
GLM_IMAGE_SPEC = TransformerSpec(cls=diffusers.GlmImageTransformer2DModel)
class GLMTokenProgressProcessor(transformers.LogitsProcessor):
@@ -99,6 +95,7 @@ def load_glm_image(checkpoint_info, diffusers_load_config=None):
log.debug(f'Load model: type=GLM-Image repo="{repo_id}" offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype} args={load_args}')
# Load transformer (DiT decoder - 7B) with quantization support
from pipelines.glm import GLM_IMAGE_SPEC
transformer = generic.load_transformer(
repo_id,
cls_name=diffusers.GlmImageTransformer2DModel,
+1 -4
View File
@@ -3,10 +3,6 @@ import diffusers
from modules import shared, sd_models, devices, model_quant
from modules.logger import log
from pipelines import generic
from pipelines.native_transformer import TransformerSpec
HUNYUANDIT_SPEC = TransformerSpec(cls=diffusers.HunyuanDiT2DModel)
def load_hunyuandit(checkpoint_info, diffusers_load_config=None):
@@ -23,6 +19,7 @@ def load_hunyuandit(checkpoint_info, diffusers_load_config=None):
load_args, _quant_args = model_quant.get_dit_args(diffusers_load_config)
log.debug(f'Load model: type=HunyuanDiT repo="{repo_id}" config={diffusers_load_config} offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype} args={load_args}')
from pipelines.hunyuandit import HUNYUANDIT_SPEC
transformer = generic.load_transformer(repo_id, cls_name=diffusers.HunyuanDiT2DModel, load_config=diffusers_load_config, native_spec=HUNYUANDIT_SPEC)
repo_te = 'Tencent-Hunyuan/HunyuanDiT-v1.2-Diffusers' if 'HunyuanDiT-v1' in repo_id else repo_id
text_encoder_2 = generic.load_text_encoder(repo_te, cls_name=transformers.T5EncoderModel, load_config=diffusers_load_config, subfolder="text_encoder_2", allow_shared=False) # this is not normal t5
+1 -4
View File
@@ -5,10 +5,6 @@ import diffusers
from modules import shared, sd_models, devices, model_quant, sd_hijack_te, sd_hijack_vae
from modules.logger import log
from pipelines import generic
from pipelines.native_transformer import TransformerSpec
HUNYUANIMAGE_SPEC = TransformerSpec(cls=diffusers.HunyuanImageTransformer2DModel)
def load_hyimage(checkpoint_info, diffusers_load_config=None): # pylint: disable=unused-argument
@@ -20,6 +16,7 @@ def load_hyimage(checkpoint_info, diffusers_load_config=None): # pylint: disable
load_args, _quant_args = model_quant.get_dit_args(diffusers_load_config)
log.debug(f'Load model: type=HunyuanImage21 repo="{repo_id}" config={diffusers_load_config} offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype} args={load_args}')
from pipelines.hyimage import HUNYUANIMAGE_SPEC
transformer = generic.load_transformer(repo_id, cls_name=diffusers.HunyuanImageTransformer2DModel, load_config=diffusers_load_config, subfolder="transformer", native_spec=HUNYUANIMAGE_SPEC)
text_encoder = generic.load_text_encoder(repo_id, cls_name=transformers.Qwen2_5_VLForConditionalGeneration, load_config=diffusers_load_config, subfolder="text_encoder")
text_encoder_2 = generic.load_text_encoder(repo_id, cls_name=transformers.T5EncoderModel, load_config=diffusers_load_config, subfolder="text_encoder_2", allow_shared=False)
+1 -4
View File
@@ -3,10 +3,6 @@ import transformers
from modules import shared, devices, sd_models, model_quant, sd_hijack_te, sd_hijack_vae
from modules.logger import log
from pipelines import generic
from pipelines.native_transformer import TransformerSpec
JOY_SPEC = TransformerSpec(cls=diffusers.JoyImageEditTransformer3DModel)
def load_joy(checkpoint_info, diffusers_load_config=None):
@@ -18,6 +14,7 @@ def load_joy(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=JoyImageEdit repo="{repo_id}" config={diffusers_load_config} offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype} args={load_args}')
from pipelines.joy import JOY_SPEC
transformer = generic.load_transformer(
repo_id,
cls_name=diffusers.JoyImageEditTransformer3DModel,
+2 -5
View File
@@ -3,11 +3,6 @@ import diffusers
from modules import shared, sd_models, devices, model_quant, sd_hijack_te, sd_hijack_vae
from modules.logger import log
from pipelines import generic
from pipelines.native_transformer import TransformerSpec
KANDINSKY3_UNET_SPEC = TransformerSpec(cls=diffusers.Kandinsky3UNet, subfolder='unet')
KANDINSKY5_SPEC = TransformerSpec(cls=diffusers.Kandinsky5Transformer3DModel)
def load_kandinsky21(checkpoint_info, diffusers_load_config=None):
@@ -55,6 +50,7 @@ def load_kandinsky3(checkpoint_info, diffusers_load_config=None):
load_args, _quant_args = model_quant.get_dit_args(diffusers_load_config)
log.debug(f'Load model: type=Kandinsky30 repo="{repo_id}" config={diffusers_load_config} offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype} args={load_args}')
from pipelines.kandinsky import KANDINSKY3_UNET_SPEC
unet = generic.load_transformer(repo_id, cls_name=diffusers.Kandinsky3UNet, load_config=diffusers_load_config, subfolder="unet", variant="fp16", native_spec=KANDINSKY3_UNET_SPEC)
text_encoder = generic.load_text_encoder(repo_id, cls_name=transformers.T5EncoderModel, load_config=diffusers_load_config, subfolder="text_encoder", variant="fp16", allow_shared=False)
@@ -88,6 +84,7 @@ def load_kandinsky5(checkpoint_info, diffusers_load_config=None):
load_args, _quant_args = model_quant.get_dit_args(diffusers_load_config)
log.debug(f'Load model: type=Kandinsky50 repo="{repo_id}" config={diffusers_load_config} offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype} args={load_args}')
from pipelines.kandinsky import KANDINSKY5_SPEC
transformer = generic.load_transformer(repo_id, cls_name=diffusers.Kandinsky5Transformer3DModel, load_config=diffusers_load_config, native_spec=KANDINSKY5_SPEC)
text_encoder = generic.load_text_encoder(repo_id, cls_name=transformers.Qwen2_5_VLForConditionalGeneration, load_config=diffusers_load_config)
+1 -4
View File
@@ -3,10 +3,6 @@ import diffusers
from modules import shared, devices, sd_models, model_quant, sd_hijack_te
from modules.logger import log
from pipelines import generic
from pipelines.native_transformer import TransformerSpec
LONGCAT_SPEC = TransformerSpec(cls=diffusers.LongCatImageTransformer2DModel)
def load_longcat(checkpoint_info, diffusers_load_config=None):
@@ -18,6 +14,7 @@ def load_longcat(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=LongCat repo="{repo_id}" config={diffusers_load_config} offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype} args={diffusers_load_config}')
from pipelines.longcat import LONGCAT_SPEC
transformer = generic.load_transformer(repo_id, cls_name=diffusers.LongCatImageTransformer2DModel, load_config=diffusers_load_config, native_spec=LONGCAT_SPEC)
text_encoder = generic.load_text_encoder(repo_id, cls_name=transformers.Qwen2_5_VLForConditionalGeneration, load_config=diffusers_load_config)
text_processor = transformers.Qwen2VLProcessor.from_pretrained(repo_id, subfolder='tokenizer', cache_dir=shared.opts.hfcache_dir)
+1 -4
View File
@@ -3,10 +3,6 @@ import transformers
from modules import shared, devices, sd_models, model_quant, sd_hijack_te, sd_hijack_vae
from modules.logger import log
from pipelines import generic
from pipelines.native_transformer import TransformerSpec
NUCLEUS_SPEC = TransformerSpec(cls=diffusers.NucleusMoEImageTransformer2DModel)
def load_nucleus(checkpoint_info, diffusers_load_config=None):
@@ -18,6 +14,7 @@ def load_nucleus(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=NucleusMoEImage repo="{repo_id}" offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype} args={load_args}')
from pipelines.nucleus import NUCLEUS_SPEC
transformer = generic.load_transformer(
repo_id,
cls_name=diffusers.NucleusMoEImageTransformer2DModel,
+1 -4
View File
@@ -3,10 +3,6 @@ import diffusers
from modules import shared, devices, sd_models, model_quant, sd_hijack_te
from modules.logger import log
from pipelines import generic
from pipelines.native_transformer import TransformerSpec
OVIS_SPEC = TransformerSpec(cls=diffusers.OvisImageTransformer2DModel)
def load_ovis(checkpoint_info, diffusers_load_config=None):
@@ -18,6 +14,7 @@ def load_ovis(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=OvisImage repo="{repo_id}" config={diffusers_load_config} offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype} args={diffusers_load_config}')
from pipelines.ovis import OVIS_SPEC
transformer = generic.load_transformer(repo_id, cls_name=diffusers.OvisImageTransformer2DModel, load_config=diffusers_load_config, native_spec=OVIS_SPEC)
text_encoder = generic.load_text_encoder(repo_id, cls_name=transformers.Qwen3Model, load_config=diffusers_load_config)
+1 -4
View File
@@ -4,10 +4,6 @@ from huggingface_hub import file_exists
from modules import shared, devices, sd_models, model_quant, sd_hijack_te
from modules.logger import log
from pipelines import generic
from pipelines.native_transformer import TransformerSpec
PIXART_SPEC = TransformerSpec(cls=diffusers.PixArtTransformer2DModel)
def load_pixart(checkpoint_info, diffusers_load_config=None):
@@ -28,6 +24,7 @@ def load_pixart(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=PixArtSigma repo="{repo_id}" config={diffusers_load_config} offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype} args={load_args}')
from pipelines.pixart import PIXART_SPEC
transformer = generic.load_transformer(repo_id, cls_name=diffusers.PixArtTransformer2DModel, load_config=diffusers_load_config, native_spec=PIXART_SPEC)
text_encoder = generic.load_text_encoder(repo_id_tenc, cls_name=transformers.T5EncoderModel, load_config=diffusers_load_config)
+1 -4
View File
@@ -2,10 +2,6 @@ import diffusers
from modules import shared, devices, sd_models, model_quant, sd_hijack_te
from modules.logger import log
from pipelines import generic
from pipelines.native_transformer import TransformerSpec
PRX_SPEC = TransformerSpec(cls=diffusers.PRXTransformer2DModel)
def load_prx(checkpoint_info, diffusers_load_config=None):
@@ -18,6 +14,7 @@ def load_prx(checkpoint_info, diffusers_load_config=None):
log.debug(f'Load model: type=PRX repo="{repo_id}" config={diffusers_load_config} offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype} args={load_args}')
from transformers.models.t5gemma.modeling_t5gemma import T5GemmaEncoder
from pipelines.prx import PRX_SPEC
transformer = generic.load_transformer(repo_id, cls_name=diffusers.PRXTransformer2DModel, load_config=diffusers_load_config, native_spec=PRX_SPEC)
text_encoder = generic.load_text_encoder(repo_id, cls_name=T5GemmaEncoder, load_config=diffusers_load_config)
+15
View File
@@ -0,0 +1,15 @@
"""Nucleus MoE-Image pipeline package.
Exports :data:`NUCLEUS_SPEC`. The minimum
``TransformerSpec(cls=NucleusMoEImageTransformer2DModel)`` works because
Nucleus community files use BFL-style ``model.diffusion_model.``
prefixed keys whose names match the diffusers state_dict verbatim after
prefix strip.
"""
import diffusers
from pipelines.native_transformer import TransformerSpec
NUCLEUS_SPEC = TransformerSpec(cls=diffusers.NucleusMoEImageTransformer2DModel)
+15
View File
@@ -0,0 +1,15 @@
"""Ovis-Image pipeline package.
Exports :data:`OVIS_SPEC`. The minimum
``TransformerSpec(cls=OvisImageTransformer2DModel)`` works because
Ovis community files use BFL-style ``model.diffusion_model.``-prefixed
keys whose names match the diffusers state_dict verbatim after prefix
strip.
"""
import diffusers
from pipelines.native_transformer import TransformerSpec
OVIS_SPEC = TransformerSpec(cls=diffusers.OvisImageTransformer2DModel)
+14
View File
@@ -0,0 +1,14 @@
"""PixArt pipeline package.
Exports :data:`PIXART_SPEC`. The minimum
``TransformerSpec(cls=PixArtTransformer2DModel)`` works because PixArt
community files use BFL-style ``model.diffusion_model.``-prefixed keys
whose names match the diffusers state_dict verbatim after prefix strip.
"""
import diffusers
from pipelines.native_transformer import TransformerSpec
PIXART_SPEC = TransformerSpec(cls=diffusers.PixArtTransformer2DModel)
+14
View File
@@ -0,0 +1,14 @@
"""PRX pipeline package.
Exports :data:`PRX_SPEC`. The minimum
``TransformerSpec(cls=PRXTransformer2DModel)`` works because PRX
community files use BFL-style ``model.diffusion_model.``-prefixed keys
whose names match the diffusers state_dict verbatim after prefix strip.
"""
import diffusers
from pipelines.native_transformer import TransformerSpec
PRX_SPEC = TransformerSpec(cls=diffusers.PRXTransformer2DModel)