| 416 | raise ValueError(error_msg) |
| 417 | |
| 418 | def _make_one_branch(self, branch_index, block, num_blocks, num_channels, stride=1): |
| 419 | downsample = None |
| 420 | if stride != 1 or self.num_inchannels[branch_index] != num_channels[branch_index] * block.expansion: |
| 421 | downsample = nn.Sequential( |
| 422 | nn.Conv2d( |
| 423 | self.num_inchannels[branch_index], num_channels[branch_index] * block.expansion, |
| 424 | kernel_size=1, stride=stride, bias=False), |
| 425 | nn.BatchNorm2d(num_channels[branch_index] * block.expansion, momentum=_BN_MOMENTUM), |
| 426 | ) |
| 427 | |
| 428 | layers = [block(self.num_inchannels[branch_index], num_channels[branch_index], stride, downsample)] |
| 429 | self.num_inchannels[branch_index] = num_channels[branch_index] * block.expansion |
| 430 | for i in range(1, num_blocks[branch_index]): |
| 431 | layers.append(block(self.num_inchannels[branch_index], num_channels[branch_index])) |
| 432 | |
| 433 | return nn.Sequential(*layers) |
| 434 | |
| 435 | def _make_branches(self, num_branches, block, num_blocks, num_channels): |
| 436 | branches = [] |