recoilme commited on
Commit
8165e1e
·
1 Parent(s): 1409b69

Support DiffusionPipeline.from_pretrained(repo, trust_remote_code=True)

Browse files

model_index.json uses the [file, class] form, pipeline.py imports transformer.py relatively so the
dynamic-module loader fetches it too, and the transformer class is published on the diffusers module
(the stock component loader has no auto_map/trust_remote_code path). Tokenizer resolved from the
processor before register_modules, VAE pinned to fp32 on every entry point.

Files changed (3) hide show
  1. README.md +20 -5
  2. model_index.json +4 -5
  3. pipeline.py +34 -4
README.md CHANGED
@@ -74,10 +74,11 @@ Every image below is generated by this pipeline with 30 steps at 1024 px.
74
 
75
  ```python
76
  import torch
77
- from pipeline import ZenImageEditPipeline # shipped in this repo
78
 
79
- pipe = ZenImageEditPipeline.from_pretrained(".", dtype=torch.float16)
80
- pipe.enable_model_cpu_offload() # 14.5 GB DiT + fp32 VAE decoder do not co-reside on 32 GB
 
81
 
82
  # text-to-image
83
  image = pipe(prompt="a red fox in a snowy forest at dusk, cinematic, 85mm",
@@ -92,8 +93,21 @@ image = pipe(prompt="Replace the woman in <image2> with the woman from <image1>;
92
  generator=torch.Generator("cuda").manual_seed(1234)).images[0]
93
  ```
94
 
 
 
 
 
 
 
 
 
95
  CLI: `python example.py --prompt "..." [--image a.png b.png] --out out.png`
96
 
 
 
 
 
 
97
  ### Files
98
 
99
  ```
@@ -110,8 +124,9 @@ media/ the examples above
110
  ```
111
 
112
  `QwenImage21FusionTransformer2DModel` is a custom class defined in `transformer.py`, not registered
113
- inside `diffusers`, so plain `DiffusionPipeline.from_pretrained` does not resolve it. Load through the
114
- shipped pipeline with this folder on `sys.path`.
 
115
 
116
  ### Limitations
117
 
 
74
 
75
  ```python
76
  import torch
77
+ from diffusers import DiffusionPipeline
78
 
79
+ pipe = DiffusionPipeline.from_pretrained("AiArtLab/zen-image-edit", trust_remote_code=True,
80
+ dtype=torch.float16)
81
+ pipe.enable_model_cpu_offload() # 14.5 GB DiT + fp32 VAE decoder do not co-reside on 32 GB
82
 
83
  # text-to-image
84
  image = pipe(prompt="a red fox in a snowy forest at dusk, cinematic, 85mm",
 
93
  generator=torch.Generator("cuda").manual_seed(1234)).images[0]
94
  ```
95
 
96
+ `trust_remote_code=True` pulls `pipeline.py` and `transformer.py` from this repo and runs them, so
97
+ no clone is needed. Cloning works too and gives the class directly:
98
+
99
+ ```python
100
+ from pipeline import ZenImageEditPipeline
101
+ pipe = ZenImageEditPipeline.from_pretrained(".", dtype=torch.float16)
102
+ ```
103
+
104
  CLI: `python example.py --prompt "..." [--image a.png b.png] --out out.png`
105
 
106
+ Requirements: `torch`, `transformers`, `accelerate` and a `diffusers` built with Qwen-Image-2.1
107
+ (`pip install git+https://github.com/huggingface/diffusers`) — the transformer subclasses
108
+ `QwenImage21Transformer2DModel`. `trust_remote_code` saves the clone, it does **not** save the 17 GB
109
+ of weights.
110
+
111
  ### Files
112
 
113
  ```
 
124
  ```
125
 
126
  `QwenImage21FusionTransformer2DModel` is a custom class defined in `transformer.py`, not registered
127
+ inside `diffusers`, so the pipeline publishes it on the `diffusers` module at import time. That is
128
+ what makes the `trust_remote_code=True` one-liner above work; without it the stock component loader
129
+ would not find the DiT class.
130
 
131
  ### Limitations
132
 
