simit 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -0,0 +1,872 @@
1
+ # SPDX-License-Identifier: Apache-2.0
2
+ # SPDX-FileCopyrightText: Copyright contributors to the vLLM project (vllm-omni)
3
+ # Copied into simit with the vLLM logger replaced by stdlib logging.
4
+ #
5
+ # Ported from ByteDance Lance upstream (https://github.com/bytedance/Lance,
6
+ # modeling/vae/wan/{vae2_2.py,model.py}). Upstream copyright:
7
+ #
8
+ # Copyright (c) 2025 ByteDance Ltd. and/or its affiliates.
9
+ # Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
10
+ # Licensed under the Apache License, Version 2.0.
11
+ #
12
+ # Lance ships its VAE as a raw ``Wan2.2_VAE.pth`` whose state dict is keyed for
13
+ # the upstream ``WanVAE_`` nn.Module here (not for diffusers' ``AutoencoderKLWan``).
14
+ # We port that module verbatim so we can load the .pth directly, and wrap it with
15
+ # :class:`LanceWanVAE` which exposes BAGEL's 4-D ``encode(BCHW)/decode(BCHW)``
16
+ # surface (and a 5-D path for the video checkpoint).
17
+ """Wan2.2 VAE used by Lance, ported from upstream so ``Wan2.2_VAE.pth`` loads
18
+ natively without state-dict surgery."""
19
+
20
+ from __future__ import annotations
21
+
22
+ import torch
23
+ import torch.nn as nn
24
+ import torch.nn.functional as F
25
+ import logging
26
+
27
+ from einops import rearrange
28
+
29
+ logger = logging.getLogger(__name__)
30
+
31
+ CACHE_T = 2
32
+
33
+
34
+ class CausalConv3d(nn.Conv3d):
35
+ """Causal 3D conv with feature-map caching across temporal chunks."""
36
+
37
+ def __init__(self, *args, **kwargs):
38
+ super().__init__(*args, **kwargs)
39
+ self._padding = (
40
+ self.padding[2],
41
+ self.padding[2],
42
+ self.padding[1],
43
+ self.padding[1],
44
+ 2 * self.padding[0],
45
+ 0,
46
+ )
47
+ self.padding = (0, 0, 0)
48
+
49
+ def forward(self, x, cache_x=None):
50
+ padding = list(self._padding)
51
+ if cache_x is not None and self._padding[4] > 0:
52
+ cache_x = cache_x.to(x.device)
53
+ x = torch.cat([cache_x, x], dim=2)
54
+ padding[4] -= cache_x.shape[2]
55
+ x = F.pad(x, padding)
56
+ return super().forward(x)
57
+
58
+
59
+ class RMS_norm(nn.Module):
60
+ def __init__(self, dim, channel_first=True, images=True, bias=False):
61
+ super().__init__()
62
+ broadcastable_dims = (1, 1, 1) if not images else (1, 1)
63
+ shape = (dim, *broadcastable_dims) if channel_first else (dim,)
64
+ self.channel_first = channel_first
65
+ self.scale = dim**0.5
66
+ self.gamma = nn.Parameter(torch.ones(shape))
67
+ self.bias = nn.Parameter(torch.zeros(shape)) if bias else 0.0
68
+
69
+ def forward(self, x):
70
+ return F.normalize(x, dim=(1 if self.channel_first else -1)) * self.scale * self.gamma + self.bias
71
+
72
+
73
+ class Upsample(nn.Upsample):
74
+ def forward(self, x):
75
+ return super().forward(x.float()).type_as(x)
76
+
77
+
78
+ class Resample(nn.Module):
79
+ def __init__(self, dim, mode):
80
+ assert mode in ("none", "upsample2d", "upsample3d", "downsample2d", "downsample3d")
81
+ super().__init__()
82
+ self.dim = dim
83
+ self.mode = mode
84
+
85
+ if mode == "upsample2d":
86
+ self.resample = nn.Sequential(
87
+ Upsample(scale_factor=(2.0, 2.0), mode="nearest-exact"),
88
+ nn.Conv2d(dim, dim, 3, padding=1),
89
+ )
90
+ elif mode == "upsample3d":
91
+ self.resample = nn.Sequential(
92
+ Upsample(scale_factor=(2.0, 2.0), mode="nearest-exact"),
93
+ nn.Conv2d(dim, dim, 3, padding=1),
94
+ )
95
+ self.time_conv = CausalConv3d(dim, dim * 2, (3, 1, 1), padding=(1, 0, 0))
96
+ elif mode == "downsample2d":
97
+ self.resample = nn.Sequential(nn.ZeroPad2d((0, 1, 0, 1)), nn.Conv2d(dim, dim, 3, stride=(2, 2)))
98
+ elif mode == "downsample3d":
99
+ self.resample = nn.Sequential(nn.ZeroPad2d((0, 1, 0, 1)), nn.Conv2d(dim, dim, 3, stride=(2, 2)))
100
+ self.time_conv = CausalConv3d(dim, dim, (3, 1, 1), stride=(2, 1, 1), padding=(0, 0, 0))
101
+ else:
102
+ self.resample = nn.Identity()
103
+
104
+ def forward(self, x, feat_cache=None, feat_idx=[0]):
105
+ b, c, t, h, w = x.size()
106
+ if self.mode == "upsample3d":
107
+ if feat_cache is not None:
108
+ idx = feat_idx[0]
109
+ if feat_cache[idx] is None:
110
+ feat_cache[idx] = "Rep"
111
+ feat_idx[0] += 1
112
+ else:
113
+ cache_x = x[:, :, -CACHE_T:, :, :].clone()
114
+ if cache_x.shape[2] < 2 and feat_cache[idx] is not None and feat_cache[idx] != "Rep":
115
+ cache_x = torch.cat(
116
+ [feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), cache_x],
117
+ dim=2,
118
+ )
119
+ if cache_x.shape[2] < 2 and feat_cache[idx] is not None and feat_cache[idx] == "Rep":
120
+ cache_x = torch.cat([torch.zeros_like(cache_x).to(cache_x.device), cache_x], dim=2)
121
+ if feat_cache[idx] == "Rep":
122
+ x = self.time_conv(x)
123
+ else:
124
+ x = self.time_conv(x, feat_cache[idx])
125
+ feat_cache[idx] = cache_x
126
+ feat_idx[0] += 1
127
+ x = x.reshape(b, 2, c, t, h, w)
128
+ x = torch.stack((x[:, 0, :, :, :, :], x[:, 1, :, :, :, :]), 3)
129
+ x = x.reshape(b, c, t * 2, h, w)
130
+ t = x.shape[2]
131
+ x = rearrange(x, "b c t h w -> (b t) c h w")
132
+ x = self.resample(x)
133
+ x = rearrange(x, "(b t) c h w -> b c t h w", t=t)
134
+
135
+ if self.mode == "downsample3d":
136
+ if feat_cache is not None:
137
+ idx = feat_idx[0]
138
+ if feat_cache[idx] is None:
139
+ feat_cache[idx] = x.clone()
140
+ feat_idx[0] += 1
141
+ else:
142
+ cache_x = x[:, :, -1:, :, :].clone()
143
+ x = self.time_conv(torch.cat([feat_cache[idx][:, :, -1:, :, :], x], 2))
144
+ feat_cache[idx] = cache_x
145
+ feat_idx[0] += 1
146
+ return x
147
+
148
+
149
+ class ResidualBlock(nn.Module):
150
+ def __init__(self, in_dim, out_dim, dropout=0.0):
151
+ super().__init__()
152
+ self.in_dim = in_dim
153
+ self.out_dim = out_dim
154
+ self.residual = nn.Sequential(
155
+ RMS_norm(in_dim, images=False),
156
+ nn.SiLU(),
157
+ CausalConv3d(in_dim, out_dim, 3, padding=1),
158
+ RMS_norm(out_dim, images=False),
159
+ nn.SiLU(),
160
+ nn.Dropout(dropout),
161
+ CausalConv3d(out_dim, out_dim, 3, padding=1),
162
+ )
163
+ self.shortcut = CausalConv3d(in_dim, out_dim, 1) if in_dim != out_dim else nn.Identity()
164
+
165
+ def forward(self, x, feat_cache=None, feat_idx=[0]):
166
+ h = self.shortcut(x)
167
+ for layer in self.residual:
168
+ if isinstance(layer, CausalConv3d) and feat_cache is not None:
169
+ idx = feat_idx[0]
170
+ cache_x = x[:, :, -CACHE_T:, :, :].clone()
171
+ if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
172
+ cache_x = torch.cat(
173
+ [feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), cache_x],
174
+ dim=2,
175
+ )
176
+ x = layer(x, feat_cache[idx])
177
+ feat_cache[idx] = cache_x
178
+ feat_idx[0] += 1
179
+ else:
180
+ x = layer(x)
181
+ return x + h
182
+
183
+
184
+ class AttentionBlock(nn.Module):
185
+ """Single-head causal self-attention over spatial tokens, per frame."""
186
+
187
+ def __init__(self, dim):
188
+ super().__init__()
189
+ self.dim = dim
190
+ self.norm = RMS_norm(dim)
191
+ self.to_qkv = nn.Conv2d(dim, dim * 3, 1)
192
+ self.proj = nn.Conv2d(dim, dim, 1)
193
+ nn.init.zeros_(self.proj.weight)
194
+
195
+ def forward(self, x):
196
+ identity = x
197
+ b, c, t, h, w = x.size()
198
+ x = rearrange(x, "b c t h w -> (b t) c h w")
199
+ x = self.norm(x)
200
+ q, k, v = self.to_qkv(x).reshape(b * t, 1, c * 3, -1).permute(0, 1, 3, 2).contiguous().chunk(3, dim=-1)
201
+ x = F.scaled_dot_product_attention(q, k, v)
202
+ x = x.squeeze(1).permute(0, 2, 1).reshape(b * t, c, h, w)
203
+ x = self.proj(x)
204
+ x = rearrange(x, "(b t) c h w-> b c t h w", t=t)
205
+ return x + identity
206
+
207
+
208
+ def _patchify(x, patch_size):
209
+ if patch_size == 1:
210
+ return x
211
+ if x.dim() == 4:
212
+ return rearrange(x, "b c (h q) (w r) -> b (c r q) h w", q=patch_size, r=patch_size)
213
+ if x.dim() == 5:
214
+ return rearrange(x, "b c f (h q) (w r) -> b (c r q) f h w", q=patch_size, r=patch_size)
215
+ raise ValueError(f"Invalid input shape: {x.shape}")
216
+
217
+
218
+ def _unpatchify(x, patch_size):
219
+ if patch_size == 1:
220
+ return x
221
+ if x.dim() == 4:
222
+ return rearrange(x, "b (c r q) h w -> b c (h q) (w r)", q=patch_size, r=patch_size)
223
+ if x.dim() == 5:
224
+ return rearrange(x, "b (c r q) f h w -> b c f (h q) (w r)", q=patch_size, r=patch_size)
225
+ return x
226
+
227
+
228
+ class AvgDown3D(nn.Module):
229
+ def __init__(self, in_channels, out_channels, factor_t, factor_s=1):
230
+ super().__init__()
231
+ self.in_channels = in_channels
232
+ self.out_channels = out_channels
233
+ self.factor_t = factor_t
234
+ self.factor_s = factor_s
235
+ self.factor = self.factor_t * self.factor_s * self.factor_s
236
+ assert in_channels * self.factor % out_channels == 0
237
+ self.group_size = in_channels * self.factor // out_channels
238
+
239
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
240
+ pad_t = (self.factor_t - x.shape[2] % self.factor_t) % self.factor_t
241
+ x = F.pad(x, (0, 0, 0, 0, pad_t, 0))
242
+ B, C, T, H, W = x.shape
243
+ x = x.view(
244
+ B,
245
+ C,
246
+ T // self.factor_t,
247
+ self.factor_t,
248
+ H // self.factor_s,
249
+ self.factor_s,
250
+ W // self.factor_s,
251
+ self.factor_s,
252
+ )
253
+ x = x.permute(0, 1, 3, 5, 7, 2, 4, 6).contiguous()
254
+ x = x.view(B, C * self.factor, T // self.factor_t, H // self.factor_s, W // self.factor_s)
255
+ x = x.view(B, self.out_channels, self.group_size, T // self.factor_t, H // self.factor_s, W // self.factor_s)
256
+ return x.mean(dim=2)
257
+
258
+
259
+ class DupUp3D(nn.Module):
260
+ def __init__(self, in_channels: int, out_channels: int, factor_t, factor_s=1):
261
+ super().__init__()
262
+ self.in_channels = in_channels
263
+ self.out_channels = out_channels
264
+ self.factor_t = factor_t
265
+ self.factor_s = factor_s
266
+ self.factor = self.factor_t * self.factor_s * self.factor_s
267
+ assert out_channels * self.factor % in_channels == 0
268
+ self.repeats = out_channels * self.factor // in_channels
269
+
270
+ def forward(self, x: torch.Tensor, first_chunk=False) -> torch.Tensor:
271
+ x = x.repeat_interleave(self.repeats, dim=1)
272
+ x = x.view(
273
+ x.size(0),
274
+ self.out_channels,
275
+ self.factor_t,
276
+ self.factor_s,
277
+ self.factor_s,
278
+ x.size(2),
279
+ x.size(3),
280
+ x.size(4),
281
+ )
282
+ x = x.permute(0, 1, 5, 2, 6, 3, 7, 4).contiguous()
283
+ x = x.view(
284
+ x.size(0),
285
+ self.out_channels,
286
+ x.size(2) * self.factor_t,
287
+ x.size(4) * self.factor_s,
288
+ x.size(6) * self.factor_s,
289
+ )
290
+ if first_chunk:
291
+ x = x[:, :, self.factor_t - 1 :, :, :]
292
+ return x
293
+
294
+
295
+ class Down_ResidualBlock(nn.Module):
296
+ def __init__(self, in_dim, out_dim, dropout, mult, temperal_downsample=False, down_flag=False):
297
+ super().__init__()
298
+ self.avg_shortcut = AvgDown3D(
299
+ in_dim, out_dim, factor_t=2 if temperal_downsample else 1, factor_s=2 if down_flag else 1
300
+ )
301
+ downsamples = []
302
+ for _ in range(mult):
303
+ downsamples.append(ResidualBlock(in_dim, out_dim, dropout))
304
+ in_dim = out_dim
305
+ if down_flag:
306
+ mode = "downsample3d" if temperal_downsample else "downsample2d"
307
+ downsamples.append(Resample(out_dim, mode=mode))
308
+ self.downsamples = nn.Sequential(*downsamples)
309
+
310
+ def forward(self, x, feat_cache=None, feat_idx=[0]):
311
+ x_copy = x.clone()
312
+ for module in self.downsamples:
313
+ x = module(x, feat_cache, feat_idx)
314
+ return x + self.avg_shortcut(x_copy)
315
+
316
+
317
+ class Up_ResidualBlock(nn.Module):
318
+ def __init__(self, in_dim, out_dim, dropout, mult, temporal_upsample=False, up_flag=False):
319
+ super().__init__()
320
+ if up_flag:
321
+ self.avg_shortcut = DupUp3D(
322
+ in_dim, out_dim, factor_t=2 if temporal_upsample else 1, factor_s=2 if up_flag else 1
323
+ )
324
+ else:
325
+ self.avg_shortcut = None
326
+ upsamples = []
327
+ for _ in range(mult):
328
+ upsamples.append(ResidualBlock(in_dim, out_dim, dropout))
329
+ in_dim = out_dim
330
+ if up_flag:
331
+ mode = "upsample3d" if temporal_upsample else "upsample2d"
332
+ upsamples.append(Resample(out_dim, mode=mode))
333
+ self.upsamples = nn.Sequential(*upsamples)
334
+
335
+ def forward(self, x, feat_cache=None, feat_idx=[0], first_chunk=False):
336
+ x_main = x.clone()
337
+ for module in self.upsamples:
338
+ x_main = module(x_main, feat_cache, feat_idx)
339
+ if self.avg_shortcut is not None:
340
+ return x_main + self.avg_shortcut(x, first_chunk)
341
+ return x_main
342
+
343
+
344
+ class Encoder3d(nn.Module):
345
+ def __init__(
346
+ self,
347
+ dim=128,
348
+ z_dim=4,
349
+ dim_mult=(1, 2, 4, 4),
350
+ num_res_blocks=2,
351
+ attn_scales=(),
352
+ temperal_downsample=(True, True, False),
353
+ dropout=0.0,
354
+ ):
355
+ super().__init__()
356
+ self.dim = dim
357
+ self.z_dim = z_dim
358
+ self.dim_mult = list(dim_mult)
359
+ self.num_res_blocks = num_res_blocks
360
+ self.attn_scales = list(attn_scales)
361
+ self.temperal_downsample = list(temperal_downsample)
362
+
363
+ dims = [dim * u for u in [1] + self.dim_mult]
364
+ self.conv1 = CausalConv3d(12, dims[0], 3, padding=1)
365
+
366
+ downsamples = []
367
+ for i, (in_dim, out_dim) in enumerate(zip(dims[:-1], dims[1:])):
368
+ t_down_flag = self.temperal_downsample[i] if i < len(self.temperal_downsample) else False
369
+ downsamples.append(
370
+ Down_ResidualBlock(
371
+ in_dim=in_dim,
372
+ out_dim=out_dim,
373
+ dropout=dropout,
374
+ mult=num_res_blocks,
375
+ temperal_downsample=t_down_flag,
376
+ down_flag=i != len(self.dim_mult) - 1,
377
+ )
378
+ )
379
+ self.downsamples = nn.Sequential(*downsamples)
380
+
381
+ self.middle = nn.Sequential(
382
+ ResidualBlock(out_dim, out_dim, dropout),
383
+ AttentionBlock(out_dim),
384
+ ResidualBlock(out_dim, out_dim, dropout),
385
+ )
386
+ self.head = nn.Sequential(
387
+ RMS_norm(out_dim, images=False),
388
+ nn.SiLU(),
389
+ CausalConv3d(out_dim, z_dim, 3, padding=1),
390
+ )
391
+
392
+ def forward(self, x, feat_cache=None, feat_idx=[0]):
393
+ if feat_cache is not None:
394
+ idx = feat_idx[0]
395
+ cache_x = x[:, :, -CACHE_T:, :, :].clone()
396
+ if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
397
+ cache_x = torch.cat(
398
+ [feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), cache_x],
399
+ dim=2,
400
+ )
401
+ x = self.conv1(x, feat_cache[idx])
402
+ feat_cache[idx] = cache_x
403
+ feat_idx[0] += 1
404
+ else:
405
+ x = self.conv1(x)
406
+
407
+ for layer in self.downsamples:
408
+ if feat_cache is not None:
409
+ x = layer(x, feat_cache, feat_idx)
410
+ else:
411
+ x = layer(x)
412
+
413
+ for layer in self.middle:
414
+ if isinstance(layer, ResidualBlock) and feat_cache is not None:
415
+ x = layer(x, feat_cache, feat_idx)
416
+ else:
417
+ x = layer(x)
418
+
419
+ for layer in self.head:
420
+ if isinstance(layer, CausalConv3d) and feat_cache is not None:
421
+ idx = feat_idx[0]
422
+ cache_x = x[:, :, -CACHE_T:, :, :].clone()
423
+ if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
424
+ cache_x = torch.cat(
425
+ [feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), cache_x],
426
+ dim=2,
427
+ )
428
+ x = layer(x, feat_cache[idx])
429
+ feat_cache[idx] = cache_x
430
+ feat_idx[0] += 1
431
+ else:
432
+ x = layer(x)
433
+ return x
434
+
435
+
436
+ class Decoder3d(nn.Module):
437
+ def __init__(
438
+ self,
439
+ dim=128,
440
+ z_dim=4,
441
+ dim_mult=(1, 2, 4, 4),
442
+ num_res_blocks=2,
443
+ attn_scales=(),
444
+ temporal_upsample=(False, True, True),
445
+ dropout=0.0,
446
+ ):
447
+ super().__init__()
448
+ self.dim = dim
449
+ self.z_dim = z_dim
450
+ self.dim_mult = list(dim_mult)
451
+ self.num_res_blocks = num_res_blocks
452
+ self.attn_scales = list(attn_scales)
453
+ self.temporal_upsample = list(temporal_upsample)
454
+
455
+ dims = [dim * u for u in [self.dim_mult[-1]] + self.dim_mult[::-1]]
456
+ self.conv1 = CausalConv3d(z_dim, dims[0], 3, padding=1)
457
+ self.middle = nn.Sequential(
458
+ ResidualBlock(dims[0], dims[0], dropout),
459
+ AttentionBlock(dims[0]),
460
+ ResidualBlock(dims[0], dims[0], dropout),
461
+ )
462
+
463
+ upsamples = []
464
+ for i, (in_dim, out_dim) in enumerate(zip(dims[:-1], dims[1:])):
465
+ t_up_flag = self.temporal_upsample[i] if i < len(self.temporal_upsample) else False
466
+ upsamples.append(
467
+ Up_ResidualBlock(
468
+ in_dim=in_dim,
469
+ out_dim=out_dim,
470
+ dropout=dropout,
471
+ mult=num_res_blocks + 1,
472
+ temporal_upsample=t_up_flag,
473
+ up_flag=i != len(self.dim_mult) - 1,
474
+ )
475
+ )
476
+ self.upsamples = nn.Sequential(*upsamples)
477
+ self.head = nn.Sequential(
478
+ RMS_norm(out_dim, images=False),
479
+ nn.SiLU(),
480
+ CausalConv3d(out_dim, 12, 3, padding=1),
481
+ )
482
+
483
+ def forward(self, x, feat_cache=None, feat_idx=[0], first_chunk=False):
484
+ if feat_cache is not None:
485
+ idx = feat_idx[0]
486
+ cache_x = x[:, :, -CACHE_T:, :, :].clone()
487
+ if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
488
+ cache_x = torch.cat(
489
+ [feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), cache_x],
490
+ dim=2,
491
+ )
492
+ x = self.conv1(x, feat_cache[idx])
493
+ feat_cache[idx] = cache_x
494
+ feat_idx[0] += 1
495
+ else:
496
+ x = self.conv1(x)
497
+
498
+ for layer in self.middle:
499
+ if isinstance(layer, ResidualBlock) and feat_cache is not None:
500
+ x = layer(x, feat_cache, feat_idx)
501
+ else:
502
+ x = layer(x)
503
+
504
+ for layer in self.upsamples:
505
+ if feat_cache is not None:
506
+ x = layer(x, feat_cache, feat_idx, first_chunk)
507
+ else:
508
+ x = layer(x)
509
+
510
+ for layer in self.head:
511
+ if isinstance(layer, CausalConv3d) and feat_cache is not None:
512
+ idx = feat_idx[0]
513
+ cache_x = x[:, :, -CACHE_T:, :, :].clone()
514
+ if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
515
+ cache_x = torch.cat(
516
+ [feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), cache_x],
517
+ dim=2,
518
+ )
519
+ x = layer(x, feat_cache[idx])
520
+ feat_cache[idx] = cache_x
521
+ feat_idx[0] += 1
522
+ else:
523
+ x = layer(x)
524
+ return x
525
+
526
+
527
+ def _count_conv3d(model):
528
+ return sum(1 for m in model.modules() if isinstance(m, CausalConv3d))
529
+
530
+
531
+ class WanVAE_(nn.Module):
532
+ """Upstream Wan2.2 VAE module — encoder3d/decoder3d sandwich with 2x patchify
533
+ on input. State-dict-compatible with ``Wan2.2_VAE.pth``."""
534
+
535
+ def __init__(
536
+ self,
537
+ dim=160,
538
+ dec_dim=256,
539
+ z_dim=16,
540
+ dim_mult=(1, 2, 4, 4),
541
+ num_res_blocks=2,
542
+ attn_scales=(),
543
+ temperal_downsample=(True, True, False),
544
+ dropout=0.0,
545
+ ):
546
+ super().__init__()
547
+ self.dim = dim
548
+ self.z_dim = z_dim
549
+ self.dim_mult = list(dim_mult)
550
+ self.num_res_blocks = num_res_blocks
551
+ self.attn_scales = list(attn_scales)
552
+ self.temperal_downsample = list(temperal_downsample)
553
+ self.temporal_upsample = list(temperal_downsample)[::-1]
554
+
555
+ self.encoder = Encoder3d(
556
+ dim, z_dim * 2, self.dim_mult, num_res_blocks, self.attn_scales, self.temperal_downsample, dropout
557
+ )
558
+ self.conv1 = CausalConv3d(z_dim * 2, z_dim * 2, 1)
559
+ self.conv2 = CausalConv3d(z_dim, z_dim, 1)
560
+ self.decoder = Decoder3d(
561
+ dec_dim, z_dim, self.dim_mult, num_res_blocks, self.attn_scales, self.temporal_upsample, dropout
562
+ )
563
+
564
+ def encode(self, x, scale):
565
+ self.clear_cache()
566
+ x = _patchify(x, patch_size=2)
567
+ t = x.shape[2]
568
+ iter_ = 1 + (t - 1) // 4
569
+ out = None
570
+ for i in range(iter_):
571
+ self._enc_conv_idx = [0]
572
+ if i == 0:
573
+ out = self.encoder(x[:, :, :1, :, :], feat_cache=self._enc_feat_map, feat_idx=self._enc_conv_idx)
574
+ else:
575
+ out_ = self.encoder(
576
+ x[:, :, 1 + 4 * (i - 1) : 1 + 4 * i, :, :],
577
+ feat_cache=self._enc_feat_map,
578
+ feat_idx=self._enc_conv_idx,
579
+ )
580
+ out = torch.cat([out, out_], 2)
581
+ mu, log_var = self.conv1(out).chunk(2, dim=1)
582
+ if isinstance(scale[0], torch.Tensor):
583
+ mu = (mu - scale[0].view(1, self.z_dim, 1, 1, 1)) * scale[1].view(1, self.z_dim, 1, 1, 1)
584
+ else:
585
+ mu = (mu - scale[0]) * scale[1]
586
+ self.clear_cache()
587
+ return mu, log_var
588
+
589
+ def decode(self, z, scale):
590
+ self.clear_cache()
591
+ if isinstance(scale[0], torch.Tensor):
592
+ z = z / scale[1].view(1, self.z_dim, 1, 1, 1) + scale[0].view(1, self.z_dim, 1, 1, 1)
593
+ else:
594
+ z = z / scale[1] + scale[0]
595
+ iter_ = z.shape[2]
596
+ x = self.conv2(z)
597
+ out = None
598
+ for i in range(iter_):
599
+ self._conv_idx = [0]
600
+ if i == 0:
601
+ out = self.decoder(
602
+ x[:, :, i : i + 1, :, :], feat_cache=self._feat_map, feat_idx=self._conv_idx, first_chunk=True
603
+ )
604
+ else:
605
+ out_ = self.decoder(x[:, :, i : i + 1, :, :], feat_cache=self._feat_map, feat_idx=self._conv_idx)
606
+ out = torch.cat([out, out_], 2)
607
+ out = _unpatchify(out, patch_size=2)
608
+ self.clear_cache()
609
+ return out
610
+
611
+ def clear_cache(self):
612
+ self._conv_num = _count_conv3d(self.decoder)
613
+ self._conv_idx = [0]
614
+ self._feat_map = [None] * self._conv_num
615
+ self._enc_conv_num = _count_conv3d(self.encoder)
616
+ self._enc_conv_idx = [0]
617
+ self._enc_feat_map = [None] * self._enc_conv_num
618
+
619
+
620
+ # ----------------------------------------------------------------------------
621
+ # Wan2.2 VAE wrapper + BAGEL-surface adapter
622
+ # ----------------------------------------------------------------------------
623
+
624
+
625
+ # Per-channel mean/std for the 48-channel Wan2.2 latent space (upstream constants).
626
+ _WAN22_LATENT_MEAN: tuple[float, ...] = (
627
+ -0.2289,
628
+ -0.0052,
629
+ -0.1323,
630
+ -0.2339,
631
+ -0.2799,
632
+ 0.0174,
633
+ 0.1838,
634
+ 0.1557,
635
+ -0.1382,
636
+ 0.0542,
637
+ 0.2813,
638
+ 0.0891,
639
+ 0.1570,
640
+ -0.0098,
641
+ 0.0375,
642
+ -0.1825,
643
+ -0.2246,
644
+ -0.1207,
645
+ -0.0698,
646
+ 0.5109,
647
+ 0.2665,
648
+ -0.2108,
649
+ -0.2158,
650
+ 0.2502,
651
+ -0.2055,
652
+ -0.0322,
653
+ 0.1109,
654
+ 0.1567,
655
+ -0.0729,
656
+ 0.0899,
657
+ -0.2799,
658
+ -0.1230,
659
+ -0.0313,
660
+ -0.1649,
661
+ 0.0117,
662
+ 0.0723,
663
+ -0.2839,
664
+ -0.2083,
665
+ -0.0520,
666
+ 0.3748,
667
+ 0.0152,
668
+ 0.1957,
669
+ 0.1433,
670
+ -0.2944,
671
+ 0.3573,
672
+ -0.0548,
673
+ -0.1681,
674
+ -0.0667,
675
+ )
676
+ _WAN22_LATENT_STD: tuple[float, ...] = (
677
+ 0.4765,
678
+ 1.0364,
679
+ 0.4514,
680
+ 1.1677,
681
+ 0.5313,
682
+ 0.4990,
683
+ 0.4818,
684
+ 0.5013,
685
+ 0.8158,
686
+ 1.0344,
687
+ 0.5894,
688
+ 1.0901,
689
+ 0.6885,
690
+ 0.6165,
691
+ 0.8454,
692
+ 0.4978,
693
+ 0.5759,
694
+ 0.3523,
695
+ 0.7135,
696
+ 0.6804,
697
+ 0.5833,
698
+ 1.4146,
699
+ 0.8986,
700
+ 0.5659,
701
+ 0.7069,
702
+ 0.5338,
703
+ 0.4889,
704
+ 0.4917,
705
+ 0.4069,
706
+ 0.4999,
707
+ 0.6866,
708
+ 0.4093,
709
+ 0.5709,
710
+ 0.6065,
711
+ 0.6415,
712
+ 0.4944,
713
+ 0.5726,
714
+ 1.2042,
715
+ 0.5458,
716
+ 1.6887,
717
+ 0.3971,
718
+ 1.0600,
719
+ 0.3943,
720
+ 0.5537,
721
+ 0.5444,
722
+ 0.4089,
723
+ 0.7468,
724
+ 0.7744,
725
+ )
726
+
727
+
728
+ class LanceWanVAE(nn.Module):
729
+ """Wan2.2 VAE wrapped for BAGEL's pipeline.
730
+
731
+ Exposes BAGEL's image-VAE surface — ``encode(BCHW) -> BC_zHW`` and
732
+ ``decode(BC_zHW) -> BCHW`` — by treating each image as a 1-frame video clip.
733
+ A 5-D ``encode_video``/``decode_video`` path is also provided for the
734
+ ``Lance_3B_Video`` checkpoint.
735
+
736
+ Construction is lazy: the heavy ``WanVAE_`` and ``Wan2.2_VAE.pth`` are not
737
+ materialized until first use. Once built, the inner module is registered as
738
+ a submodule so ``self.parameters()``, ``self.to(device)`` and ``vae_dtype =
739
+ next(vae.parameters()).dtype`` (used by BAGEL's decode path) all behave.
740
+ """
741
+
742
+ z_channels: int = 48
743
+ downsample_spatial: int = 16
744
+ downsample_temporal: int = 4
745
+
746
+ def __init__(
747
+ self,
748
+ vae_path: str,
749
+ dtype: torch.dtype = torch.bfloat16,
750
+ device: torch.device | str | None = None,
751
+ *,
752
+ lazy: bool = False,
753
+ ):
754
+ super().__init__()
755
+ self._vae_path = vae_path
756
+ self._dtype = dtype
757
+ self._device = torch.device(device) if device is not None else None
758
+ self._built = False
759
+ # Registered lazily once the .pth is loaded:
760
+ # self.model: WanVAE_ (set in _ensure_built)
761
+ if not lazy:
762
+ # BAGEL's decode path inspects ``next(vae.parameters()).dtype`` before
763
+ # the first encode/decode call, so we have to register the submodule
764
+ # eagerly when a real device is provided. ``lazy=True`` is reserved
765
+ # for config-only / no-GPU paths.
766
+ self._ensure_built()
767
+
768
+ def _ensure_built(self) -> None:
769
+ if self._built:
770
+ return
771
+ device = self._device or (
772
+ next(self.parameters(), torch.empty(0)).device
773
+ if any(True for _ in self.parameters())
774
+ else torch.device("cpu")
775
+ )
776
+ # Wan2.2 VAE config from upstream Wan2_2_VAE defaults:
777
+ # z_dim=48 (channels), c_dim=160 (encoder width), dec_dim=256 (decoder width),
778
+ # dim_mult=[1,2,4,4], temperal_downsample=[False,True,True].
779
+ model = WanVAE_(
780
+ dim=160,
781
+ dec_dim=256,
782
+ z_dim=self.z_channels,
783
+ dim_mult=(1, 2, 4, 4),
784
+ num_res_blocks=2,
785
+ attn_scales=(),
786
+ temperal_downsample=(False, True, True),
787
+ dropout=0.0,
788
+ )
789
+ logger.info("Loading Wan2.2 VAE state dict from %s", self._vae_path)
790
+ state = torch.load(self._vae_path, map_location="cpu", weights_only=True)
791
+ missing, unexpected = model.load_state_dict(state, strict=False)
792
+ if missing:
793
+ logger.warning("Wan2.2 VAE missing keys (%d): %s", len(missing), missing[:5])
794
+ if unexpected:
795
+ logger.warning("Wan2.2 VAE unexpected keys (%d): %s", len(unexpected), unexpected[:5])
796
+ model = model.to(device=device, dtype=self._dtype).eval()
797
+ model.requires_grad_(False)
798
+ # Register as submodule so .to() / .parameters() pick it up.
799
+ self.model = model
800
+ # Per-channel normalization tensors (matches upstream Wan2_2_VAE.scale = [mean, 1/std]).
801
+ mean = torch.tensor(_WAN22_LATENT_MEAN, dtype=self._dtype, device=device)
802
+ inv_std = torch.tensor([1.0 / s for s in _WAN22_LATENT_STD], dtype=self._dtype, device=device)
803
+ # Buffers, so .to(device) follows along.
804
+ self.register_buffer("_latent_mean", mean, persistent=False)
805
+ self.register_buffer("_latent_inv_std", inv_std, persistent=False)
806
+ self._built = True
807
+
808
+ @torch.inference_mode()
809
+ def encode_video(self, video: torch.Tensor, *, use_sample: bool = True) -> torch.Tensor:
810
+ """Encode a 5-D clip ``[B, 3, T, H, W]`` -> latent ``[B, 48, t, h, w]``."""
811
+ self._ensure_built()
812
+ video = video.to(self.model.encoder.conv1.weight.dtype)
813
+ mu, log_var = self.model.encode(video, [self._latent_mean, self._latent_inv_std])
814
+ if use_sample:
815
+ std = torch.exp(0.5 * log_var)
816
+ mu = mu + std * torch.randn_like(std)
817
+ return mu
818
+
819
+ @torch.inference_mode()
820
+ def decode_video(self, latent: torch.Tensor) -> torch.Tensor:
821
+ """Decode a 5-D latent ``[B, 48, t, h, w]`` -> video ``[B, 3, T, H, W]``."""
822
+ self._ensure_built()
823
+ latent = latent.to(self.model.decoder.conv1.weight.dtype)
824
+ out = self.model.decode(latent, [self._latent_mean, self._latent_inv_std])
825
+ return out.clamp_(-1.0, 1.0)
826
+
827
+ # ----- BAGEL image-VAE surface (4-D) -------------------------------- #
828
+
829
+ @torch.inference_mode()
830
+ def encode(self, padded_images: torch.Tensor) -> torch.Tensor:
831
+ """``[B, 3, H, W]`` -> ``[B, 48, H/16, W/16]`` (single-frame image path).
832
+
833
+ Each image is wrapped as a 1-frame clip, encoded, and the temporal axis
834
+ is squeezed back out so the result matches BAGEL's iteration pattern.
835
+ """
836
+ if padded_images.dim() == 4:
837
+ video = padded_images.unsqueeze(2) # B,C,1,H,W
838
+ squeeze_time = True # image path: caller iterates 3-D latents
839
+ elif padded_images.dim() == 5:
840
+ video = padded_images
841
+ # video path: caller expects 4-D latents (C,T,H,W), even when T==1
842
+ # (e.g. video2video with a 1-frame reference clip). Squeezing here
843
+ # would drop the temporal axis and break the downstream indexing
844
+ # ``latent[:, :t_lat, ...]`` in ``forward_cache_update_vae``.
845
+ squeeze_time = False
846
+ else:
847
+ raise ValueError(f"LanceWanVAE.encode expects 4-D BCHW or 5-D BCTHW, got {tuple(padded_images.shape)}")
848
+ latent = self.encode_video(video, use_sample=True)
849
+ # B,48,t,h,w -> B,48,h,w when t == 1 (image path only)
850
+ if squeeze_time and latent.shape[2] == 1:
851
+ return latent.squeeze(2)
852
+ return latent
853
+
854
+ @torch.inference_mode()
855
+ def decode(self, latent: torch.Tensor) -> torch.Tensor:
856
+ """``[B, 48, h, w]`` -> ``[B, 3, H, W]`` (single-frame image path)."""
857
+ if latent.dim() == 4:
858
+ latent = latent.unsqueeze(2) # B,48,1,h,w
859
+ elif latent.dim() != 5:
860
+ raise ValueError(f"LanceWanVAE.decode expects 4-D BCHW or 5-D BCTHW, got {tuple(latent.shape)}")
861
+ video = self.decode_video(latent) # B,3,T,H,W
862
+ if video.shape[2] == 1:
863
+ return video.squeeze(2)
864
+ return video
865
+
866
+
867
+ def build_wan22_vae(vae_path: str, dtype: torch.dtype = torch.bfloat16, device=None) -> LanceWanVAE:
868
+ """Convenience factory: lazy-construct a :class:`LanceWanVAE` adapter."""
869
+ return LanceWanVAE(vae_path=vae_path, dtype=dtype, device=device)
870
+
871
+
872
+ __all__ = ["LanceWanVAE", "WanVAE_", "build_wan22_vae"]