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

Class Decoder3d

wan/models/wan_vae3_8.py:621–728  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

619
620
621class 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:

Callers 1

__init__Method · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected