(self, input, mask)
| 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 |
| 92 | output = self.corr(input, kernel, **self.kwargs) # B x (r*out) x W x H |
| 93 | output = output.view(output.size(0), self.region_num, -1, output.size(2), output.size(3)) # B x r x out x W x H |
| 94 | |
| 95 | mask = F.interpolate(mask.detach(), size=input.size()[2:], mode='nearest') |
| 96 | mask = mask.unsqueeze(1) |
| 97 | inv_msak = 1 - mask |
| 98 | guide_mask = torch.cat((mask, inv_msak), 1) |
| 99 | |
| 100 | output = torch.sum(output * guide_mask, dim=1) |
| 101 | return output |
| 102 | |
| 103 |
nothing calls this directly
no outgoing calls
no test coverage detected