| 619 | |
| 620 | |
| 621 | class Decoder3d(nn.Module): |
| 622 | |
| 623 | def __init__( |
| 624 | self, |
| 625 | dim=128, |
| 626 | z_dim=4, |
| 627 | dim_mult=[1, 2, 4, 4], |
| 628 | num_res_blocks=2, |
| 629 | attn_scales=[], |
| 630 | temperal_upsample=[False, True, True], |
| 631 | dropout=0.0, |
| 632 | ): |
| 633 | super().__init__() |
| 634 | self.dim = dim |
| 635 | self.z_dim = z_dim |
| 636 | self.dim_mult = dim_mult |
| 637 | self.num_res_blocks = num_res_blocks |
| 638 | self.attn_scales = attn_scales |
| 639 | self.temperal_upsample = temperal_upsample |
| 640 | |
| 641 | # dimensions |
| 642 | dims = [dim * u for u in [dim_mult[-1]] + dim_mult[::-1]] |
| 643 | scale = 1.0 / 2**(len(dim_mult) - 2) |
| 644 | # init block |
| 645 | self.conv1 = CausalConv3d(z_dim, dims[0], 3, padding=1) |
| 646 | |
| 647 | # middle blocks |
| 648 | self.middle = nn.Sequential( |
| 649 | ResidualBlock(dims[0], dims[0], dropout), |
| 650 | AttentionBlock(dims[0]), |
| 651 | ResidualBlock(dims[0], dims[0], dropout), |
| 652 | ) |
| 653 | |
| 654 | # upsample blocks |
| 655 | upsamples = [] |
| 656 | for i, (in_dim, out_dim) in enumerate(zip(dims[:-1], dims[1:])): |
| 657 | t_up_flag = temperal_upsample[i] if i < len( |
| 658 | temperal_upsample) else False |
| 659 | upsamples.append( |
| 660 | Up_ResidualBlock( |
| 661 | in_dim=in_dim, |
| 662 | out_dim=out_dim, |
| 663 | dropout=dropout, |
| 664 | mult=num_res_blocks + 1, |
| 665 | temperal_upsample=t_up_flag, |
| 666 | up_flag=i != len(dim_mult) - 1, |
| 667 | )) |
| 668 | self.upsamples = nn.Sequential(*upsamples) |
| 669 | |
| 670 | # output blocks |
| 671 | self.head = nn.Sequential( |
| 672 | RMS_norm(out_dim, images=False), |
| 673 | nn.SiLU(), |
| 674 | CausalConv3d(out_dim, 12, 3, padding=1), |
| 675 | ) |
| 676 | |
| 677 | def forward(self, x, feat_cache=None, feat_idx=[0], first_chunk=False): |
| 678 | if feat_cache is not None: |