| 199 | } |
| 200 | |
| 201 | class HighResolutionNet(nn.Module): |
| 202 | |
| 203 | def __init__(self, cfg, **kwargs): |
| 204 | self.inplanes = 64 |
| 205 | super(HighResolutionNet, self).__init__() |
| 206 | use_old_impl = cfg.get('use_old_impl') |
| 207 | self.use_old_impl = use_old_impl |
| 208 | |
| 209 | # stem net |
| 210 | self.conv1 = nn.Conv2d(3, 64, kernel_size=3, stride=2, padding=1, |
| 211 | bias=False) |
| 212 | self.bn1 = nn.BatchNorm2d(64, momentum=BN_MOMENTUM) |
| 213 | self.conv2 = nn.Conv2d(64, 64, kernel_size=3, stride=2, padding=1, |
| 214 | bias=False) |
| 215 | self.bn2 = nn.BatchNorm2d(64, momentum=BN_MOMENTUM) |
| 216 | self.relu = nn.ReLU(inplace=True) |
| 217 | |
| 218 | self.stage1_cfg = cfg.get('stage1', {}) |
| 219 | num_channels = self.stage1_cfg['num_channels'][0] |
| 220 | block = blocks_dict[self.stage1_cfg['block']] |
| 221 | num_blocks = self.stage1_cfg['num_blocks'][0] |
| 222 | self.layer1 = self._make_layer(block, num_channels, num_blocks) |
| 223 | stage1_out_channel = block.expansion * num_channels |
| 224 | |
| 225 | self.stage2_cfg = cfg.get('stage2', {}) |
| 226 | num_channels = self.stage2_cfg.get('num_channels', (32, 64)) |
| 227 | block = blocks_dict[self.stage2_cfg.get('block')] |
| 228 | num_channels = [ |
| 229 | num_channels[i] * block.expansion for i in range(len(num_channels)) |
| 230 | ] |
| 231 | stage2_num_channels = num_channels |
| 232 | self.transition1 = self._make_transition_layer( |
| 233 | [stage1_out_channel], num_channels) |
| 234 | self.stage2, pre_stage_channels = self._make_stage( |
| 235 | self.stage2_cfg, num_channels) |
| 236 | |
| 237 | self.stage3_cfg = cfg.get('stage3') |
| 238 | num_channels = self.stage3_cfg['num_channels'] |
| 239 | block = blocks_dict[self.stage3_cfg['block']] |
| 240 | num_channels = [ |
| 241 | num_channels[i] * block.expansion for i in range(len(num_channels)) |
| 242 | ] |
| 243 | stage3_num_channels = num_channels |
| 244 | self.transition2 = self._make_transition_layer( |
| 245 | pre_stage_channels, num_channels) |
| 246 | self.stage3, pre_stage_channels = self._make_stage( |
| 247 | self.stage3_cfg, num_channels) |
| 248 | |
| 249 | self.stage4_cfg = cfg.get('stage4') |
| 250 | num_channels = self.stage4_cfg['num_channels'] |
| 251 | block = blocks_dict[self.stage4_cfg['block']] |
| 252 | num_channels = [ |
| 253 | num_channels[i] * block.expansion for i in range(len(num_channels)) |
| 254 | ] |
| 255 | self.transition3 = self._make_transition_layer( |
| 256 | pre_stage_channels, num_channels) |
| 257 | stage_4_out_channels = num_channels |
| 258 | |