| 60 | return init_fun |
| 61 | |
| 62 | class PartialConv(nn.Module): |
| 63 | def __init__(self, in_channels, out_channels, kernel_size, stride=1, |
| 64 | padding=0, dilation=1, groups=1, bias=True): |
| 65 | super().__init__() |
| 66 | self.input_conv = nn.Conv2d(in_channels, out_channels, kernel_size, |
| 67 | stride, padding, dilation, groups, bias) |
| 68 | self.mask_conv = nn.Conv2d(in_channels, out_channels, kernel_size, |
| 69 | stride, padding, dilation, groups, False) |
| 70 | self.input_conv.apply(weights_init('kaiming')) |
| 71 | self.slide_winsize = in_channels * kernel_size * kernel_size |
| 72 | |
| 73 | torch.nn.init.constant_(self.mask_conv.weight, 1.0) |
| 74 | |
| 75 | # mask is not updated |
| 76 | for param in self.mask_conv.parameters(): |
| 77 | param.requires_grad = False |
| 78 | |
| 79 | def forward(self, input, mask): |
| 80 | # http://masc.cs.gmu.edu/wiki/partialconv |
| 81 | # C(X) = W^T * X + b, C(0) = b, D(M) = 1 * M + 0 = sum(M) |
| 82 | # W^T* (M .* X) / sum(M) + b = [C(M .* X) – C(0)] / D(M) + C(0) |
| 83 | output = self.input_conv(input * mask) |
| 84 | if self.input_conv.bias is not None: |
| 85 | output_bias = self.input_conv.bias.view(1, -1, 1, 1).expand_as( |
| 86 | output) |
| 87 | else: |
| 88 | output_bias = torch.zeros_like(output) |
| 89 | |
| 90 | with torch.no_grad(): |
| 91 | output_mask = self.mask_conv(mask) |
| 92 | |
| 93 | no_update_holes = output_mask == 0 |
| 94 | |
| 95 | mask_sum = output_mask.masked_fill_(no_update_holes, 1.0) |
| 96 | |
| 97 | output_pre = ((output - output_bias) * self.slide_winsize) / mask_sum + output_bias |
| 98 | output = output_pre.masked_fill_(no_update_holes, 0.0) |
| 99 | |
| 100 | new_mask = torch.ones_like(output) |
| 101 | new_mask = new_mask.masked_fill_(no_update_holes, 0.0) |
| 102 | |
| 103 | return output, new_mask |
| 104 | |
| 105 | |
| 106 | class PCBActiv(nn.Module): |