"""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, )