| 158 | |
| 159 | |
| 160 | class ETM(nn.Module): |
| 161 | def __init__(self, in_channels, out_channels): |
| 162 | super(ETM, self).__init__() |
| 163 | self.relu = nn.ReLU(True) |
| 164 | self.branch0 = BasicConv2d(in_channels, out_channels, 1) |
| 165 | self.branch1 = nn.Sequential( |
| 166 | BasicConv2d(in_channels, out_channels, 1), |
| 167 | BasicConv2d(out_channels, out_channels, kernel_size=(1, 3), padding=(0, 1)), |
| 168 | BasicConv2d(out_channels, out_channels, kernel_size=(3, 1), padding=(1, 0)), |
| 169 | BasicConv2d(out_channels, out_channels, 3, padding=3, dilation=3) |
| 170 | ) |
| 171 | self.branch2 = nn.Sequential( |
| 172 | BasicConv2d(in_channels, out_channels, 1), |
| 173 | BasicConv2d(out_channels, out_channels, kernel_size=(1, 5), padding=(0, 2)), |
| 174 | BasicConv2d(out_channels, out_channels, kernel_size=(5, 1), padding=(2, 0)), |
| 175 | BasicConv2d(out_channels, out_channels, 3, padding=5, dilation=5) |
| 176 | ) |
| 177 | self.branch3 = nn.Sequential( |
| 178 | BasicConv2d(in_channels, out_channels, 1), |
| 179 | BasicConv2d(out_channels, out_channels, kernel_size=(1, 7), padding=(0, 3)), |
| 180 | BasicConv2d(out_channels, out_channels, kernel_size=(7, 1), padding=(3, 0)), |
| 181 | BasicConv2d(out_channels, out_channels, 3, padding=7, dilation=7) |
| 182 | ) |
| 183 | self.conv_cat = BasicConv2d(4 * out_channels, out_channels, 3, padding=1) |
| 184 | self.conv_res = BasicConv2d(in_channels, out_channels, 1) |
| 185 | |
| 186 | def forward(self, x): |
| 187 | x0 = self.branch0(x) |
| 188 | x1 = self.branch1(x) |
| 189 | x2 = self.branch2(x) |
| 190 | x3 = self.branch3(x) |
| 191 | x_cat = self.conv_cat(torch.cat((x0, x1, x2, x3), 1)) |
| 192 | |
| 193 | x = self.relu(x_cat + self.conv_res(x)) |
| 194 | return x |
| 195 | |
| 196 | |
| 197 | |