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
| """Qwen-Image-2.1 DiT with the text adapter built in: `text_fusion` right before `txt_in`. | |
| The adapter (`image21-08b-text-encoder-adapter` v11) is a regression from Qwen3.5-0.8B hidden | |
| states to the input of the native `txt_in = QwenImage21TextProjection`. It therefore sits | |
| *before* `txt_in`, exactly like `Krea2TextFusion` sits before `Krea2TextProjection` in Krea 2: | |
| student slices (B, L, K*d) -> text_fusion -> (B, L, 4096) -> txt_in -> joint stream | |
| y = MLP(x) + Attn(x) + Mixer(x) | |
| | | +-- attention over the slice axis (K student layers) | |
| | +------------ self-attention over tokens, key padding mask | |
| +---------------------- ln-per-layer over slices (norms grow ~40x with depth) | |
| Two details matter and are easy to get wrong: | |
| * The 14 system-prompt tokens are sliced off **after** the fusion, not before. The adapter was | |
| trained on the full sequence (`drop_first=0`) and its attention branch indexes a learned | |
| position table by absolute token index, so cropping first would shift every position. | |
| * The config key is named `text_fusion_config`, not `text_fusion`: `ModelMixin.__getattr__` | |
| returns a config value before `nn.Module` can hand back a submodule, so a same-named key made | |
| `model.text_fusion` a `dict` and weight loading failed. | |
| """ | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| from diffusers.models.transformers.transformer_qwenimage21 import ( | |
| QwenImage21Transformer2DModel, | |
| ) | |
| class PerLayerNorm(nn.Module): | |
| """Normalize every slice separately, then the concatenation. | |
| Without per-slice normalization the early layers barely reach the output: hidden-state norms | |
| grow roughly 40x from layer 2 to layer 27. | |
| """ | |
| def __init__(self, in_dim, n_slices): | |
| super().__init__() | |
| assert in_dim % n_slices == 0, f"{in_dim} is not divisible by {n_slices}" | |
| self.n, self.d = n_slices, in_dim // n_slices | |
| self.per = nn.LayerNorm(self.d, elementwise_affine=False) | |
| self.all = nn.LayerNorm(in_dim) | |
| def forward(self, x): | |
| b, l, _ = x.shape | |
| return self.all(self.per(x.view(b, l, self.n, self.d)).reshape(b, l, -1)) | |
| class AttnBlock(nn.Module): | |
| """Pre-norm self-attention + FFN. No causality: text is not autoregressive and the DiT | |
| already sees the whole sequence at once.""" | |
| def __init__(self, d_model, n_heads, ffn_mult=4): | |
| super().__init__() | |
| self.n1 = nn.LayerNorm(d_model) | |
| self.attn = nn.MultiheadAttention(d_model, n_heads, batch_first=True) | |
| self.n2 = nn.LayerNorm(d_model) | |
| self.ffn = nn.Sequential(nn.Linear(d_model, d_model * ffn_mult), | |
| nn.GELU(approximate="tanh"), | |
| nn.Linear(d_model * ffn_mult, d_model)) | |
| def forward(self, x, key_padding_mask=None): | |
| h = self.n1(x) | |
| x = x + self.attn(h, h, h, need_weights=False, key_padding_mask=key_padding_mask)[0] | |
| return x + self.ffn(self.n2(x)) | |
| class AttnMixer(nn.Module): | |
| """Token branch: norm -> project to d_model -> N blocks -> project to out_dim. | |
| The last projection is zero-initialized, so at step 0 the branch adds nothing. | |
| """ | |
| def __init__(self, in_dim, out_dim, d_model=1024, n_heads=8, blocks=2, max_len=256): | |
| super().__init__() | |
| self.norm = nn.LayerNorm(in_dim) | |
| self.inp = nn.Linear(in_dim, d_model) | |
| self.pos = nn.Parameter(torch.zeros(1, max_len, d_model)) | |
| self.blocks = nn.ModuleList([AttnBlock(d_model, n_heads) for _ in range(blocks)]) | |
| self.out = nn.Linear(d_model, out_dim) | |
| nn.init.zeros_(self.out.weight) | |
| nn.init.zeros_(self.out.bias) | |
| def forward(self, x, key_padding_mask=None): | |
| h = self.inp(self.norm(x)) | |
| n = h.shape[1] | |
| pos = self.pos | |
| if n > pos.shape[1]: | |
| # The position table was trained on `max_len` slots; the tail is padded with zeros | |
| # (long sequences simply carry no positional signal there, but nothing crashes). | |
| pos = F.pad(pos, (0, 0, 0, n - pos.shape[1])) | |
| h = h + pos[:, :n] | |
| for block in self.blocks: | |
| h = block(h, key_padding_mask) | |
| return self.out(h) | |
| class SliceMixer(nn.Module): | |
| """Attention over the slice axis (the `layerwise_blocks` + `projector` part of Krea 2 fusion). | |
| Slices are normalized one by one and run as a sequence of length K through N self-attention | |
| blocks; `Linear(K -> 1)` then collapses the axis. The result is added to the MLP output | |
| (zero-initialized output projection). | |
| """ | |
| def __init__(self, n_slices, d, out_dim, n_heads=8, blocks=2, ffn_mult=2): | |
| super().__init__() | |
| self.n, self.d = n_slices, d | |
| self.per = nn.LayerNorm(d, elementwise_affine=False) | |
| self.blocks = nn.ModuleList([AttnBlock(d, n_heads, ffn_mult) for _ in range(blocks)]) | |
| self.proj = nn.Linear(n_slices, 1, bias=False) | |
| self.out = nn.Linear(d, out_dim) | |
| nn.init.zeros_(self.out.weight) | |
| nn.init.zeros_(self.out.bias) | |
| def forward(self, x): | |
| b, l, _ = x.shape | |
| h = self.per(x.view(b, l, self.n, self.d)).reshape(b * l, self.n, self.d) | |
| for block in self.blocks: | |
| h = block(h) | |
| h = self.proj(h.permute(0, 2, 1)).squeeze(-1) | |
| return self.out(h).reshape(b, l, -1) | |
| class AttnAdapter(nn.Module): | |
| """Point-wise MLP (`self.mlp`) plus residual branches over tokens (`self.attn`) and slices (`self.mixer`).""" | |
| def __init__(self, mods, attn=None, mixer=None): | |
| super().__init__() | |
| self.mlp = nn.Sequential(*mods) | |
| self.attn = attn | |
| self.mixer = mixer | |
| def forward(self, x, mask=None): | |
| """`mask`: `(B, L)` bool, True = real token. The attention branches get `key_padding_mask = ~mask`, | |
| otherwise attention would look into the padding.""" | |
| out = self.mlp(x) | |
| if self.attn is not None: | |
| kpm = None if mask is None else ~mask.bool() | |
| out = out + self.attn(x, kpm) | |
| if self.mixer is not None: | |
| out = out + self.mixer(x) | |
| return out | |
| def build_fusion(in_dim, out_dim, hidden=4096, proj_layers=2, norm="none", n_slices=1, | |
| attention=0, attn_dim=1024, attn_heads=8, max_len=256, | |
| mixer=0, mixer_heads=8, mixer_ffn=2, **_ignored): | |
| """Build the fusion block from the transformer config (`text_fusion_config`). | |
| `proj_layers` linear layers with GELU(tanh) between them; `attention`/`mixer` are the number of | |
| residual branches. Extra config keys (`student_layers`, `drop_idx`) are ignored here: the | |
| pipeline uses them, the block does not. | |
| """ | |
| if proj_layers < 2: | |
| raise ValueError("at least 2 linear layers are required") | |
| mods = [] | |
| if norm == "ln": | |
| mods.append(nn.LayerNorm(in_dim)) | |
| elif norm == "ln-per-layer": | |
| mods.append(PerLayerNorm(in_dim, n_slices)) | |
| elif norm == "rms": | |
| mods.append(nn.RMSNorm(in_dim)) | |
| elif norm != "none": | |
| raise ValueError(f"unknown norm: {norm}") | |
| mods += [nn.Linear(in_dim, hidden), nn.GELU(approximate="tanh")] | |
| for _ in range(proj_layers - 2): | |
| mods += [nn.Linear(hidden, hidden), nn.GELU(approximate="tanh")] | |
| mods.append(nn.Linear(hidden, out_dim)) | |
| attn = AttnMixer(in_dim, out_dim, attn_dim, attn_heads, attention, max_len) if attention else None | |
| mix = SliceMixer(n_slices, in_dim // n_slices, out_dim, mixer_heads, mixer, mixer_ffn) if mixer else None | |
| if attn is not None or mix is not None: | |
| return AttnAdapter(mods, attn, mix) | |
| return nn.Sequential(*mods) | |
| class QwenImage21FusionTransformer2DModel(QwenImage21Transformer2DModel): | |
| """`QwenImage21Transformer2DModel` + `text_fusion` (student stack -> 4096 condition). | |
| The config carries an extra key `text_fusion_config` (arguments of `build_fusion` plus | |
| `student_layers` and `drop_idx` for the pipeline), so the checkpoint is self-contained: the | |
| block weights live in the same file under the `text_fusion.` prefix. | |
| The `__init__` signature intentionally repeats the parent's. `ConfigMixin.extract_init_dict` | |
| builds `init_dict` from named parameters only and ignores `**kwargs`, so forwarding the config | |
| through `**kwargs` would drop the parent keys and rebuild the model from defaults. | |
| """ | |
| _no_split_modules = QwenImage21Transformer2DModel._no_split_modules + ["AttnAdapter"] | |
| def __init__( | |
| self, | |
| patch_size: int = 1, | |
| in_channels: int = 64, | |
| out_channels: int | None = 64, | |
| num_layers: int = 32, | |
| attention_head_dim: int = 128, | |
| num_attention_heads: int = 32, | |
| context_in_dim: int = 4096, | |
| mlp_ratio: int = 3, | |
| axes_dims_rope: tuple[int, int, int] = (16, 56, 56), | |
| eps: float = 1e-6, | |
| causal_condition: bool = True, | |
| text_fusion_config: dict | None = None, | |
| ): | |
| super().__init__( | |
| patch_size=patch_size, | |
| in_channels=in_channels, | |
| out_channels=out_channels, | |
| num_layers=num_layers, | |
| attention_head_dim=attention_head_dim, | |
| num_attention_heads=num_attention_heads, | |
| context_in_dim=context_in_dim, | |
| mlp_ratio=mlp_ratio, | |
| axes_dims_rope=axes_dims_rope, | |
| eps=eps, | |
| causal_condition=causal_condition, | |
| ) | |
| if text_fusion_config is None: | |
| raise ValueError( | |
| "`text_fusion_config` is required in config.json: without it there is nowhere to " | |
| "attach the adapter, and `encoder_hidden_states` are expected to be context_in_dim" | |
| ) | |
| self.text_fusion = build_fusion(**text_fusion_config) | |
| self.register_to_config(text_fusion_config=dict(text_fusion_config)) | |
| self.condition_drop_idx = int(text_fusion_config.get("drop_idx", 0)) | |
| def text_fusion_dtype(self): | |
| return next(self.text_fusion.parameters()).dtype | |
| def _condition(self, encoder_hidden_states, encoder_hidden_states_mask, img_mask): | |
| """Student slices -> DiT condition: fusion over the full sequence, then crop the prefix. | |
| `img_mask` arrives from the pipeline already concatenated with the target-image slots | |
| (`append_target_slots`), so only the conditioning part is cropped — otherwise the target | |
| slots would slide out of place. | |
| """ | |
| length = encoder_hidden_states.shape[1] | |
| mask = None if encoder_hidden_states_mask is None else encoder_hidden_states_mask.bool() | |
| fused = self.text_fusion(encoder_hidden_states.to(self.text_fusion_dtype), mask) | |
| drop = self.condition_drop_idx | |
| if not drop: | |
| return fused, encoder_hidden_states_mask, img_mask | |
| fused = fused[:, drop:] | |
| if encoder_hidden_states_mask is not None: | |
| encoder_hidden_states_mask = encoder_hidden_states_mask[:, drop:] | |
| img_mask = torch.cat([img_mask[:, drop:length], img_mask[:, length:]], dim=1) | |
| return fused, encoder_hidden_states_mask, img_mask | |
| def forward( | |
| self, | |
| hidden_states, | |
| encoder_hidden_states, | |
| timestep, | |
| img_shapes, | |
| img_mask, | |
| encoder_hidden_states_mask=None, | |
| **kwargs, | |
| ): | |
| """`encoder_hidden_states` here is the student stack `(B, L, K*d)`, not a ready condition. | |
| The fusion output is cast back to the latent dtype: the block is stored in fp16 while | |
| `txt_in` and the transformer blocks expect the model dtype. | |
| """ | |
| fused, encoder_hidden_states_mask, img_mask = self._condition( | |
| encoder_hidden_states, encoder_hidden_states_mask, img_mask | |
| ) | |
| return super().forward( | |
| hidden_states=hidden_states, | |
| encoder_hidden_states=fused.to(hidden_states.dtype), | |
| timestep=timestep, | |
| img_shapes=img_shapes, | |
| img_mask=img_mask, | |
| encoder_hidden_states_mask=encoder_hidden_states_mask, | |
| **kwargs, | |
| ) | |