| 32 | |
| 33 | |
| 34 | class Inception_Trans_Block_V1(nn.Module): |
| 35 | def __init__(self, in_channels, out_channels, stride=1, num_kernels=6, init_weight=True): |
| 36 | super(Inception_Trans_Block_V1, self).__init__() |
| 37 | self.in_channels = in_channels |
| 38 | self.out_channels = out_channels |
| 39 | self.num_kernels = num_kernels |
| 40 | self.stride = stride |
| 41 | |
| 42 | kernels = [] |
| 43 | for i in range(self.num_kernels): |
| 44 | kernels.append( |
| 45 | nn.ConvTranspose2d(in_channels, out_channels, kernel_size=2 * i + 1, padding=i, stride=stride)) |
| 46 | self.kernels = nn.ModuleList(kernels) |
| 47 | if init_weight: |
| 48 | self._initialize_weights() |
| 49 | |
| 50 | def _initialize_weights(self): |
| 51 | for m in self.modules(): |
| 52 | if isinstance(m, nn.ConvTranspose2d): |
| 53 | nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu') |
| 54 | if m.bias is not None: |
| 55 | nn.init.constant_(m.bias, 0) |
| 56 | |
| 57 | def forward(self, x, output_size): |
| 58 | res_list = [] |
| 59 | for i in range(self.num_kernels): |
| 60 | res_list.append(self.kernels[i](x, output_size=output_size)) |
| 61 | res = torch.stack(res_list, dim=-1).mean(-1) |
| 62 | return res |
nothing calls this directly
no outgoing calls
no test coverage detected