"""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 (`` 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 from transformer import 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__() 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<|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)) tokenizer = tokenizer if tokenizer is not None else getattr(processor, "tokenizer", None) 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 = "<|vision_start|><|image_pad|><|vision_end|>" for i in range(2, len(image) + 1): replace += f" <|vision_start|><|image_pad|><|vision_end|>" template = self.prompt_template_ti2i.replace( "<|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) @classmethod def from_pretrained(cls, root=".", dtype=torch.float16, **kwargs): """Load the pipeline from a model folder (the layout shipped in this repo).""" from diffusers import AutoencoderKLQwenImage21, FlowMatchEulerDiscreteScheduler from transformers import AutoProcessor, AutoTokenizer, Qwen3_5ForConditionalGeneration if not os.path.isdir(root): raise ValueError(f"expected a 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, )