| 72 | |
| 73 | class DRConv2d(nn.Module): |
| 74 | def __init__(self, in_channels, out_channels, kernel_size, region_num=2, **kwargs): |
| 75 | super(DRConv2d, self).__init__() |
| 76 | self.region_num = 2 |
| 77 | |
| 78 | self.conv_kernel = nn.Sequential( |
| 79 | nn.AdaptiveAvgPool2d((kernel_size, kernel_size)), |
| 80 | nn.Conv2d(in_channels, region_num * region_num, kernel_size=1), |
| 81 | nn.Sigmoid(), |
| 82 | nn.Conv2d(region_num * region_num, region_num * in_channels * out_channels, kernel_size=1, groups=region_num) |
| 83 | ) |
| 84 | self.conv_guide = nn.Conv2d(in_channels, region_num, kernel_size=kernel_size, **kwargs) |
| 85 | |
| 86 | self.corr = Correlation(use_slow=False) |
| 87 | self.kwargs = kwargs |
| 88 | self.act = nn.Sigmoid() |
| 89 | def forward(self, input, mask): |
| 90 | kernel = self.conv_kernel(input) |
| 91 | kernel = kernel.view(kernel.size(0), -1, kernel.size(2), kernel.size(3)) # B x (r*in*out) x W X H |