MCPcopy Create free account
hub / github.com/AMAP-ML/Eevee / Decoder3d

Class Decoder3d

models/vae.py:736–838  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

734
735
736class Decoder3d(nn.Module):
737
738 def __init__(self,
739 dim=128,
740 z_dim=4,
741 dim_mult=[1, 2, 4, 4],
742 num_res_blocks=2,
743 attn_scales=[],
744 temperal_upsample=[False, True, True],
745 dropout=0.0):
746 super().__init__()
747 self.dim = dim
748 self.z_dim = z_dim
749 self.dim_mult = dim_mult
750 self.num_res_blocks = num_res_blocks
751 self.attn_scales = attn_scales
752 self.temperal_upsample = temperal_upsample
753
754 # dimensions
755 dims = [dim * u for u in [dim_mult[-1]] + dim_mult[::-1]]
756 scale = 1.0 / 2**(len(dim_mult) - 2)
757
758 # init block
759 self.conv1 = CausalConv3d(z_dim, dims[0], 3, padding=1)
760
761 # middle blocks
762 self.middle = nn.Sequential(ResidualBlock(dims[0], dims[0], dropout),
763 AttentionBlock(dims[0]),
764 ResidualBlock(dims[0], dims[0], dropout))
765
766 # upsample blocks
767 upsamples = []
768 for i, (in_dim, out_dim) in enumerate(zip(dims[:-1], dims[1:])):
769 # residual (+attention) blocks
770 if i == 1 or i == 2 or i == 3:
771 in_dim = in_dim // 2
772 for _ in range(num_res_blocks + 1):
773 upsamples.append(ResidualBlock(in_dim, out_dim, dropout))
774 if scale in attn_scales:
775 upsamples.append(AttentionBlock(out_dim))
776 in_dim = out_dim
777
778 # upsample block
779 if i != len(dim_mult) - 1:
780 mode = 'upsample3d' if temperal_upsample[i] else 'upsample2d'
781 upsamples.append(Resample(out_dim, mode=mode))
782 scale *= 2.0
783 self.upsamples = nn.Sequential(*upsamples)
784
785 # output blocks
786 self.head = nn.Sequential(RMS_norm(out_dim, images=False), nn.SiLU(),
787 CausalConv3d(out_dim, 3, 3, padding=1))
788
789 def forward(self, x, feat_cache=None, feat_idx=[0]):
790 ## conv1
791 if feat_cache is not None:
792 idx = feat_idx[0]
793 cache_x = x[:, :, -CACHE_T:, :, :].clone()

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected