Spaces:
No application file
No application file
File size: 5,791 Bytes
a5deb1f | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 | 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
|