model_index.json CHANGED
@@ -1,5 +1,8 @@
1
  {
2
- "_class_name": "ZenImageEditPipeline",
 
 
 
3
  "_diffusers_version": "0.41.0.dev0",
4
  "processor": [
5
  "transformers",
@@ -13,10 +16,6 @@
13
  "transformers",
14
  "Qwen3_5ForConditionalGeneration"
15
  ],
16
- "tokenizer": [
17
- "transformers",
18
- "Qwen2Tokenizer"
19
- ],
20
  "transformer": [
21
  "diffusers",
22
  "QwenImage21FusionTransformer2DModel"
 
1
  {
2
+ "_class_name": [
3
+ "pipeline",
4
+ "ZenImageEditPipeline"
5
+ ],
6
  "_diffusers_version": "0.41.0.dev0",
7
  "processor": [
8
  "transformers",
 
16
  "transformers",
17
  "Qwen3_5ForConditionalGeneration"
18
  ],
 
 
 
 
19
  "transformer": [
20
  "diffusers",
21
  "QwenImage21FusionTransformer2DModel"
pipeline.py CHANGED
@@ -16,7 +16,22 @@ from PIL import Image as PILImage
16
  from diffusers.image_processor import VaeImageProcessor
17
  from diffusers.pipelines.qwenimage21.pipeline_qwenimage21 import QwenImage21Pipeline
18
 
19
- from transformer import QwenImage21FusionTransformer2DModel
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
20
 
21
 
22
  class ZenImageEditPipeline(QwenImage21Pipeline):
@@ -44,6 +59,18 @@ class ZenImageEditPipeline(QwenImage21Pipeline):
44
  # `processor.apply_chat_template(system)`, and the Qwen3.5 processor raises
45
  # "No user query found" on a system-only message. The rest is its own code.
46
  super(QwenImage21Pipeline, self).__init__()
 
 
 
 
 
 
 
 
 
 
 
 
47
  self.register_modules(
48
  vae=vae,
49
  text_encoder=text_encoder,
@@ -75,7 +102,6 @@ class ZenImageEditPipeline(QwenImage21Pipeline):
75
  if not self.student_layers:
76
  raise ValueError("transformer config has no text_fusion_config.student_layers")
77
  self._drop_idx = int(fusion.get("drop_idx", 0))
78
- tokenizer = tokenizer if tokenizer is not None else getattr(processor, "tokenizer", None)
79
  self._img_token_id = int(tokenizer.encode("<|image_pad|>", add_special_tokens=False)[0])
80
  # The prefix crop happens inside the transformer; if the tokenizer and the config ever
81
  # disagree the condition silently slides by a few positions. Check once, at build time.
@@ -162,12 +188,16 @@ class ZenImageEditPipeline(QwenImage21Pipeline):
162
 
163
  @classmethod
164
  def from_pretrained(cls, root=".", dtype=torch.float16, **kwargs):
165
- """Load the pipeline from a model folder (the layout shipped in this repo)."""
 
 
 
 
166
  from diffusers import AutoencoderKLQwenImage21, FlowMatchEulerDiscreteScheduler
167
  from transformers import AutoProcessor, AutoTokenizer, Qwen3_5ForConditionalGeneration
168
 
169
  if not os.path.isdir(root):
170
- raise ValueError(f"expected a model folder with the components, got {root!r}")
171
  transformer = QwenImage21FusionTransformer2DModel.from_pretrained(
172
  os.path.join(root, "transformer"), torch_dtype=dtype
173
  )
 
16
  from diffusers.image_processor import VaeImageProcessor
17
  from diffusers.pipelines.qwenimage21.pipeline_qwenimage21 import QwenImage21Pipeline
18
 
19
+ try:
20
+ # Loaded as a Hub dynamic module (`trust_remote_code=True`): the loader follows relative imports
21
+ # and fetches `transformer.py` from the same repo next to this file.
22
+ from .transformer import QwenImage21FusionTransformer2DModel
23
+ except ImportError:
24
+ # Plain module on `sys.path` (local use, `example.py`).
25
+ from transformer import QwenImage21FusionTransformer2DModel
26
+
27
+ # `DiffusionPipeline.from_pretrained(repo, trust_remote_code=True)` resolves every component by
28
+ # looking its class up in the library named by `model_index.json`, and diffusers has no
29
+ # `auto_map`/`trust_remote_code` path for model components (`models/model_loading_utils.py`).
30
+ # Publishing the class on the `diffusers` module makes the stock loader find it: this module is
31
+ # imported while the pipeline class is resolved, i.e. before any component is loaded.
32
+ import diffusers # noqa: E402
33
+
34
+ diffusers.QwenImage21FusionTransformer2DModel = QwenImage21FusionTransformer2DModel
35
 
36
 
37
  class ZenImageEditPipeline(QwenImage21Pipeline):
 
59
  # `processor.apply_chat_template(system)`, and the Qwen3.5 processor raises
60
  # "No user query found" on a system-only message. The rest is its own code.
61
  super(QwenImage21Pipeline, self).__init__()
62
+ if isinstance(tokenizer, (list, tuple)):
63
+ # Loading through `DiffusionPipeline.from_pretrained` hands tokenizer components over as
64
+ # a `[slow, fast]` pair; we only need one.
65
+ tokenizer = next((t for t in tokenizer if t is not None), None)
66
+ if tokenizer is None:
67
+ # The Hub loader does not pass a tokenizer (it lives inside the processor); resolve it
68
+ # before `register_modules`, otherwise `self.tokenizer` would silently stay None.
69
+ tokenizer = getattr(processor, "tokenizer", None)
70
+ if vae is not None:
71
+ # The repo ships the VAE in fp32, and loading through the Hub applies one `torch_dtype`
72
+ # to every component. Pin it here so both entry points behave identically.
73
+ vae.to(torch.float32)
74
  self.register_modules(
75
  vae=vae,
76
  text_encoder=text_encoder,
 
102
  if not self.student_layers:
103
  raise ValueError("transformer config has no text_fusion_config.student_layers")
104
  self._drop_idx = int(fusion.get("drop_idx", 0))
 
105
  self._img_token_id = int(tokenizer.encode("<|image_pad|>", add_special_tokens=False)[0])
106
  # The prefix crop happens inside the transformer; if the tokenizer and the config ever
107
  # disagree the condition silently slides by a few positions. Check once, at build time.
 
188
 
189
  @classmethod
190
  def from_pretrained(cls, root=".", dtype=torch.float16, **kwargs):
191
+ """Load the pipeline from a local model folder (the layout shipped in this repo).
192
+
193
+ For a Hub repo use `DiffusionPipeline.from_pretrained(repo_id, trust_remote_code=True)`,
194
+ which resolves this class through `model_index.json` and loads the components itself.
195
+ """
196
  from diffusers import AutoencoderKLQwenImage21, FlowMatchEulerDiscreteScheduler
197
  from transformers import AutoProcessor, AutoTokenizer, Qwen3_5ForConditionalGeneration
198
 
199
  if not os.path.isdir(root):
200
+ raise ValueError(f"expected a local model folder with the components, got {root!r}")
201
  transformer = QwenImage21FusionTransformer2DModel.from_pretrained(
202
  os.path.join(root, "transformer"), torch_dtype=dtype
203
  )