MCPcopy Create free account
hub / github.com/Francis-Rings/FlashPortrait / Encoder3d

Class Encoder3d

wan/models/wan_vae.py:269–370  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

267
268
269class Encoder3d(nn.Module):
270
271 def __init__(self,
272 dim=128,
273 z_dim=4,
274 dim_mult=[1, 2, 4, 4],
275 num_res_blocks=2,
276 attn_scales=[],
277 temperal_downsample=[True, True, False],
278 dropout=0.0):
279 super().__init__()
280 self.dim = dim
281 self.z_dim = z_dim
282 self.dim_mult = dim_mult
283 self.num_res_blocks = num_res_blocks
284 self.attn_scales = attn_scales
285 self.temperal_downsample = temperal_downsample
286
287 # dimensions
288 dims = [dim * u for u in [1] + dim_mult]
289 scale = 1.0
290
291 # init block
292 self.conv1 = CausalConv3d(3, dims[0], 3, padding=1)
293
294 # downsample blocks
295 downsamples = []
296 for i, (in_dim, out_dim) in enumerate(zip(dims[:-1], dims[1:])):
297 # residual (+attention) blocks
298 for _ in range(num_res_blocks):
299 downsamples.append(ResidualBlock(in_dim, out_dim, dropout))
300 if scale in attn_scales:
301 downsamples.append(AttentionBlock(out_dim))
302 in_dim = out_dim
303
304 # downsample block
305 if i != len(dim_mult) - 1:
306 mode = 'downsample3d' if temperal_downsample[
307 i] else 'downsample2d'
308 downsamples.append(Resample(out_dim, mode=mode))
309 scale /= 2.0
310 self.downsamples = nn.Sequential(*downsamples)
311
312 # middle blocks
313 self.middle = nn.Sequential(
314 ResidualBlock(out_dim, out_dim, dropout), AttentionBlock(out_dim),
315 ResidualBlock(out_dim, out_dim, dropout))
316
317 # output blocks
318 self.head = nn.Sequential(
319 RMS_norm(out_dim, images=False), nn.SiLU(),
320 CausalConv3d(out_dim, z_dim, 3, padding=1))
321
322 def forward(self, x, feat_cache=None, feat_idx=[0]):
323 if feat_cache is not None:
324 idx = feat_idx[0]
325 cache_x = x[:, :, -CACHE_T:, :, :].clone()
326 if cache_x.shape[2] < 2 and feat_cache[idx] is not None:

Callers 1

__init__Method · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected