File size: 12,083 Bytes
3a93d0e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
"""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,
        )