| 38 | |
| 39 | |
| 40 | class AttING(nn.Module): |
| 41 | def __init__(self, in_channels, channels): |
| 42 | super(AttING, self).__init__() |
| 43 | self.conv1 = nn.Conv2d(in_channels, channels, kernel_size=1, stride=1, padding=0, bias=False) |
| 44 | self.conv2_1 = nn.Conv2d(channels, channels, kernel_size=3, stride=1, padding=1, bias=False) |
| 45 | self.conv2_2 = nn.Conv2d(channels, channels, kernel_size=3, stride=1, padding=1, bias=False) |
| 46 | self.instance = nn.InstanceNorm2d(channels, affine=True) |
| 47 | self.interative = nn.Sequential( |
| 48 | nn.Conv2d(channels*2, channels, kernel_size=3, stride=1, padding=1, bias=False), |
| 49 | nn.LeakyReLU(0.1), |
| 50 | nn.Sigmoid() |
| 51 | ) |
| 52 | self.act = nn.LeakyReLU(0.1) |
| 53 | self.avgpool = nn.AdaptiveAvgPool2d(1) |
| 54 | self.contrast = stdv_channels |
| 55 | self.process = nn.Sequential(nn.Conv2d(channels*2, channels//2, kernel_size=3, padding=1, bias=True), |
| 56 | nn.LeakyReLU(0.1), |
| 57 | nn.Conv2d(channels//2, channels*2, kernel_size=3, padding=1, bias=True), |
| 58 | nn.Sigmoid()) |
| 59 | self.conv1x1 = nn.Conv2d(2*channels, channels, 1, 1, 0) |
| 60 | |
| 61 | def forward(self, x): |
| 62 | x1 = self.conv1(x) |
| 63 | out_instance = self.instance(x1) |
| 64 | out_identity = x1 |
| 65 | # feature_save(out_identity, '1') |
| 66 | # feature_save(out_instance, '2') |
| 67 | out1 = self.conv2_1(out_instance) |
| 68 | out2 = self.conv2_2(out_identity) |
| 69 | out = torch.cat((out1, out2), 1) |
| 70 | xp1 = self.interative(out)*out2 + out1 |
| 71 | xp2 = (1-self.interative(out))*out1 + out2 |
| 72 | xp = torch.cat((xp1, xp2), 1) |
| 73 | xp = self.process(self.contrast(xp)+self.avgpool(xp))*xp |
| 74 | xp = self.conv1x1(xp) |
| 75 | xout = xp |
| 76 | |
| 77 | return xout,out_instance |
| 78 | |
| 79 | |
| 80 | |