Architecture of MODNet
| 199 | #------------------------------------------------------------------------------ |
| 200 | |
| 201 | class MODNet(nn.Module): |
| 202 | """ Architecture of MODNet |
| 203 | """ |
| 204 | |
| 205 | def __init__(self, in_channels=3, hr_channels=32, backbone_arch='mobilenetv2', backbone_pretrained=True): |
| 206 | super(MODNet, self).__init__() |
| 207 | |
| 208 | self.in_channels = in_channels |
| 209 | self.hr_channels = hr_channels |
| 210 | self.backbone_arch = backbone_arch |
| 211 | self.backbone_pretrained = backbone_pretrained |
| 212 | |
| 213 | self.backbone = SUPPORTED_BACKBONES[self.backbone_arch](self.in_channels) |
| 214 | |
| 215 | self.lr_branch = LRBranch(self.backbone) |
| 216 | self.hr_branch = HRBranch(self.hr_channels, self.backbone.enc_channels) |
| 217 | self.f_branch = FusionBranch(self.hr_channels, self.backbone.enc_channels) |
| 218 | |
| 219 | for m in self.modules(): |
| 220 | if isinstance(m, nn.Conv2d): |
| 221 | self._init_conv(m) |
| 222 | elif isinstance(m, nn.BatchNorm2d) or isinstance(m, nn.InstanceNorm2d): |
| 223 | self._init_norm(m) |
| 224 | |
| 225 | if self.backbone_pretrained: |
| 226 | self.backbone.load_pretrained_ckpt() |
| 227 | |
| 228 | def forward(self, img): |
| 229 | # NOTE |
| 230 | lr_out = self.lr_branch(img) |
| 231 | lr8x = lr_out[0] |
| 232 | enc2x = lr_out[1] |
| 233 | enc4x = lr_out[2] |
| 234 | |
| 235 | hr2x = self.hr_branch(img, enc2x, enc4x, lr8x) |
| 236 | |
| 237 | pred_matte = self.f_branch(img, lr8x, hr2x) |
| 238 | |
| 239 | return pred_matte |
| 240 | |
| 241 | def freeze_norm(self): |
| 242 | norm_types = [nn.BatchNorm2d, nn.InstanceNorm2d] |
| 243 | for m in self.modules(): |
| 244 | for n in norm_types: |
| 245 | if isinstance(m, n): |
| 246 | m.eval() |
| 247 | continue |
| 248 | |
| 249 | def _init_conv(self, conv): |
| 250 | nn.init.kaiming_uniform_( |
| 251 | conv.weight, a=0, mode='fan_in', nonlinearity='relu') |
| 252 | if conv.bias is not None: |
| 253 | nn.init.constant_(conv.bias, 0) |
| 254 | |
| 255 | def _init_norm(self, norm): |
| 256 | if norm.weight is not None: |
| 257 | nn.init.constant_(norm.weight, 1) |
| 258 | nn.init.constant_(norm.bias, 0) |
nothing calls this directly
no outgoing calls
no test coverage detected