| 139 | raise ValueError(error_msg) |
| 140 | |
| 141 | def _make_one_branch(self, branch_index, block, num_blocks, num_channels, |
| 142 | stride=1): |
| 143 | downsample = None |
| 144 | if stride != 1 or \ |
| 145 | self.num_inchannels[branch_index] != num_channels[branch_index] * block.expansion: |
| 146 | downsample = nn.Sequential( |
| 147 | nn.Conv2d( |
| 148 | self.num_inchannels[branch_index], |
| 149 | num_channels[branch_index] * block.expansion, |
| 150 | kernel_size=1, stride=stride, bias=False |
| 151 | ), |
| 152 | nn.BatchNorm2d( |
| 153 | num_channels[branch_index] * block.expansion, |
| 154 | momentum=BN_MOMENTUM |
| 155 | ), |
| 156 | ) |
| 157 | |
| 158 | layers = [] |
| 159 | layers.append( |
| 160 | block( |
| 161 | self.num_inchannels[branch_index], |
| 162 | num_channels[branch_index], |
| 163 | stride, |
| 164 | downsample |
| 165 | ) |
| 166 | ) |
| 167 | self.num_inchannels[branch_index] = \ |
| 168 | num_channels[branch_index] * block.expansion |
| 169 | for i in range(1, num_blocks[branch_index]): |
| 170 | layers.append( |
| 171 | block( |
| 172 | self.num_inchannels[branch_index], |
| 173 | num_channels[branch_index] |
| 174 | ) |
| 175 | ) |
| 176 | |
| 177 | return nn.Sequential(*layers) |
| 178 | |
| 179 | def _make_branches(self, num_branches, block, num_blocks, num_channels): |
| 180 | branches = [] |