Image-to-Image
Diffusers
Safetensors
ZenImageEditPipeline
text-to-image
image-editing
qwen-image
text-encoder
adapter
Instructions to use AiArtLab/zen-image-edit with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Diffusers
How to use AiArtLab/zen-image-edit with Diffusers:
pip install -U diffusers transformers accelerate
import torch from diffusers import DiffusionPipeline from diffusers.utils import load_image # switch to "mps" for apple devices pipe = DiffusionPipeline.from_pretrained("AiArtLab/zen-image-edit", dtype=torch.bfloat16, device_map="cuda") prompt = "Turn this cat into a dog" input_image = load_image("https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/cat.png") image = pipe(image=input_image, prompt=prompt).images[0] - Notebooks
- Google Colab
- Kaggle
Download pipeline.py from AiArtLab/zen-image-edit: direct link, hf CLI and curl.
- Browser
- Download file 11.4 kB
-
https://huggingface.co/AiArtLab/zen-image-edit/resolve/88a93c15aa070da8fbde00158c92e2d9ee5bcb7f/pipeline.py
- Command line
-
hf download hf://AiArtLab/zen-image-edit@88a93c15aa070da8fbde00158c92e2d9ee5bcb7f/pipeline.py
-
curl -L -o pipeline.py https://huggingface.co/AiArtLab/zen-image-edit/resolve/88a93c15aa070da8fbde00158c92e2d9ee5bcb7f/pipeline.py
11.4 kB
| """Inference pipeline for zen-image-edit: Qwen3.5-0.8B text encoder, adapter inside the DiT. | |
| Compared to `QwenImage21Pipeline` exactly one method differs: `_get_qwen_prompt_embeds` returns the | |
| Qwen3.5-0.8B hidden-state stack `(B, L, K*d)` plus masks instead of `(B, L, 4096)` from Qwen3-VL-8B. | |
| The transformer's own `text_fusion` turns that into the condition, and it also crops the system | |
| prompt — the adapter was trained on the full sequence (`drop_first=0`), so cropping before the | |
| fusion would shift every position. | |
| Latents, VAE, scheduler, edit geometry (`<imageN>` blocks, one vision slot per 2x2 latent group) and | |
| the prefix KV cache are inherited unchanged. | |
| """ | |
| import os | |
| import torch | |
| from PIL import Image as PILImage | |
| from diffusers.image_processor import VaeImageProcessor | |
| from diffusers.pipelines.qwenimage21.pipeline_qwenimage21 import QwenImage21Pipeline | |
| try: | |
| # Loaded as a Hub dynamic module (`trust_remote_code=True`): the loader follows relative imports | |
| # and fetches `transformer.py` from the same repo next to this file. | |
| from .transformer import QwenImage21FusionTransformer2DModel | |
| except ImportError: | |
| # Plain module on `sys.path` (local use, `example.py`). | |
| from transformer import QwenImage21FusionTransformer2DModel | |
| # `DiffusionPipeline.from_pretrained(repo, trust_remote_code=True)` resolves every component by | |
| # looking its class up in the library named by `model_index.json`, and diffusers has no | |
| # `auto_map`/`trust_remote_code` path for model components (`models/model_loading_utils.py`). | |
| # Publishing the class on the `diffusers` module makes the stock loader find it: this module is | |
| # imported while the pipeline class is resolved, i.e. before any component is loaded. | |
| import diffusers # noqa: E402 | |
| diffusers.QwenImage21FusionTransformer2DModel = QwenImage21FusionTransformer2DModel | |
| class ZenImageEditPipeline(QwenImage21Pipeline): | |
| """Qwen-Image-2.1 with Qwen3.5-0.8B instead of Qwen3-VL-8B; one class for both modes. | |
| Like `QwenImage21Pipeline` itself, the mode is picked by the arguments: `image=None` is | |
| text-to-image, `image=[...]` is editing. | |
| Components: `transformer` — `QwenImage21FusionTransformer2DModel` (stock DiT + `text_fusion`), | |
| `text_encoder` — `Qwen3_5ForConditionalGeneration`, `processor` — the student processor | |
| (images and tokenization), `tokenizer` — the student tokenizer, `vae`/`scheduler` — the | |
| Qwen-Image-2.1 ones. | |
| """ | |
| def __init__( | |
| self, | |
| scheduler, | |
| vae, | |
| text_encoder, | |
| processor, | |
| transformer, | |
| tokenizer=None, | |
| ): | |
| # `QwenImage21Pipeline.__init__` is not reused: it derives `_drop_idx` from | |
| # `processor.apply_chat_template(system)`, and the Qwen3.5 processor raises | |
| # "No user query found" on a system-only message. The rest is its own code. | |
| super(QwenImage21Pipeline, self).__init__() | |
| if isinstance(tokenizer, (list, tuple)): | |
| # Loading through `DiffusionPipeline.from_pretrained` hands tokenizer components over as | |
| # a `[slow, fast]` pair; we only need one. | |
| tokenizer = next((t for t in tokenizer if t is not None), None) | |
| if tokenizer is None: | |
| # The Hub loader does not pass a tokenizer (it lives inside the processor); resolve it | |
| # before `register_modules`, otherwise `self.tokenizer` would silently stay None. | |
| tokenizer = getattr(processor, "tokenizer", None) | |
| if vae is not None: | |
| # The repo ships the VAE in fp32, and loading through the Hub applies one `torch_dtype` | |
| # to every component. Pin it here so both entry points behave identically. | |
| vae.to(torch.float32) | |
| self.register_modules( | |
| vae=vae, | |
| text_encoder=text_encoder, | |
| processor=processor, | |
| tokenizer=tokenizer, | |
| transformer=transformer, | |
| scheduler=scheduler, | |
| ) | |
| self.vae_scale_factor = 16 | |
| self.latent_channels = self.vae.config.z_dim if getattr(self, "vae", None) else 64 | |
| self.image_processor = VaeImageProcessor( | |
| vae_scale_factor=self.vae_scale_factor, vae_latent_channels=self.latent_channels | |
| ) | |
| self.sys_prompt = "Comprehend and analyze the provided prompt." | |
| # The prompt is a raw template string, not `apply_chat_template`: the checkpoint was | |
| # trained on this one. | |
| self.prompt_template_t2i = ( | |
| f"<|im_start|>system\n{self.sys_prompt}<|im_end|>\n" | |
| f"<|im_start|>user\n{{}}<|im_end|>\n" | |
| f"<|im_start|>assistant\n" | |
| ) | |
| self.prompt_template_ti2i = ( | |
| f"<|im_start|>system\n{self.sys_prompt}<|im_end|>\n" | |
| f"<|im_start|>user\n<image1><|vision_start|><|image_pad|><|vision_end|>{{}}<|im_end|>\n" | |
| f"<|im_start|>assistant\n" | |
| ) | |
| fusion = dict(getattr(transformer.config, "text_fusion_config", {}) or {}) | |
| self.student_layers = [int(i) for i in fusion.get("student_layers", ())] | |
| if not self.student_layers: | |
| raise ValueError("transformer config has no text_fusion_config.student_layers") | |
| self._drop_idx = int(fusion.get("drop_idx", 0)) | |
| self._img_token_id = int(tokenizer.encode("<|image_pad|>", add_special_tokens=False)[0]) | |
| # The prefix crop happens inside the transformer; if the tokenizer and the config ever | |
| # disagree the condition silently slides by a few positions. Check once, at build time. | |
| prefix = f"<|im_start|>system\n{self.sys_prompt}<|im_end|>\n" | |
| got = len(tokenizer(prefix, add_special_tokens=False)["input_ids"]) | |
| if got != self._drop_idx: | |
| raise ValueError( | |
| f"text_fusion_config.drop_idx={self._drop_idx}, but the tokenizer's system prefix " | |
| f"is {got} tokens" | |
| ) | |
| def _get_qwen_prompt_embeds(self, prompt=None, image=None, device=None): | |
| """Student slices -> flat `(B, L, K*d)`; masks are full-length, the prefix is not cropped here.""" | |
| device = device or self._execution_device | |
| prompt = [prompt] if isinstance(prompt, str) else prompt | |
| # Qwen has no BOS token, so an empty string would leave the encoder with nothing to read. | |
| prompt = [" " if not p else p for p in prompt] | |
| is_t2i = image is None | |
| if is_t2i: | |
| prompts = [self.prompt_template_t2i.format(t) for t in prompt] | |
| condition_pil_list = [] | |
| else: | |
| replace = "<image1><|vision_start|><|image_pad|><|vision_end|>" | |
| for i in range(2, len(image) + 1): | |
| replace += f" <image{i}><|vision_start|><|image_pad|><|vision_end|>" | |
| template = self.prompt_template_ti2i.replace( | |
| "<image1><|vision_start|><|image_pad|><|vision_end|>", replace | |
| ) | |
| prompts = [template.format(t) for t in prompt] | |
| condition_pil_list = [] | |
| for _ in prompt: | |
| for img in image: | |
| if not isinstance(img, PILImage.Image): | |
| img = PILImage.fromarray(img) | |
| if img.mode == "RGBA": | |
| # The checkpoint saw alpha composited over white. Encoder only: the VAE | |
| # still reads all four channels. | |
| white = PILImage.new("RGB", img.size, (255, 255, 255)) | |
| white.paste(img, mask=img.getchannel("A")) | |
| img = white | |
| condition_pil_list.append(img) | |
| # Right padding, as in adapter training: positions of the real tokens do not shift, and the | |
| # fusion attention branch indexes its position table by absolute index. | |
| processor_kwargs = { | |
| "text": prompts, | |
| "padding": True, | |
| "padding_side": "right", | |
| "return_tensors": "pt", | |
| } | |
| if not is_t2i: | |
| processor_kwargs["images"] = condition_pil_list | |
| model_inputs = self.processor(**processor_kwargs).to(device) | |
| forward_kwargs = { | |
| "input_ids": model_inputs.input_ids, | |
| "attention_mask": model_inputs.attention_mask, | |
| "output_hidden_states": True, | |
| } | |
| if not is_t2i and hasattr(model_inputs, "pixel_values"): | |
| forward_kwargs["pixel_values"] = model_inputs.pixel_values.to(self.text_encoder.dtype) | |
| forward_kwargs["image_grid_thw"] = model_inputs.image_grid_thw | |
| if hasattr(model_inputs, "mm_token_type_ids"): | |
| forward_kwargs["mm_token_type_ids"] = model_inputs.mm_token_type_ids | |
| outputs = self.text_encoder(**forward_kwargs) | |
| hidden = outputs.hidden_states | |
| # (B, K, L, d) -> (B, L, K*d): exactly the flat input the adapter was trained on. | |
| x = torch.stack([hidden[i] for i in self.student_layers], dim=1) | |
| x = x.permute(0, 2, 1, 3).reshape(x.shape[0], x.shape[2], -1) | |
| return x, model_inputs.attention_mask.bool(), (model_inputs.input_ids == self._img_token_id) | |
| def _encode_vae_image(self, image: torch.Tensor, generator: torch.Generator): | |
| """The VAE is fp32 while condition latents must be in the model dtype (fp16). | |
| Without the round trip, `torch.cat([input_images_latents, latents])` in `__call__` promotes | |
| the input to fp32 and the DiT fails on a dtype mismatch. The encoder itself runs in fp32 | |
| (more accurate), the result is handed back in fp16. | |
| """ | |
| latents = super()._encode_vae_image(image.to(self.vae.dtype), generator) | |
| return latents.to(next(self.transformer.parameters()).dtype) | |
| def from_pretrained(cls, root=".", dtype=torch.float16, **kwargs): | |
| """Load the pipeline from a local model folder (the layout shipped in this repo). | |
| For a Hub repo use `DiffusionPipeline.from_pretrained(repo_id, trust_remote_code=True)`, | |
| which resolves this class through `model_index.json` and loads the components itself. | |
| """ | |
| from diffusers import AutoencoderKLQwenImage21, FlowMatchEulerDiscreteScheduler | |
| from transformers import AutoProcessor, AutoTokenizer, Qwen3_5ForConditionalGeneration | |
| if not os.path.isdir(root): | |
| raise ValueError(f"expected a local model folder with the components, got {root!r}") | |
| transformer = QwenImage21FusionTransformer2DModel.from_pretrained( | |
| os.path.join(root, "transformer"), torch_dtype=dtype | |
| ) | |
| text_encoder = Qwen3_5ForConditionalGeneration.from_pretrained( | |
| os.path.join(root, "text_encoder"), dtype=dtype | |
| ) | |
| # The VAE stays fp32 (1.3 GB): it is the component that misbehaves in fp16. | |
| vae = AutoencoderKLQwenImage21.from_pretrained(os.path.join(root, "vae"), torch_dtype=torch.float32) | |
| processor = AutoProcessor.from_pretrained(os.path.join(root, "processor")) | |
| tokenizer = AutoTokenizer.from_pretrained(os.path.join(root, "tokenizer")) | |
| scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained(os.path.join(root, "scheduler")) | |
| return cls( | |
| scheduler=scheduler, | |
| vae=vae, | |
| text_encoder=text_encoder, | |
| processor=processor, | |
| transformer=transformer, | |
| tokenizer=tokenizer, | |
| **kwargs, | |
| ) | |