zen-image-edit / transformer.py
recoilme's picture
Zen Image Edit: Qwen-Image-2.1 on Qwen3.5-0.8B with a text-fusion adapter inside the DiT
3a93d0e
Raw
History Blame
12.1 kB
"""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))
@property
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,
)