MCPcopy Create free account
hub / github.com/tdrussell/diffusion-pipe / Decoder3d

Class Decoder3d

models/wan/vae2_2.py:616–723  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

614
615
616class 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:

Callers 1

__init__Method · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected