MCPcopy Create free account
hub / github.com/MaureenZOU/TSAM / UpSampleResNetSkipGated

Class UpSampleResNetSkipGated

src/model/modules.py:223–286  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

221 return c11
222
223class 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)

Callers 1

__init__Method · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected