| 440 | return nn.ModuleList(branches) |
| 441 | |
| 442 | def _make_fuse_layers(self): |
| 443 | if self.num_branches == 1: |
| 444 | return nn.Identity() |
| 445 | |
| 446 | num_branches = self.num_branches |
| 447 | num_inchannels = self.num_inchannels |
| 448 | fuse_layers = [] |
| 449 | for i in range(num_branches if self.multi_scale_output else 1): |
| 450 | fuse_layer = [] |
| 451 | for j in range(num_branches): |
| 452 | if j > i: |
| 453 | fuse_layer.append(nn.Sequential( |
| 454 | nn.Conv2d(num_inchannels[j], num_inchannels[i], 1, 1, 0, bias=False), |
| 455 | nn.BatchNorm2d(num_inchannels[i], momentum=_BN_MOMENTUM), |
| 456 | nn.Upsample(scale_factor=2 ** (j - i), mode='nearest'))) |
| 457 | elif j == i: |
| 458 | fuse_layer.append(nn.Identity()) |
| 459 | else: |
| 460 | conv3x3s = [] |
| 461 | for k in range(i - j): |
| 462 | if k == i - j - 1: |
| 463 | num_outchannels_conv3x3 = num_inchannels[i] |
| 464 | conv3x3s.append(nn.Sequential( |
| 465 | nn.Conv2d(num_inchannels[j], num_outchannels_conv3x3, 3, 2, 1, bias=False), |
| 466 | nn.BatchNorm2d(num_outchannels_conv3x3, momentum=_BN_MOMENTUM))) |
| 467 | else: |
| 468 | num_outchannels_conv3x3 = num_inchannels[j] |
| 469 | conv3x3s.append(nn.Sequential( |
| 470 | nn.Conv2d(num_inchannels[j], num_outchannels_conv3x3, 3, 2, 1, bias=False), |
| 471 | nn.BatchNorm2d(num_outchannels_conv3x3, momentum=_BN_MOMENTUM), |
| 472 | nn.ReLU(False))) |
| 473 | fuse_layer.append(nn.Sequential(*conv3x3s)) |
| 474 | fuse_layers.append(nn.ModuleList(fuse_layer)) |
| 475 | |
| 476 | return nn.ModuleList(fuse_layers) |
| 477 | |
| 478 | def get_num_inchannels(self): |
| 479 | return self.num_inchannels |