(
self,
in_channels: int = 3,
dim=128,
z_dim=4,
dim_mult=[1, 2, 4, 4],
num_res_blocks=2,
attn_scales=[],
temperal_downsample=[True, True, False],
dropout=0.0,
non_linearity: str = "silu",
is_residual: bool = False, # wan 2.2 vae use a residual downblock
)
| 518 | """ |
| 519 | |
| 520 | def __init__( |
| 521 | self, |
| 522 | in_channels: int = 3, |
| 523 | dim=128, |
| 524 | z_dim=4, |
| 525 | dim_mult=[1, 2, 4, 4], |
| 526 | num_res_blocks=2, |
| 527 | attn_scales=[], |
| 528 | temperal_downsample=[True, True, False], |
| 529 | dropout=0.0, |
| 530 | non_linearity: str = "silu", |
| 531 | is_residual: bool = False, # wan 2.2 vae use a residual downblock |
| 532 | ): |
| 533 | super().__init__() |
| 534 | self.dim = dim |
| 535 | self.z_dim = z_dim |
| 536 | self.dim_mult = dim_mult |
| 537 | self.num_res_blocks = num_res_blocks |
| 538 | self.attn_scales = attn_scales |
| 539 | self.temperal_downsample = temperal_downsample |
| 540 | self.nonlinearity = get_activation(non_linearity) |
| 541 | |
| 542 | # dimensions |
| 543 | dims = [dim * u for u in [1] + dim_mult] |
| 544 | scale = 1.0 |
| 545 | |
| 546 | # init block |
| 547 | self.conv_in = WanCausalConv3d(in_channels, dims[0], 3, padding=1) |
| 548 | |
| 549 | # downsample blocks |
| 550 | self.down_blocks = nn.ModuleList([]) |
| 551 | for i, (in_dim, out_dim) in enumerate(zip(dims[:-1], dims[1:])): |
| 552 | # residual (+attention) blocks |
| 553 | if is_residual: |
| 554 | self.down_blocks.append( |
| 555 | WanResidualDownBlock( |
| 556 | in_dim, |
| 557 | out_dim, |
| 558 | dropout, |
| 559 | num_res_blocks, |
| 560 | temperal_downsample=temperal_downsample[i] if i != len(dim_mult) - 1 else False, |
| 561 | down_flag=i != len(dim_mult) - 1, |
| 562 | ) |
| 563 | ) |
| 564 | else: |
| 565 | for _ in range(num_res_blocks): |
| 566 | self.down_blocks.append(WanResidualBlock(in_dim, out_dim, dropout)) |
| 567 | if scale in attn_scales: |
| 568 | self.down_blocks.append(WanAttentionBlock(out_dim)) |
| 569 | in_dim = out_dim |
| 570 | |
| 571 | # downsample block |
| 572 | if i != len(dim_mult) - 1: |
| 573 | mode = "downsample3d" if temperal_downsample[i] else "downsample2d" |
| 574 | self.down_blocks.append(WanResample(out_dim, mode=mode)) |
| 575 | scale /= 2.0 |
| 576 | |
| 577 | # middle blocks |
nothing calls this directly
no test coverage detected