| 253 | |
| 254 | # Channel attention block (CAB) |
| 255 | class CAB(nn.Module): |
| 256 | def __init__(self, in_channels, out_channels=None, ratio=16, activation='relu'): |
| 257 | super(CAB, self).__init__() |
| 258 | |
| 259 | self.in_channels = in_channels |
| 260 | self.out_channels = out_channels |
| 261 | if self.in_channels < ratio: |
| 262 | ratio = self.in_channels |
| 263 | self.reduced_channels = self.in_channels // ratio |
| 264 | if self.out_channels == None: |
| 265 | self.out_channels = in_channels |
| 266 | |
| 267 | self.avg_pool = nn.AdaptiveAvgPool2d(1) |
| 268 | self.max_pool = nn.AdaptiveMaxPool2d(1) |
| 269 | self.activation = act_layer(activation, inplace=True) |
| 270 | self.fc1 = nn.Conv2d(self.in_channels, self.reduced_channels, 1, bias=False) |
| 271 | self.fc2 = nn.Conv2d(self.reduced_channels, self.out_channels, 1, bias=False) |
| 272 | |
| 273 | self.sigmoid = nn.Sigmoid() |
| 274 | |
| 275 | self.init_weights('normal') |
| 276 | |
| 277 | def init_weights(self, scheme=''): |
| 278 | named_apply(partial(_init_weights, scheme=scheme), self) |
| 279 | |
| 280 | def forward(self, x): |
| 281 | avg_pool_out = self.avg_pool(x) |
| 282 | avg_out = self.fc2(self.activation(self.fc1(avg_pool_out))) |
| 283 | |
| 284 | max_pool_out= self.max_pool(x) |
| 285 | max_out = self.fc2(self.activation(self.fc1(max_pool_out))) |
| 286 | |
| 287 | out = avg_out + max_out |
| 288 | return self.sigmoid(out) |
| 289 | |
| 290 | # Spatial attention block (SAB) |
| 291 | class SAB(nn.Module): |