(self)
| 70 | |
| 71 | class SVDVAEDecoder(torch.nn.Module): |
| 72 | def __init__(self): |
| 73 | super().__init__() |
| 74 | self.scaling_factor = 0.18215 |
| 75 | self.conv_in = torch.nn.Conv2d(4, 512, kernel_size=3, padding=1) |
| 76 | |
| 77 | self.blocks = torch.nn.ModuleList([ |
| 78 | # UNetMidBlock |
| 79 | ResnetBlock(512, 512, eps=1e-6), |
| 80 | TemporalResnetBlock(512, 512, eps=1e-6), |
| 81 | VAEAttentionBlock(1, 512, 512, 1, eps=1e-6), |
| 82 | ResnetBlock(512, 512, eps=1e-6), |
| 83 | TemporalResnetBlock(512, 512, eps=1e-6), |
| 84 | # UpDecoderBlock |
| 85 | ResnetBlock(512, 512, eps=1e-6), |
| 86 | TemporalResnetBlock(512, 512, eps=1e-6), |
| 87 | ResnetBlock(512, 512, eps=1e-6), |
| 88 | TemporalResnetBlock(512, 512, eps=1e-6), |
| 89 | ResnetBlock(512, 512, eps=1e-6), |
| 90 | TemporalResnetBlock(512, 512, eps=1e-6), |
| 91 | UpSampler(512), |
| 92 | # UpDecoderBlock |
| 93 | ResnetBlock(512, 512, eps=1e-6), |
| 94 | TemporalResnetBlock(512, 512, eps=1e-6), |
| 95 | ResnetBlock(512, 512, eps=1e-6), |
| 96 | TemporalResnetBlock(512, 512, eps=1e-6), |
| 97 | ResnetBlock(512, 512, eps=1e-6), |
| 98 | TemporalResnetBlock(512, 512, eps=1e-6), |
| 99 | UpSampler(512), |
| 100 | # UpDecoderBlock |
| 101 | ResnetBlock(512, 256, eps=1e-6), |
| 102 | TemporalResnetBlock(256, 256, eps=1e-6), |
| 103 | ResnetBlock(256, 256, eps=1e-6), |
| 104 | TemporalResnetBlock(256, 256, eps=1e-6), |
| 105 | ResnetBlock(256, 256, eps=1e-6), |
| 106 | TemporalResnetBlock(256, 256, eps=1e-6), |
| 107 | UpSampler(256), |
| 108 | # UpDecoderBlock |
| 109 | ResnetBlock(256, 128, eps=1e-6), |
| 110 | TemporalResnetBlock(128, 128, eps=1e-6), |
| 111 | ResnetBlock(128, 128, eps=1e-6), |
| 112 | TemporalResnetBlock(128, 128, eps=1e-6), |
| 113 | ResnetBlock(128, 128, eps=1e-6), |
| 114 | TemporalResnetBlock(128, 128, eps=1e-6), |
| 115 | ]) |
| 116 | |
| 117 | self.conv_norm_out = torch.nn.GroupNorm(num_channels=128, num_groups=32, eps=1e-5) |
| 118 | self.conv_act = torch.nn.SiLU() |
| 119 | self.conv_out = torch.nn.Conv2d(128, 3, kernel_size=3, padding=1) |
| 120 | self.time_conv_out = torch.nn.Conv3d(3, 3, kernel_size=(3, 1, 1), padding=(1, 0, 0)) |
| 121 | |
| 122 | |
| 123 | def forward(self, sample): |
no test coverage detected