| 3 | |
| 4 | |
| 5 | class Inception_Block_V1(nn.Module): |
| 6 | def __init__(self, in_channels, out_channels, stride=1, num_kernels=6, init_weight=True): |
| 7 | super(Inception_Block_V1, self).__init__() |
| 8 | self.in_channels = in_channels |
| 9 | self.out_channels = out_channels |
| 10 | self.num_kernels = num_kernels |
| 11 | self.stride = stride |
| 12 | kernels = [] |
| 13 | for i in range(self.num_kernels): |
| 14 | kernels.append(nn.Conv2d(in_channels, out_channels, kernel_size=2 * i + 1, padding=i, stride=stride)) |
| 15 | self.kernels = nn.ModuleList(kernels) |
| 16 | if init_weight: |
| 17 | self._initialize_weights() |
| 18 | |
| 19 | def _initialize_weights(self): |
| 20 | for m in self.modules(): |
| 21 | if isinstance(m, nn.Conv2d): |
| 22 | nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu') |
| 23 | if m.bias is not None: |
| 24 | nn.init.constant_(m.bias, 0) |
| 25 | |
| 26 | def forward(self, x): |
| 27 | res_list = [] |
| 28 | for i in range(self.num_kernels): |
| 29 | res_list.append(self.kernels[i](x)) |
| 30 | res = torch.stack(res_list, dim=-1).mean(-1) |
| 31 | return res |
| 32 | |
| 33 | |
| 34 | class Inception_Trans_Block_V1(nn.Module): |
nothing calls this directly
no outgoing calls
no test coverage detected