import torch import torch.nn as nn import torch.nn.functional as F import math """ Encoder should take a batch of image, return a batch of embedding """ class codegeneration(torch.nn.Module): def __init__(self): super(codegeneration, self).__init__() self.conv1 = nn.Sequential(nn.Conv2d(3,64,7,2,3, bias=True), nn.LeakyReLU(negative_slope=0.1), nn.Conv2d(64,64,3,1,1, bias=True), nn.LeakyReLU(negative_slope=0.1)) self.layer1 = nn.Sequential(nn.Conv2d(64,64,3,1,1, bias=True), nn.LeakyReLU(negative_slope=0.1), selfattention(64), nn.Conv2d(64,64,3,1,1, bias=True), nn.LeakyReLU(negative_slope=0.1)) #64 self.layer2_1 = nn.Sequential(nn.Conv2d(64,128,3,2,1, bias=True), nn.LeakyReLU(negative_slope=0.1), selfattention(128), nn.Conv2d(128,128,3,1,1, bias=True), nn.LeakyReLU(negative_slope=0.1),) #64 self.resblock1 = BasicBlockNormal(128,128) self.resblock2 = BasicBlockNormal(128,128) self.layer2_2 = nn.Sequential(nn.Conv2d(128,128,3,2,1, bias=True), nn.LeakyReLU(negative_slope=0.1), nn.Conv2d(128,128,3,1,1, bias=True), nn.LeakyReLU(negative_slope=0.1),) #64 self.layer3_1 = nn.Sequential(nn.Conv2d(128,256,3,2,1, bias=True), nn.LeakyReLU(negative_slope=0.1), nn.Conv2d(256,256,3,1,1, bias=True), nn.LeakyReLU(negative_slope=0.1),) #64 self.layer3_2 = nn.Sequential(nn.Conv2d(256,256,3,1,1, bias=True), # stride 2 for 128x128 nn.LeakyReLU(negative_slope=0.1), nn.Conv2d(256,128,3,1,1, bias=True), nn.LeakyReLU(negative_slope=0.1)) #64 self.expresscode = nn.Sequential(nn.Linear(2048,512), nn.LeakyReLU(negative_slope=0.1), nn.Linear(512,256)) for m in self.modules(): if isinstance(m, nn.Conv2d): n = m.kernel_size[0] * m.kernel_size[1] * m.out_channels m.weight.data.normal_(0, math.sqrt(2. / n)) if m.bias is not None: m.bias.data.zero_() elif isinstance(m, nn.BatchNorm2d): m.weight.data.fill_(1) m.bias.data.zero_() elif isinstance(m, nn.Linear): m.weight.data.normal_(0, 0.01) m.bias.data.zero_() def forward(self, x): #encoder out_1 = self.conv1(x) out_1 = self.layer1(out_1) out_2 = self.layer2_1(out_1) out_2 = self.resblock1(out_2) out_2 = self.resblock2(out_2) out_2 = self.layer2_2(out_2) out_3 = self.layer3_1(out_2) out_3 = self.layer3_2(out_3) out_3 = out_3.view(x.size()[0],-1) expcode = self.expresscode(out_3) expcode = expcode.view(x.size()[0],-1,1,1) expcode = F.tanh(expcode) return expcode class BasicBlockNormal(nn.Module): expansion = 1 def __init__(self, inplanes, planes, stride=1, downsample=None): super(BasicBlockNormal, self).__init__() # Both self.conv1 and self.downsample layers downsample the input when stride != 1 self.conv1 = nn.Conv2d(inplanes,planes,3,stride,1) self.relu = nn.LeakyReLU(negative_slope=0.1,inplace=True) self.conv2 = nn.Conv2d(planes,planes,3,1,1) self.downsample = downsample self.stride = stride def forward(self, x): identity = x out = self.conv1(x) out = self.relu(out) out = self.conv2(out) #out = self.relu(out) if self.downsample is not None: identity = self.downsample(x) out = (out + identity) return self.relu(out) class selfattention(nn.Module): def __init__(self, inplanes): super(selfattention, self).__init__() self.interchannel = inplanes self.inplane = inplanes self.g = nn.Conv2d(inplanes, inplanes, kernel_size=1, stride=1, padding=0) self.theta = nn.Conv2d(inplanes, self.interchannel, kernel_size=1, stride=1, padding=0) self.phi = nn.Conv2d(inplanes, self.interchannel, kernel_size=1, stride=1, padding=0) self.act = nn.LeakyReLU(0.1) def forward(self, x): b,c,h,w = x.size() g_y = self.g(x).view(b, c, -1) #BXcXN theta_x = self.theta(x).view(b, self.interchannel, -1) theta_x = F.softmax(theta_x, dim = -1) # softmax on N theta_x = theta_x.permute(0,2,1).contiguous() #BXNXC' phi_x = self.phi(x).view(b, self.interchannel, -1) #BXC'XN similarity = torch.bmm(phi_x, theta_x) #BXc'Xc' g_y = F.softmax(g_y, dim = 1) attention = torch.bmm(similarity, g_y) #BXCXN attention = attention.view(b,c,h,w).contiguous() y = self.act(x + attention) return y