ajh-code commited on
Commit
0d00722
·
verified ·
1 Parent(s): f864990

Add vendor/mage_flow/models/modules/mage_vae.py

Browse files
vendor/mage_flow/models/modules/mage_vae.py ADDED
@@ -0,0 +1,651 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ MageVAE: DConvEncoder + DConvDenoiser (with CoD Decoder) wrapper.
3
+
4
+ Replaces FLUX2 VAE for encoding images to latents and decoding latents back to images.
5
+ Supports only the kl0.1 CoD ckpt layout:
6
+ encoder weights: 'state_dict' → 'student.dconv_encoder.*' (packed mean+logvar, out_ch_mult=2)
7
+ decoder weights: 'state_dict' → 'pipeline.*' (denoiser + y_embedder.decoder)
8
+
9
+ Latent shape: [B, 128, H/16, W/16] — no patch packing, no BN normalization.
10
+ """
11
+
12
+ import math
13
+ import os
14
+ from functools import lru_cache
15
+
16
+ import torch
17
+ import torch.nn as nn
18
+ import torch.nn.functional as F
19
+ from loguru import logger
20
+
21
+
22
+ # ---------------------------------------------------------------------------
23
+ # Primitive layers (vendored from GenCodec, inference subset)
24
+ # ---------------------------------------------------------------------------
25
+ def nonlinearity(x):
26
+ return x * torch.sigmoid(x)
27
+
28
+
29
+ def Normalize(in_channels):
30
+ return torch.nn.GroupNorm(num_groups=32, num_channels=in_channels, eps=1e-6, affine=True)
31
+
32
+
33
+ def modulate(x, shift, scale):
34
+ if x.dim() == 4:
35
+ b, c = x.shape[:2]
36
+ return x * (1 + scale.view(b, c, 1, 1)) + shift.view(b, c, 1, 1)
37
+ return x * (1 + scale.unsqueeze(1)) + shift.unsqueeze(1)
38
+
39
+
40
+ class LayerNorm2d(nn.LayerNorm):
41
+ def __init__(self, num_channels, eps=1e-6, affine=True):
42
+ super().__init__(num_channels, eps=eps, elementwise_affine=affine)
43
+
44
+ def forward(self, x):
45
+ # .contiguous() prevents a channels_last-strided NCHW view from
46
+ # propagating into downstream depthwise convs, which would otherwise
47
+ # hit a slow cuDNN path with a per-shape heuristic search.
48
+ x = x.permute(0, 2, 3, 1).contiguous()
49
+ x = F.layer_norm(x, self.normalized_shape, self.weight, self.bias, self.eps)
50
+ return x.permute(0, 3, 1, 2).contiguous()
51
+
52
+
53
+ class _EncoderLayerNorm2d(LayerNorm2d):
54
+ pass
55
+
56
+
57
+ class RMSNorm(nn.Module):
58
+ def __init__(self, hidden_size, eps=1e-6):
59
+ super().__init__()
60
+ self.weight = nn.Parameter(torch.ones(hidden_size))
61
+ self.variance_epsilon = eps
62
+
63
+ def forward(self, x):
64
+ in_dtype = x.dtype
65
+ x = x.to(torch.float32)
66
+ var = x.pow(2).mean(-1, keepdim=True)
67
+ x = x * torch.rsqrt(var + self.variance_epsilon)
68
+ return self.weight * x.to(in_dtype)
69
+
70
+
71
+ class TimestepEmbedder(nn.Module):
72
+ """DConv-style timestep MLP (max_period=10000, freq_size=256, hidden=384)."""
73
+
74
+ def __init__(self, hidden_size, frequency_embedding_size=256):
75
+ super().__init__()
76
+ self.mlp = nn.Sequential(
77
+ nn.Linear(frequency_embedding_size, hidden_size, bias=True),
78
+ nn.SiLU(),
79
+ nn.Linear(hidden_size, hidden_size, bias=True),
80
+ )
81
+ self.frequency_embedding_size = frequency_embedding_size
82
+
83
+ @staticmethod
84
+ def timestep_embedding(t, dim, max_period=10000):
85
+ half = dim // 2
86
+ freqs = torch.exp(
87
+ -math.log(max_period) * torch.arange(0, half, dtype=torch.float32) / half
88
+ ).to(t.device)
89
+ args = t[:, None].float() * freqs[None]
90
+ emb = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
91
+ if dim % 2:
92
+ emb = torch.cat([emb, torch.zeros_like(emb[:, :1])], dim=-1)
93
+ return emb
94
+
95
+ def forward(self, t):
96
+ emb = self.timestep_embedding(t, self.frequency_embedding_size)
97
+ return self.mlp(emb.to(self.mlp[0].weight.dtype))
98
+
99
+
100
+ class BottleneckPatchEmbed(nn.Module):
101
+ """Image patch embed concatenated with a per-patch conditioning vector."""
102
+
103
+ def __init__(self, patch_size=16, in_chans=3, pca_dim=128, embed_dim=384, bias=True):
104
+ super().__init__()
105
+ self.proj1 = nn.Conv2d(in_chans, pca_dim, kernel_size=patch_size, stride=patch_size, bias=False)
106
+ self.proj2 = nn.Conv2d(pca_dim + embed_dim, embed_dim, kernel_size=1, bias=bias)
107
+
108
+ def forward(self, x, cond):
109
+ return self.proj2(torch.cat([self.proj1(x), cond], dim=1))
110
+
111
+
112
+ class DiCoBlock(nn.Module):
113
+ """DConv block with adaLN modulation."""
114
+
115
+ def __init__(self, hidden_size, mlp_ratio=4.0):
116
+ super().__init__()
117
+ self.conv1 = nn.Conv2d(hidden_size, hidden_size, 1, bias=True)
118
+ self.conv2 = nn.Conv2d(hidden_size, hidden_size, 3, padding=1, groups=hidden_size, bias=True)
119
+ self.conv3 = nn.Conv2d(hidden_size, hidden_size, 1, bias=True)
120
+
121
+ self.ca = nn.Sequential(
122
+ nn.AdaptiveAvgPool2d(1),
123
+ nn.Conv2d(hidden_size, hidden_size, 1, bias=True),
124
+ nn.Sigmoid(),
125
+ )
126
+
127
+ ffn = int(mlp_ratio * hidden_size)
128
+ self.conv4 = nn.Conv2d(hidden_size, ffn, 1, bias=True)
129
+ self.conv5 = nn.Conv2d(ffn, hidden_size, 1, bias=True)
130
+
131
+ self.norm1 = LayerNorm2d(hidden_size, affine=False)
132
+ self.norm2 = LayerNorm2d(hidden_size, affine=False)
133
+
134
+ self.adaLN_modulation = nn.Sequential(
135
+ nn.SiLU(),
136
+ nn.Linear(hidden_size, 6 * hidden_size, bias=True),
137
+ )
138
+
139
+ def forward(self, inp, c):
140
+ shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.adaLN_modulation(c).chunk(6, dim=1)
141
+ x = modulate(self.norm1(inp), shift_msa, scale_msa)
142
+ x = F.gelu(self.conv2(self.conv1(x)))
143
+ x = x * self.ca(x)
144
+ x = self.conv3(x)
145
+ x = inp + gate_msa[..., None, None] * x
146
+ x = x + gate_mlp[..., None, None] * self.conv5(
147
+ F.gelu(self.conv4(modulate(self.norm2(x), shift_mlp, scale_mlp)))
148
+ )
149
+ return x
150
+
151
+
152
+ class _EncoderDiCoBlock(nn.Module):
153
+ """DiCoBlock without adaLN, for the encoder pathway."""
154
+
155
+ def __init__(self, hidden_size, mlp_ratio=4.0):
156
+ super().__init__()
157
+ self.conv1 = nn.Conv2d(hidden_size, hidden_size, 1, bias=True)
158
+ self.conv2 = nn.Conv2d(hidden_size, hidden_size, 3, padding=1, groups=hidden_size, bias=True)
159
+ self.conv3 = nn.Conv2d(hidden_size, hidden_size, 1, bias=True)
160
+ self.ca = nn.Sequential(
161
+ nn.AdaptiveAvgPool2d(1),
162
+ nn.Conv2d(hidden_size, hidden_size, 1, bias=True),
163
+ nn.Sigmoid(),
164
+ )
165
+ ffn = int(mlp_ratio * hidden_size)
166
+ self.conv4 = nn.Conv2d(hidden_size, ffn, 1, bias=True)
167
+ self.conv5 = nn.Conv2d(ffn, hidden_size, 1, bias=True)
168
+ self.norm1 = _EncoderLayerNorm2d(hidden_size)
169
+ self.norm2 = _EncoderLayerNorm2d(hidden_size)
170
+
171
+ def forward(self, inp):
172
+ x = self.norm1(inp)
173
+ x = F.gelu(self.conv2(self.conv1(x)))
174
+ x = x * self.ca(x)
175
+ x = self.conv3(x)
176
+ x = inp + x
177
+ return x + self.conv5(F.gelu(self.conv4(self.norm2(x))))
178
+
179
+
180
+ class NerfEmbedder(nn.Module):
181
+ """Patch-position embedder used by the DConv decoder x-pathway."""
182
+
183
+ def __init__(self, in_channels, hidden_size_input, max_freqs=8):
184
+ super().__init__()
185
+ self.max_freqs = max_freqs
186
+ self.embedder = nn.Sequential(
187
+ nn.Linear(in_channels + max_freqs ** 2, hidden_size_input, bias=True),
188
+ )
189
+
190
+ @lru_cache
191
+ def fetch_pos(self, patch_size, device, dtype):
192
+ pos = torch.linspace(0, 1, patch_size, device=device, dtype=dtype)
193
+ pos_y, pos_x = torch.meshgrid(pos, pos, indexing="ij")
194
+ pos_x = pos_x.reshape(-1, 1, 1)
195
+ pos_y = pos_y.reshape(-1, 1, 1)
196
+ freqs = torch.linspace(0, self.max_freqs, self.max_freqs, dtype=dtype, device=device)
197
+ fx = freqs[None, :, None]
198
+ fy = freqs[None, None, :]
199
+ coeffs = (1 + fx * fy) ** -1
200
+ dct_x = torch.cos(pos_x * fx * torch.pi)
201
+ dct_y = torch.cos(pos_y * fy * torch.pi)
202
+ return (dct_x * dct_y * coeffs).view(1, -1, self.max_freqs ** 2)
203
+
204
+ def forward(self, x):
205
+ B, P2, _ = x.shape
206
+ ps = int(P2 ** 0.5)
207
+ dct = self.fetch_pos(ps, x.device, x.dtype).expand(B, -1, -1)
208
+ return self.embedder(torch.cat([x, dct], dim=-1))
209
+
210
+
211
+ class NerfFinalLayer(nn.Module):
212
+ def __init__(self, hidden_size, out_channels):
213
+ super().__init__()
214
+ self.norm = RMSNorm(hidden_size)
215
+ self.linear = nn.Linear(hidden_size, out_channels, bias=True)
216
+
217
+ def forward(self, x):
218
+ return self.linear(self.norm(x))
219
+
220
+
221
+ class SimpleMLPAdaLN(nn.Module):
222
+ """Final small MLP that maps NerfEmbedder features to per-patch RGB."""
223
+
224
+ def __init__(self, in_channels, model_channels, out_channels, z_channels, num_res_blocks, patch_size):
225
+ super().__init__()
226
+ self.in_channels = in_channels
227
+ self.model_channels = model_channels
228
+ self.out_channels = out_channels
229
+ self.num_res_blocks = num_res_blocks
230
+ self.patch_size = patch_size
231
+
232
+ self.cond_embed = nn.Linear(z_channels, patch_size ** 2 * model_channels)
233
+ self.input_proj = nn.Linear(in_channels, model_channels)
234
+
235
+ self.res_blocks = nn.ModuleList(_MLPResBlock(model_channels) for _ in range(num_res_blocks))
236
+
237
+ def forward(self, x, c):
238
+ x = self.input_proj(x)
239
+ c = self.cond_embed(c).reshape(c.shape[0], self.patch_size ** 2, -1)
240
+ for block in self.res_blocks:
241
+ x = block(x, c)
242
+ return x
243
+
244
+
245
+ class _MLPResBlock(nn.Module):
246
+ def __init__(self, channels):
247
+ super().__init__()
248
+ self.in_ln = nn.LayerNorm(channels, eps=1e-6)
249
+ self.mlp = nn.Sequential(
250
+ nn.Linear(channels, channels, bias=True),
251
+ nn.SiLU(),
252
+ nn.Linear(channels, channels, bias=True),
253
+ )
254
+ self.adaLN_modulation = nn.Sequential(
255
+ nn.SiLU(),
256
+ nn.Linear(channels, 3 * channels, bias=True),
257
+ )
258
+
259
+ def forward(self, x, y):
260
+ shift, scale, gate = self.adaLN_modulation(y).chunk(3, dim=-1)
261
+ h = self.in_ln(x) * (1 + scale) + shift
262
+ return x + gate * self.mlp(h)
263
+
264
+
265
+ class ResnetBlock(nn.Module):
266
+ """GroupNorm + Conv ResBlock used by the CoD Decoder."""
267
+
268
+ def __init__(self, *, in_channels, out_channels=None, dropout=0.0):
269
+ super().__init__()
270
+ out_channels = out_channels or in_channels
271
+ self.in_channels = in_channels
272
+ self.out_channels = out_channels
273
+
274
+ self.norm1 = Normalize(in_channels)
275
+ self.conv1 = nn.Conv2d(in_channels, out_channels, 3, padding=1)
276
+ self.norm2 = Normalize(out_channels)
277
+ self.dropout = nn.Dropout(dropout)
278
+ self.conv2 = nn.Conv2d(out_channels, out_channels, 3, padding=1)
279
+ if in_channels != out_channels:
280
+ self.nin_shortcut = nn.Conv2d(in_channels, out_channels, 1)
281
+
282
+ def forward(self, x):
283
+ h = self.conv1(nonlinearity(self.norm1(x)))
284
+ h = self.conv2(self.dropout(nonlinearity(self.norm2(h))))
285
+ if self.in_channels != self.out_channels:
286
+ x = self.nin_shortcut(x)
287
+ return x + h
288
+
289
+
290
+ class AttnBlock(nn.Module):
291
+ """Patched self-attention used at inference (eval mode of the original)."""
292
+
293
+ def __init__(self, in_channels, patch_size=32):
294
+ super().__init__()
295
+ self.in_channels = in_channels
296
+ self.patch_size = patch_size
297
+ self.norm = Normalize(in_channels)
298
+ self.q = nn.Conv2d(in_channels, in_channels, 1)
299
+ self.k = nn.Conv2d(in_channels, in_channels, 1)
300
+ self.v = nn.Conv2d(in_channels, in_channels, 1)
301
+ self.proj_out = nn.Conv2d(in_channels, in_channels, 1)
302
+
303
+ def forward(self, x):
304
+ h_ = self.norm(x)
305
+ Q = self.q(h_)
306
+ K = self.k(h_)
307
+ V = self.v(h_)
308
+
309
+ d = self.patch_size
310
+ b, c, H, W = Q.shape
311
+ pad_h = (d - H % d) % d
312
+ pad_w = (d - W % d) % d
313
+ if pad_h or pad_w:
314
+ Q = F.pad(Q, (0, pad_w, 0, pad_h), mode="replicate")
315
+ K = F.pad(K, (0, pad_w, 0, pad_h), mode="replicate")
316
+ V = F.pad(V, (0, pad_w, 0, pad_h), mode="replicate")
317
+ _, _, H_pad, W_pad = Q.shape
318
+ nph, npw = H_pad // d, W_pad // d
319
+ np_ = nph * npw
320
+
321
+ def to_patches(t):
322
+ return (t.reshape(b, c, nph, d, npw, d)
323
+ .permute(0, 2, 4, 1, 3, 5)
324
+ .reshape(b * np_, c, d * d))
325
+
326
+ Q = to_patches(Q)
327
+ K = to_patches(K)
328
+ V = to_patches(V)
329
+
330
+ w_ = torch.bmm(Q.permute(0, 2, 1), K) * (c ** -0.5)
331
+ w_ = F.softmax(w_, dim=2).permute(0, 2, 1)
332
+ h_ = torch.bmm(V, w_).reshape(b, nph, npw, c, d, d).permute(0, 3, 1, 4, 2, 5).reshape(b, c, H_pad, W_pad)
333
+ if pad_h or pad_w:
334
+ h_ = h_[:, :, :H, :W]
335
+ return x + self.proj_out(h_)
336
+
337
+
338
+ # ---------------------------------------------------------------------------
339
+ # adaLN constant-folding: at fixed t=0, adaLN_modulation(c) is constant.
340
+ # Replace the MLP with a buffer so DiCoBlock.forward stays unchanged and
341
+ # torch.compile can fuse the surrounding ops normally.
342
+ # ---------------------------------------------------------------------------
343
+ class _ConstAdaLN(nn.Module):
344
+ def __init__(self, modulation: torch.Tensor):
345
+ super().__init__()
346
+ self.register_buffer("modulation", modulation.detach().clone())
347
+
348
+ def forward(self, c):
349
+ b = c.shape[0]
350
+ if self.modulation.shape[0] != b:
351
+ return self.modulation.expand(b, *self.modulation.shape[1:])
352
+ return self.modulation
353
+
354
+
355
+ def _replace_adaln_with_const(module: nn.Module, c: torch.Tensor) -> int:
356
+ # Only DiCoBlock is targeted: its adaLN is conditioned solely on t.
357
+ # Other adaLN_modulation submodules (e.g. _MLPResBlock in the decoder MLP)
358
+ # take a per-position latent and must not be folded.
359
+ n = 0
360
+ for child in module.modules():
361
+ if not isinstance(child, DiCoBlock):
362
+ continue
363
+ adaln = child.adaLN_modulation
364
+ if isinstance(adaln, _ConstAdaLN):
365
+ continue
366
+ with torch.no_grad():
367
+ mod = adaln(c)
368
+ child.adaLN_modulation = _ConstAdaLN(mod)
369
+ n += 1
370
+ return n
371
+
372
+
373
+ # ---------------------------------------------------------------------------
374
+ # CoD Decoder: latent → conditioning features for the denoiser
375
+ # ---------------------------------------------------------------------------
376
+ class _Decoder(nn.Module):
377
+ """ds=16, up2x=True, light=True only."""
378
+
379
+ def __init__(self, out_ch=384, z_ch=128):
380
+ super().__init__()
381
+ self.conv_in = nn.Conv2d(z_ch, out_ch, kernel_size=3, stride=1, padding=1)
382
+ self.block = nn.Sequential(
383
+ ResnetBlock(in_channels=out_ch, out_channels=out_ch),
384
+ AttnBlock(out_ch, patch_size=32),
385
+ ResnetBlock(in_channels=out_ch, out_channels=out_ch),
386
+ AttnBlock(out_ch, patch_size=32),
387
+ ResnetBlock(in_channels=out_ch, out_channels=out_ch),
388
+ )
389
+ self.norm_out = Normalize(out_ch)
390
+ self.conv_out = nn.Conv2d(out_ch, out_ch, kernel_size=3, stride=1, padding=1)
391
+ self.ada = nn.Identity()
392
+
393
+ def forward(self, z):
394
+ h = self.block(self.conv_in(z))
395
+ h = self.conv_out(nonlinearity(self.norm_out(h)))
396
+ return self.ada(h)
397
+
398
+
399
+ # ---------------------------------------------------------------------------
400
+ # DConvEncoder: image → packed (mean, logvar) latent
401
+ # ---------------------------------------------------------------------------
402
+ class _DConvEncoder(nn.Module):
403
+ def __init__(
404
+ self,
405
+ z_ch=128,
406
+ hidden_size=384,
407
+ num_blocks=21,
408
+ patch_size=16,
409
+ mlp_ratio=4.0,
410
+ head_size=768,
411
+ num_head_blocks=2,
412
+ out_ch_mult=2,
413
+ ):
414
+ super().__init__()
415
+ self.z_ch = z_ch
416
+ self.patch_size = patch_size
417
+ self.patch_cond_embed = nn.Conv2d(3, head_size, kernel_size=patch_size, stride=patch_size, bias=True)
418
+ self.head_blocks = nn.ModuleList([
419
+ _EncoderDiCoBlock(head_size, mlp_ratio=mlp_ratio) for _ in range(num_head_blocks)
420
+ ])
421
+ self.proj_down = nn.Conv2d(head_size, hidden_size, kernel_size=1, bias=True)
422
+ self.z_proj = nn.Conv2d(z_ch, hidden_size, kernel_size=1, bias=True)
423
+ self.fuse_proj = nn.Conv2d(hidden_size * 2, hidden_size, kernel_size=1, bias=True)
424
+ self.t_embedder = TimestepEmbedder(hidden_size)
425
+ self.blocks = nn.ModuleList([
426
+ DiCoBlock(hidden_size, mlp_ratio=mlp_ratio) for _ in range(num_blocks)
427
+ ])
428
+ self.norm_out = LayerNorm2d(hidden_size)
429
+ self.proj_out = nn.Conv2d(hidden_size, z_ch * out_ch_mult, kernel_size=1, bias=True)
430
+
431
+ def forward_pred(self, z_t, t, y):
432
+ cond = self.patch_cond_embed(y)
433
+ for block in self.head_blocks:
434
+ cond = block(cond)
435
+ cond = self.proj_down(cond)
436
+
437
+ s = self.fuse_proj(torch.cat([cond, self.z_proj(z_t)], dim=1))
438
+ c = self.t_embedder(t.view(-1))
439
+ for block in self.blocks:
440
+ s = block(s, c)
441
+ return self.proj_out(self.norm_out(s))
442
+
443
+
444
+ # ---------------------------------------------------------------------------
445
+ # DConv denoiser: latent (via cond) + zero noise → reconstructed image
446
+ # ---------------------------------------------------------------------------
447
+ class _YEmbedder(nn.Module):
448
+ """Holds only the CoD decoder; the original Flux2 VAE encoder side is omitted."""
449
+
450
+ def __init__(self, ch=384, z_ch=128):
451
+ super().__init__()
452
+ self.decoder = _Decoder(out_ch=ch, z_ch=z_ch)
453
+
454
+
455
+ class _DConvDenoiser(nn.Module):
456
+ def __init__(
457
+ self,
458
+ patch_size=16,
459
+ in_channels=3,
460
+ hidden_size=384,
461
+ hidden_size_x=32,
462
+ mlp_ratio=4.0,
463
+ num_blocks=24,
464
+ num_cond_blocks=21,
465
+ bottleneck_dim=128,
466
+ ):
467
+ super().__init__()
468
+ self.in_channels = in_channels
469
+ self.patch_size = patch_size
470
+ self.hidden_size = hidden_size
471
+ self.num_cond_blocks = num_cond_blocks
472
+
473
+ self.t_embedder = TimestepEmbedder(hidden_size)
474
+ self.y_embedder_x = nn.Conv2d(hidden_size, hidden_size_x * patch_size ** 2, 1, 1, 0)
475
+ self.x_embedder = NerfEmbedder(in_channels + hidden_size_x, hidden_size_x, max_freqs=8)
476
+ self.s_embedder = BottleneckPatchEmbed(patch_size, in_channels, bottleneck_dim, hidden_size, bias=True)
477
+ self.blocks = nn.ModuleList([
478
+ DiCoBlock(hidden_size, mlp_ratio=mlp_ratio) for _ in range(num_cond_blocks)
479
+ ])
480
+ self.dec_net = SimpleMLPAdaLN(
481
+ in_channels=hidden_size_x,
482
+ model_channels=hidden_size_x,
483
+ out_channels=in_channels,
484
+ z_channels=hidden_size,
485
+ num_res_blocks=num_blocks - num_cond_blocks,
486
+ patch_size=patch_size,
487
+ )
488
+ self.final_layer = NerfFinalLayer(hidden_size_x, in_channels)
489
+ self.y_embedder = _YEmbedder(ch=hidden_size, z_ch=bottleneck_dim)
490
+
491
+ def forward(self, x, t, cond):
492
+ b, _, h, w = x.shape
493
+ c = self.t_embedder(t.view(-1))
494
+
495
+ s = self.s_embedder(x, cond)
496
+ for block in self.blocks:
497
+ s = block(s, c)
498
+
499
+ length = s.shape[-2] * s.shape[-1]
500
+ s = s.permute(0, 2, 3, 1).reshape(-1, self.hidden_size)
501
+
502
+ x = torch.nn.functional.unfold(x, kernel_size=self.patch_size, stride=self.patch_size)
503
+ x = torch.cat([x, self.y_embedder_x(cond).flatten(2)], dim=1)
504
+ x = x.reshape(b, -1, self.patch_size ** 2, length).permute(0, 3, 2, 1).flatten(0, 1)
505
+ x = self.x_embedder(x)
506
+
507
+ x = self.dec_net(x, s)
508
+ x = self.final_layer(x)
509
+ x = x.transpose(1, 2).reshape(b, length, -1)
510
+ return torch.nn.functional.fold(
511
+ x.transpose(1, 2).contiguous(), (h, w),
512
+ kernel_size=self.patch_size, stride=self.patch_size,
513
+ )
514
+
515
+
516
+ # ---------------------------------------------------------------------------
517
+ # Wrapper
518
+ # ---------------------------------------------------------------------------
519
+ def _load_state_dict(ckpt_path: str):
520
+ if ckpt_path.endswith(".safetensors"):
521
+ from safetensors.torch import load_file
522
+ return load_file(ckpt_path, device="cpu")
523
+ if os.path.exists(os.path.join(ckpt_path, "checkpoint-state_dict.pt")):
524
+ ckpt_path = os.path.join(ckpt_path, "checkpoint-state_dict.pt")
525
+ elif os.path.isdir(ckpt_path):
526
+ ckpt_path = os.path.join(ckpt_path, "checkpoint", "mp_rank_00_model_states.pt")
527
+ state = torch.load(ckpt_path, map_location="cpu")
528
+ if "module" in state:
529
+ return state["module"]
530
+ if "state_dict" in state:
531
+ return state["state_dict"]
532
+ return state
533
+
534
+
535
+ class MageVAE(nn.Module):
536
+ """
537
+ Encode: DConvEncoder (one-step diffusion) → latent [B, 128, H/16, W/16]
538
+ Decode: DConvDenoiser + CoD Decoder → image [B, 3, H, W] in [-1, 1]
539
+ """
540
+
541
+ latent_channels = 128
542
+ downsample_factor = 16
543
+
544
+ def __init__(self, ckpt_path: str, sample_posterior: bool = True):
545
+ super().__init__()
546
+ self.sample_posterior = sample_posterior
547
+
548
+ self.dconv_encoder = _DConvEncoder()
549
+ self.decoder_model = _DConvDenoiser()
550
+
551
+ sd = _load_state_dict(ckpt_path)
552
+ self._load_encoder(sd, ckpt_path)
553
+ self._load_decoder(sd, ckpt_path)
554
+
555
+ # adaLN modulation depends only on t, and we always run at t=0.
556
+ # Precompute and drop the MLPs once at construction (~37M params saved).
557
+ self._freeze_adaln_cache()
558
+
559
+ def _load_encoder(self, sd, ckpt_path):
560
+ prefix = "student.dconv_encoder."
561
+ enc_sd = {k[len(prefix):]: v for k, v in sd.items() if k.startswith(prefix)}
562
+ if not enc_sd:
563
+ raise RuntimeError(f"CoDEncoder: no '{prefix}*' keys in {ckpt_path}")
564
+ proj = enc_sd.get("proj_out.weight")
565
+ if proj is None or proj.shape[0] != 2 * self.latent_channels:
566
+ raise RuntimeError(
567
+ f"CoDEncoder: expected packed mean+logvar (proj_out out_channels="
568
+ f"{2 * self.latent_channels}), got {None if proj is None else tuple(proj.shape)}"
569
+ )
570
+ missing, unexpected = self.dconv_encoder.load_state_dict(enc_sd, strict=False)
571
+ logger.info(
572
+ f"CoDEncoder: loaded {len(enc_sd)} keys, "
573
+ f"missing={len(missing)}, unexpected={len(unexpected)}"
574
+ )
575
+ if missing:
576
+ logger.warning(f"CoDEncoder missing: {missing[:10]}")
577
+
578
+ def _load_decoder(self, sd, ckpt_path):
579
+ prefix = "pipeline."
580
+ if not any(k.startswith(prefix) for k in sd):
581
+ raise RuntimeError(f"CoDDecoder: no '{prefix}*' keys in {ckpt_path}")
582
+ model_dict = self.decoder_model.state_dict()
583
+ matched = {}
584
+ for k, v in sd.items():
585
+ if not k.startswith(prefix):
586
+ continue
587
+ new_k = k[len(prefix):]
588
+ if new_k.startswith("y_embedder.encoder.") or new_k.startswith("y_embedder.bottleneck."):
589
+ continue
590
+ if new_k in model_dict and model_dict[new_k].shape == v.shape:
591
+ matched[new_k] = v
592
+ self.decoder_model.load_state_dict(matched, strict=False)
593
+ logger.info(f"CoDDecoder: loaded {len(matched)} params (denoiser + y_embedder.decoder)")
594
+ if not matched:
595
+ raise RuntimeError(f"CoDDecoder: 0 params matched from {ckpt_path}")
596
+
597
+ @torch.no_grad()
598
+ def _moments(self, x: torch.Tensor):
599
+ B, _, H, W = x.shape
600
+ ps = self.dconv_encoder.patch_size
601
+ z_t = torch.zeros(B, self.dconv_encoder.z_ch, H // ps, W // ps, device=x.device, dtype=x.dtype)
602
+ t = torch.zeros(B, device=x.device, dtype=x.dtype)
603
+ out = self.dconv_encoder.forward_pred(z_t, t, x)
604
+ mean = out[:, : self.latent_channels]
605
+ logvar = out[:, self.latent_channels :].clamp(min=-20.0, max=10.0)
606
+ return mean, logvar
607
+
608
+ @torch.no_grad()
609
+ def _encode_moments(self, x: torch.Tensor):
610
+ # Compile target: pure deterministic part of encode (no RNG, no
611
+ # asserts), so torch.compile produces a single dynamic graph.
612
+ return self._moments(x)
613
+
614
+ @torch.no_grad()
615
+ def encode(self, x: torch.Tensor) -> torch.Tensor:
616
+ ps = self.dconv_encoder.patch_size
617
+ H, W = x.shape[-2], x.shape[-1]
618
+ if H % ps or W % ps:
619
+ raise ValueError(f"H, W must be multiples of {ps}, got ({H}, {W})")
620
+ mean, logvar = self._encode_moments(x)
621
+ if self.sample_posterior:
622
+ return mean + torch.exp(0.5 * logvar) * torch.randn_like(mean)
623
+ return mean
624
+
625
+ @torch.no_grad()
626
+ def decode(self, z: torch.Tensor) -> torch.Tensor:
627
+ cond = self.decoder_model.y_embedder.decoder(z)
628
+ B = z.shape[0]
629
+ H = z.shape[2] * self.downsample_factor
630
+ W = z.shape[3] * self.downsample_factor
631
+ noise = torch.zeros(B, 3, H, W, device=z.device, dtype=z.dtype)
632
+ t = torch.zeros(B, device=z.device, dtype=z.dtype)
633
+ return self.decoder_model.forward(noise, t, cond)
634
+
635
+ @property
636
+ def device(self):
637
+ return next(self.parameters()).device
638
+
639
+ @property
640
+ def dtype(self):
641
+ return next(self.parameters()).dtype
642
+
643
+ def _freeze_adaln_cache(self):
644
+ """Constant-fold adaLN_modulation MLPs at t=0 (encoder + decoder)."""
645
+ device = next(self.parameters()).device
646
+ dtype = next(self.parameters()).dtype
647
+ t = torch.zeros(1, device=device, dtype=dtype)
648
+ c_enc = self.dconv_encoder.t_embedder(t)
649
+ _replace_adaln_with_const(self.dconv_encoder, c_enc)
650
+ c_dec = self.decoder_model.t_embedder(t)
651
+ _replace_adaln_with_const(self.decoder_model, c_dec)