| 221 | return c11 |
| 222 | |
| 223 | class UpSampleResNetSkipGated(BaseModule): |
| 224 | def __init__(self, nc_in, nc_out, nf, use_bias, norm, conv_by, conv_type, |
| 225 | use_skip_connection=False, use_flow_tsm=False): |
| 226 | super().__init__(conv_type, 'vanilla') |
| 227 | # Upsample 1 |
| 228 | self.conv_c2 = self.ConvBlock(512, 256, kernel_size=(3, 1, 1), stride=1, |
| 229 | bias=use_bias, norm=norm, conv_by=conv_by, use_flow_tsm=use_flow_tsm) |
| 230 | self.conv_c4 = self.ConvBlock(1024, 256, kernel_size=(3, 1, 1), stride=1, |
| 231 | bias=use_bias, norm=norm, conv_by=conv_by, use_flow_tsm=use_flow_tsm) |
| 232 | |
| 233 | self.deconv1 = self.DeconvBlock(2048, 256, kernel_size=(3, 1, 1), stride=1, |
| 234 | bias=use_bias, norm=norm, conv_by="2d", use_flow_tsm=False) |
| 235 | self.conv9 = self.ConvBlock( |
| 236 | 256, 256, kernel_size=(3, 3, 3), stride=(1, 1, 1), padding=1, |
| 237 | bias=use_bias, norm=norm, conv_by=conv_by, use_flow_tsm=use_flow_tsm) |
| 238 | |
| 239 | # Upsample 2 |
| 240 | self.deconv2 = self.DeconvBlock( |
| 241 | 256, |
| 242 | 256, kernel_size=(3, 1, 1), stride=1, |
| 243 | bias=use_bias, norm=norm, conv_by="2d", use_flow_tsm=False) |
| 244 | self.conv10 = self.ConvBlock( |
| 245 | 256, 256, kernel_size=(3, 3, 3), stride=(1, 1, 1), padding=1, |
| 246 | bias=use_bias, norm=norm, conv_by=conv_by, use_flow_tsm=use_flow_tsm) |
| 247 | |
| 248 | # Upsample 3 |
| 249 | self.deconv3 = self.DeconvBlock( |
| 250 | 256, |
| 251 | 256, kernel_size=(3, 1, 1), stride=1, |
| 252 | bias=use_bias, norm=norm, conv_by="2d", use_flow_tsm=False) |
| 253 | self.conv11 = self.ConvBlock( |
| 254 | 256, nc_out, kernel_size=(3, 3, 3), stride=(1, 1, 1), |
| 255 | padding=1, bias=use_bias, norm=None, activation=None, conv_by=conv_by, use_flow_tsm=use_flow_tsm) |
| 256 | |
| 257 | for name, module in self.named_modules(): |
| 258 | if isinstance(module, (GatedConv)): |
| 259 | nn.init.kaiming_uniform_(module.featureConv.layer.weight, a=1) |
| 260 | nn.init.constant_(module.featureConv.layer.bias, 0) |
| 261 | nn.init.kaiming_uniform_(module.gatingConv.layer.weight, a=1) |
| 262 | nn.init.constant_(module.gatingConv.layer.bias, 0) |
| 263 | elif isinstance(module, (VanillaDeconv)): |
| 264 | nn.init.kaiming_uniform_(module.conv.featureConv.layer.weight, a=1) |
| 265 | nn.init.constant_(module.conv.featureConv.layer.bias, 0) |
| 266 | |
| 267 | def forward(self, inp, flows=None): |
| 268 | c1, c2, c4, c8 = inp |
| 269 | |
| 270 | c2 = self.conv_c2(c2, flows) |
| 271 | c4 = self.conv_c4(c4, flows) |
| 272 | |
| 273 | c4 = interpolate(c4, scale_factor=2, mode="nearest") |
| 274 | c2 = interpolate(c2, scale_factor=4, mode="nearest") |
| 275 | c1 = interpolate(c1, scale_factor=4, mode="nearest") |
| 276 | |
| 277 | d1 = self.deconv1(c8, flows) |
| 278 | d1 = d1 + c4 |
| 279 | c9 = self.conv9(d1, flows) |
| 280 | d2 = self.deconv2(c9, flows) |