linoyts HF Staff commited on
Commit
dad256a
·
verified ·
1 Parent(s): a834245

Register Krea2ModularPipeline + krea2 mapping for stock-diffusers init_pipeline

Browse files
Files changed (1) hide show
  1. block.py +20 -4
block.py CHANGED
@@ -1,12 +1,28 @@
1
  """Entry point for the Krea 2 reference-image edit modular blocks (loaded via `trust_remote_code`).
2
 
3
- Importing this module also registers `Krea2Transformer2DModel` into the `diffusers` namespace so that
4
- `ModularPipeline` component loading can resolve the transformer referenced as
5
- `["diffusers", "Krea2Transformer2DModel"]` in the base `krea/Krea-2-Turbo` `model_index.json`.
 
 
 
 
 
6
  """
7
 
 
 
 
 
 
 
8
  from .transformer_krea2 import Krea2Transformer2DModel # noqa: F401 (registers into diffusers namespace)
 
9
  from .modular_blocks_krea2_edit import Krea2EditBlocks
10
 
11
 
12
- __all__ = ["Krea2EditBlocks", "Krea2Transformer2DModel"]
 
 
 
 
 
1
  """Entry point for the Krea 2 reference-image edit modular blocks (loaded via `trust_remote_code`).
2
 
3
+ Importing this module registers the Krea 2 classes into the `diffusers` namespace / modular mapping so that,
4
+ on a stock `diffusers` install (which does not ship Krea 2), component loading and `init_pipeline` resolve:
5
+
6
+ - `Krea2Transformer2DModel` — referenced as `["diffusers", "Krea2Transformer2DModel"]` in the base repo's
7
+ `model_index.json`;
8
+ - `Krea2ModularPipeline` — the model-specific pipeline class (`init_pipeline` maps `model_name="krea2"` to it;
9
+ without this it would fall back to the generic `ModularPipeline`, which lacks the Krea 2 properties such as
10
+ `requires_unconditional_embeds`, `vae_scale_factor`, `num_channels_latents`).
11
  """
12
 
13
+ import diffusers
14
+ from diffusers.modular_pipelines.modular_pipeline import (
15
+ MODULAR_PIPELINE_MAPPING,
16
+ _create_default_map_fn,
17
+ )
18
+
19
  from .transformer_krea2 import Krea2Transformer2DModel # noqa: F401 (registers into diffusers namespace)
20
+ from .modular_pipeline import Krea2ModularPipeline
21
  from .modular_blocks_krea2_edit import Krea2EditBlocks
22
 
23
 
24
+ diffusers.Krea2ModularPipeline = Krea2ModularPipeline
25
+ MODULAR_PIPELINE_MAPPING.setdefault("krea2", _create_default_map_fn("Krea2ModularPipeline"))
26
+
27
+
28
+ __all__ = ["Krea2EditBlocks", "Krea2Transformer2DModel", "Krea2ModularPipeline"]