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