| 58 | |
| 59 | |
| 60 | class BottleneckBlock(nn.Module): |
| 61 | def __init__(self, in_planes, planes, norm_fn='group', stride=1): |
| 62 | super(BottleneckBlock, self).__init__() |
| 63 | |
| 64 | self.conv1 = nn.Conv2d(in_planes, planes//4, kernel_size=1, padding=0) |
| 65 | self.conv2 = nn.Conv2d(planes//4, planes//4, kernel_size=3, padding=1, stride=stride) |
| 66 | self.conv3 = nn.Conv2d(planes//4, planes, kernel_size=1, padding=0) |
| 67 | self.relu = nn.ReLU(inplace=True) |
| 68 | |
| 69 | num_groups = planes // 8 |
| 70 | |
| 71 | if norm_fn == 'group': |
| 72 | self.norm1 = nn.GroupNorm(num_groups=num_groups, num_channels=planes//4) |
| 73 | self.norm2 = nn.GroupNorm(num_groups=num_groups, num_channels=planes//4) |
| 74 | self.norm3 = nn.GroupNorm(num_groups=num_groups, num_channels=planes) |
| 75 | if not stride == 1: |
| 76 | self.norm4 = nn.GroupNorm(num_groups=num_groups, num_channels=planes) |
| 77 | |
| 78 | elif norm_fn == 'batch': |
| 79 | self.norm1 = nn.BatchNorm2d(planes//4) |
| 80 | self.norm2 = nn.BatchNorm2d(planes//4) |
| 81 | self.norm3 = nn.BatchNorm2d(planes) |
| 82 | if not stride == 1: |
| 83 | self.norm4 = nn.BatchNorm2d(planes) |
| 84 | |
| 85 | elif norm_fn == 'instance': |
| 86 | self.norm1 = nn.InstanceNorm2d(planes//4) |
| 87 | self.norm2 = nn.InstanceNorm2d(planes//4) |
| 88 | self.norm3 = nn.InstanceNorm2d(planes) |
| 89 | if not stride == 1: |
| 90 | self.norm4 = nn.InstanceNorm2d(planes) |
| 91 | |
| 92 | elif norm_fn == 'none': |
| 93 | self.norm1 = nn.Sequential() |
| 94 | self.norm2 = nn.Sequential() |
| 95 | self.norm3 = nn.Sequential() |
| 96 | if not stride == 1: |
| 97 | self.norm4 = nn.Sequential() |
| 98 | |
| 99 | if stride == 1: |
| 100 | self.downsample = None |
| 101 | |
| 102 | else: |
| 103 | self.downsample = nn.Sequential( |
| 104 | nn.Conv2d(in_planes, planes, kernel_size=1, stride=stride), self.norm4) |
| 105 | |
| 106 | |
| 107 | def forward(self, x): |
| 108 | y = x |
| 109 | y = self.relu(self.norm1(self.conv1(y))) |
| 110 | y = self.relu(self.norm2(self.conv2(y))) |
| 111 | y = self.relu(self.norm3(self.conv3(y))) |
| 112 | |
| 113 | if self.downsample is not None: |
| 114 | x = self.downsample(x) |
| 115 | |
| 116 | return self.relu(x+y) |
| 117 | |