| 60 | N X C X H X W |
| 61 | ''' |
| 62 | def __init__(self, in_channels, key_channels, scale=1): |
| 63 | super(ObjectAttentionBlock, self).__init__() |
| 64 | self.scale = scale |
| 65 | self.in_channels = in_channels |
| 66 | self.key_channels = key_channels |
| 67 | self.pool = nn.MaxPool2d(kernel_size=(scale, scale)) |
| 68 | self.f_pixel = nn.Sequential( |
| 69 | nn.Conv2d(in_channels=self.in_channels, out_channels=self.key_channels, |
| 70 | kernel_size=1, stride=1, padding=0, bias=False), |
| 71 | BNReLU(self.key_channels), |
| 72 | nn.Conv2d(in_channels=self.key_channels, out_channels=self.key_channels, |
| 73 | kernel_size=1, stride=1, padding=0, bias=False), |
| 74 | BNReLU(self.key_channels), |
| 75 | ) |
| 76 | self.f_object = nn.Sequential( |
| 77 | nn.Conv2d(in_channels=self.in_channels, out_channels=self.key_channels, |
| 78 | kernel_size=1, stride=1, padding=0, bias=False), |
| 79 | BNReLU(self.key_channels), |
| 80 | nn.Conv2d(in_channels=self.key_channels, out_channels=self.key_channels, |
| 81 | kernel_size=1, stride=1, padding=0, bias=False), |
| 82 | BNReLU(self.key_channels), |
| 83 | ) |
| 84 | self.f_down = nn.Sequential( |
| 85 | nn.Conv2d(in_channels=self.in_channels, out_channels=self.key_channels, |
| 86 | kernel_size=1, stride=1, padding=0, bias=False), |
| 87 | BNReLU(self.key_channels), |
| 88 | ) |
| 89 | self.f_up = nn.Sequential( |
| 90 | nn.Conv2d(in_channels=self.key_channels, out_channels=self.in_channels, |
| 91 | kernel_size=1, stride=1, padding=0, bias=False), |
| 92 | BNReLU(self.in_channels), |
| 93 | ) |
| 94 | |
| 95 | def forward(self, x, proxy): |
| 96 | batch_size, h, w = x.size(0), x.size(2), x.size(3) |