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.
- simit/__init__.py +39 -0
- simit/api.py +258 -0
- simit/backends/__init__.py +112 -0
- simit/backends/bagel.py +359 -0
- simit/backends/base.py +114 -0
- simit/backends/hf.py +336 -0
- simit/backends/imagegen.py +136 -0
- simit/backends/lance.py +338 -0
- simit/backends/mot/engine.py +844 -0
- simit/backends/mot/kernels.py +130 -0
- simit/backends/mot/modeling.py +514 -0
- simit/backends/mot/wan_vae.py +872 -0
- simit/backends/vllm.py +204 -0
- simit/config.py +112 -0
- simit/metrics.py +113 -0
- simit/pipeline/core.py +586 -0
- simit/pipeline/prompts.py +167 -0
- simit/skills/__init__.py +6 -0
- simit/skills/assets/MERMAID_LICENSE +21 -0
- simit/skills/assets/mermaid.min.js +3587 -0
- simit/skills/base.py +172 -0
- simit/skills/builtin.py +614 -0
- simit/skills/helpers.py +197 -0
- simit/skills/pool.py +262 -0
- simit/skills/prompts.py +2374 -0
- simit/skills/renderers.py +1292 -0
- simit/skills/routing.py +89 -0
- simit/skills/sandbox.py +188 -0
- simit/skills/web.py +253 -0
- simit/tune.py +323 -0
- simit/types.py +96 -0
- simit/utils.py +67 -0
- simit-0.1.0.dist-info/METADATA +388 -0
- simit-0.1.0.dist-info/RECORD +38 -0
- simit-0.1.0.dist-info/WHEEL +5 -0
- simit-0.1.0.dist-info/licenses/LICENSE +202 -0
- simit-0.1.0.dist-info/licenses/src/simit/skills/assets/MERMAID_LICENSE +21 -0
- simit-0.1.0.dist-info/top_level.txt +1 -0
|
@@ -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"]
|