GoodWin/Deep-Multi-scale
0
1import torch2import torch.nn as nn3from torch.nn import init4import functools5from torch.optim import lr_scheduler6import torch.nn.functional as F7from torch.nn import Parameter as P8from util import util9from torchvision import models10import scipy.io as sio11import numpy as np12import scipy.ndimage13import torch.nn.utils.spectral_norm as SpectralNorm14 15from torch.autograd import Function16from math import sqrt17import random18import os19import math20 21from sync_batchnorm import convert_model22####23 24###############################################################################25# Helper Functions26###############################################################################27def get_norm_layer(norm_type='instance'):28 if norm_type == 'batch':29 norm_layer = functools.partial(nn.BatchNorm2d, affine=True)30 elif norm_type == 'instance':31 norm_layer = functools.partial(nn.InstanceNorm2d, affine=False, track_running_stats=True)32 elif norm_type == 'none':33 norm_layer = None34 else:35 raise NotImplementedError('normalization layer [%s] is not found' % norm_type)36 37 return norm_layer38 39 40def get_scheduler(optimizer, opt):41 if opt.lr_policy == 'lambda':42 def lambda_rule(epoch):43 lr_l = 1.0 - max(0, epoch + 1 + opt.epoch_count - opt.niter) / float(opt.niter_decay + 1)44 return lr_l45 scheduler = lr_scheduler.LambdaLR(optimizer, lr_lambda=lambda_rule)46 elif opt.lr_policy == 'step':47 scheduler = lr_scheduler.StepLR(optimizer, step_size=opt.lr_decay_iters, gamma=0.1)48 elif opt.lr_policy == 'plateau':49 scheduler = lr_scheduler.ReduceLROnPlateau(optimizer, mode='min', factor=0.2, threshold=0.01, patience=5)50 else:51 return NotImplementedError('learning rate policy [%s] is not implemented', opt.lr_policy)52 53 return scheduler54 55 56def init_weights(net, init_type='normal', gain=0.02):57 def init_func(m):58 classname = m.__class__.__name__59 if hasattr(m, 'weight') and (classname.find('Conv') != -1 or classname.find('Linear') != -1):60 if init_type == 'normal':61 init.normal_(m.weight.data, 0.0, gain)62 elif init_type == 'xavier':63 init.xavier_normal_(m.weight.data, gain=gain)64 elif init_type == 'kaiming':65 init.kaiming_normal_(m.weight.data, a=0, mode='fan_in')66 elif init_type == 'orthogonal':67 init.orthogonal_(m.weight.data, gain=gain)68 else:69 raise NotImplementedError('initialization method [%s] is not implemented' % init_type)70 if hasattr(m, 'bias') and m.bias is not None:71 init.constant_(m.bias.data, 0.0)72 elif classname.find('BatchNorm2d') != -1:73 init.normal_(m.weight.data, 1.0, gain)74 init.constant_(m.bias.data, 0.0)75 76 print('initialize network with %s' % init_type)77 net.apply(init_func)78 79 80def init_net(net, init_type='normal', init_gain=0.02, gpu_ids=[], init_flag=True):81 if len(gpu_ids) > 0:82 assert(torch.cuda.is_available())83 net = convert_model(net)84 net.to(gpu_ids[0])85 net = torch.nn.DataParallel(net, gpu_ids)86 87 if init_flag:88 89 init_weights(net, init_type, gain=init_gain)90 91 return net92 93 94# compute adaptive instance norm95def calc_mean_std(feat, eps=1e-5):96 # eps is a small value added to the variance to avoid divide-by-zero.97 size = feat.size()98 assert (len(size) == 3)99 C, _ = size[:2]100 feat_var = feat.contiguous().view(C, -1).var(dim=1) + eps101 feat_std = feat_var.sqrt().view(C, 1, 1)102 feat_mean = feat.contiguous().view(C, -1).mean(dim=1).view(C, 1, 1)103 104 return feat_mean, feat_std105 106 107def adaptive_instance_normalization(content_feat, style_feat): # content_feat is degraded feature, style is ref feature108 assert (content_feat.size()[:1] == style_feat.size()[:1])109 size = content_feat.size()110 style_mean, style_std = calc_mean_std(style_feat)111 content_mean, content_std = calc_mean_std(content_feat)112 113 normalized_feat = (content_feat - content_mean.expand(114 size)) / content_std.expand(size)115 116 return normalized_feat * style_std.expand(size) + style_mean.expand(size)117 118def calc_mean_std_4D(feat, eps=1e-5):119 # eps is a small value added to the variance to avoid divide-by-zero.120 size = feat.size()121 assert (len(size) == 4)122 N, C = size[:2]123 feat_var = feat.view(N, C, -1).var(dim=2) + eps124 feat_std = feat_var.sqrt().view(N, C, 1, 1)125 feat_mean = feat.view(N, C, -1).mean(dim=2).view(N, C, 1, 1)126 return feat_mean, feat_std127 128def adaptive_instance_normalization_4D(content_feat, style_feat): # content_feat is ref feature, style is degradate feature129 # assert (content_feat.size()[:2] == style_feat.size()[:2])130 size = content_feat.size()131 style_mean, style_std = calc_mean_std_4D(style_feat)132 133 content_mean, content_std = calc_mean_std_4D(content_feat)134 normalized_feat = (content_feat - content_mean.expand(135 size)) / content_std.expand(size)136 return normalized_feat * style_std.expand(size) + style_mean.expand(size)137 138def define_G(which_model_netG, gpu_ids=[]):139 if which_model_netG == 'UNetDictFace':140 netG = UNetDictFace(64)141 init_flag = False142 else:143 raise NotImplementedError('Generator model name [%s] is not recognized' % which_model_netG)144 return init_net(netG, 'normal', 0.02, gpu_ids, init_flag)145 146 147##############################################################################148# Classes149############################################################################################################################################150 151 152def convU(in_channels, out_channels,conv_layer, norm_layer, kernel_size=3, stride=1,dilation=1, bias=True):153 return nn.Sequential(154 SpectralNorm(conv_layer(in_channels, out_channels, kernel_size=kernel_size, stride=stride, dilation=dilation, padding=((kernel_size-1)//2)*dilation, bias=bias)),155# conv_layer(in_channels, out_channels, kernel_size=kernel_size, stride=stride, dilation=dilation, padding=((kernel_size-1)//2)*dilation, bias=bias),156# nn.BatchNorm2d(out_channels),157 nn.LeakyReLU(0.2),158 SpectralNorm(conv_layer(out_channels, out_channels, kernel_size=kernel_size, stride=stride, dilation=dilation, padding=((kernel_size-1)//2)*dilation, bias=bias)),159 )160class MSDilateBlock(nn.Module):161 def __init__(self, in_channels,conv_layer=nn.Conv2d, norm_layer=nn.BatchNorm2d, kernel_size=3, dilation=[1,1,1,1], bias=True):162 super(MSDilateBlock, self).__init__()163 self.conv1 = convU(in_channels, in_channels,conv_layer, norm_layer, kernel_size,dilation=dilation[0], bias=bias)164 self.conv2 = convU(in_channels, in_channels,conv_layer, norm_layer, kernel_size,dilation=dilation[1], bias=bias)165 self.conv3 = convU(in_channels, in_channels,conv_layer, norm_layer, kernel_size,dilation=dilation[2], bias=bias)166 self.conv4 = convU(in_channels, in_channels,conv_layer, norm_layer, kernel_size,dilation=dilation[3], bias=bias)167 self.convi = SpectralNorm(conv_layer(in_channels*4, in_channels, kernel_size=kernel_size, stride=1, padding=(kernel_size-1)//2, bias=bias))168 def forward(self, x):169 conv1 = self.conv1(x)170 conv2 = self.conv2(x)171 conv3 = self.conv3(x)172 conv4 = self.conv4(x)173 cat = torch.cat([conv1, conv2, conv3, conv4], 1)174 out = self.convi(cat) + x175 return out176 177##############################UNetFace#########################178class AdaptiveInstanceNorm(nn.Module):179 def __init__(self, in_channel):180 super().__init__()181 self.norm = nn.InstanceNorm2d(in_channel)182 183 def forward(self, input, style):184 style_mean, style_std = calc_mean_std_4D(style)185 out = self.norm(input)186 size = input.size()187 out = style_std.expand(size) * out + style_mean.expand(size)188 return out189 190class BlurFunctionBackward(Function):191 @staticmethod192 def forward(ctx, grad_output, kernel, kernel_flip):193 ctx.save_for_backward(kernel, kernel_flip)194 195 grad_input = F.conv2d(196 grad_output, kernel_flip, padding=1, groups=grad_output.shape[1]197 )198 return grad_input199 200 @staticmethod201 def backward(ctx, gradgrad_output):202 kernel, kernel_flip = ctx.saved_tensors203 204 grad_input = F.conv2d(205 gradgrad_output, kernel, padding=1, groups=gradgrad_output.shape[1]206 )207 return grad_input, None, None208 209 210class BlurFunction(Function):211 @staticmethod212 def forward(ctx, input, kernel, kernel_flip):213 ctx.save_for_backward(kernel, kernel_flip)214 215 output = F.conv2d(input, kernel, padding=1, groups=input.shape[1])216 217 return output218 219 @staticmethod220 def backward(ctx, grad_output):221 kernel, kernel_flip = ctx.saved_tensors222 223 grad_input = BlurFunctionBackward.apply(grad_output, kernel, kernel_flip)224 225 return grad_input, None, None226 227blur = BlurFunction.apply228 229 230class Blur(nn.Module):231 def __init__(self, channel):232 super().__init__()233 234 weight = torch.tensor([[1, 2, 1], [2, 4, 2], [1, 2, 1]], dtype=torch.float32)235 weight = weight.view(1, 1, 3, 3)236 weight = weight / weight.sum()237 weight_flip = torch.flip(weight, [2, 3])238 239 self.register_buffer('weight', weight.repeat(channel, 1, 1, 1))240 self.register_buffer('weight_flip', weight_flip.repeat(channel, 1, 1, 1))241 242 def forward(self, input):243 return blur(input, self.weight, self.weight_flip)244 245class EqualLR:246 def __init__(self, name):247 self.name = name248 249 def compute_weight(self, module):250 weight = getattr(module, self.name + '_orig')251 fan_in = weight.data.size(1) * weight.data[0][0].numel()252 return weight * sqrt(2 / fan_in)253 @staticmethod254 def apply(module, name):255 fn = EqualLR(name)256 257 weight = getattr(module, name)258 del module._parameters[name]259 module.register_parameter(name + '_orig', nn.Parameter(weight.data))260 module.register_forward_pre_hook(fn)261 262 return fn263 264 def __call__(self, module, input):265 weight = self.compute_weight(module)266 setattr(module, self.name, weight)267 268def equal_lr(module, name='weight'):269 EqualLR.apply(module, name)270 return module271 272class EqualConv2d(nn.Module):273 def __init__(self, *args, **kwargs):274 super().__init__()275 conv = nn.Conv2d(*args, **kwargs)276 conv.weight.data.normal_()277 conv.bias.data.zero_()278 self.conv = equal_lr(conv)279 def forward(self, input):280 return self.conv(input)281 282class NoiseInjection(nn.Module):283 def __init__(self, channel):284 super().__init__()285 self.weight = nn.Parameter(torch.zeros(1, channel, 1, 1))286 def forward(self, image, noise):287 return image + self.weight * noise288 289class StyledUpBlock(nn.Module):290 def __init__(self, in_channel, out_channel, kernel_size=3, padding=1,upsample=False):291 super().__init__()292 if upsample:293 self.conv1 = nn.Sequential(294 nn.Upsample(scale_factor=2, mode='bilinear', align_corners=False),295 Blur(out_channel),296 # EqualConv2d(in_channel, out_channel, kernel_size, padding=padding),297 SpectralNorm(nn.Conv2d(in_channel, out_channel, kernel_size, padding=padding)),298 nn.LeakyReLU(0.2),299 )300 else:301 self.conv1 = nn.Sequential(302 Blur(in_channel),303 # EqualConv2d(in_channel, out_channel, kernel_size, padding=padding)304 SpectralNorm(nn.Conv2d(in_channel, out_channel, kernel_size, padding=padding)),305 nn.LeakyReLU(0.2),306 )307 self.convup = nn.Sequential(308 nn.Upsample(scale_factor=2, mode='bilinear', align_corners=False),309 # EqualConv2d(out_channel, out_channel, kernel_size, padding=padding),310 SpectralNorm(nn.Conv2d(out_channel, out_channel, kernel_size, padding=padding)),311 nn.LeakyReLU(0.2),312 # Blur(out_channel),313 )314 # self.noise1 = equal_lr(NoiseInjection(out_channel))315 # self.adain1 = AdaptiveInstanceNorm(out_channel)316 self.lrelu1 = nn.LeakyReLU(0.2)317 318 # self.conv2 = EqualConv2d(out_channel, out_channel, kernel_size, padding=padding)319 # self.noise2 = equal_lr(NoiseInjection(out_channel))320 # self.adain2 = AdaptiveInstanceNorm(out_channel)321 # self.lrelu2 = nn.LeakyReLU(0.2)322 323 self.ScaleModel1 = nn.Sequential(324 # Blur(in_channel),325 SpectralNorm(nn.Conv2d(in_channel,out_channel,3, 1, 1)),326 # nn.Conv2d(in_channel,out_channel,3, 1, 1),327 nn.LeakyReLU(0.2, True),328 SpectralNorm(nn.Conv2d(out_channel, out_channel, 3, 1, 1))329 # nn.Conv2d(out_channel, out_channel, 3, 1, 1)330 )331 self.ShiftModel1 = nn.Sequential(332 # Blur(in_channel),333 SpectralNorm(nn.Conv2d(in_channel,out_channel,3, 1, 1)),334 # nn.Conv2d(in_channel,out_channel,3, 1, 1),335 nn.LeakyReLU(0.2, True),336 SpectralNorm(nn.Conv2d(out_channel, out_channel, 3, 1, 1)),337 nn.Sigmoid(),338 # nn.Conv2d(out_channel, out_channel, 3, 1, 1)339 )340 341 def forward(self, input, style):342 out = self.conv1(input)343# out = self.noise1(out, noise)344 out = self.lrelu1(out)345 346 Shift1 = self.ShiftModel1(style)347 Scale1 = self.ScaleModel1(style)348 out = out * Scale1 + Shift1349 # out = self.adain1(out, style)350 outup = self.convup(out)351 352 return outup353 354##############################################################################355##Face Dictionary356##############################################################################357class VGGFeat(torch.nn.Module):358 """359 Input: (B, C, H, W), RGB, [-1, 1]360 """361 def __init__(self, weight_path='./weights/vgg19.pth'):362 super().__init__()363 self.model = models.vgg19(pretrained=False)364 self.build_vgg_layers()365 366 self.model.load_state_dict(torch.load(weight_path))367 368 self.register_parameter("RGB_mean", nn.Parameter(torch.Tensor([0.485, 0.456, 0.406]).view(1, 3, 1, 1)))369 self.register_parameter("RGB_std", nn.Parameter(torch.Tensor([0.229, 0.224, 0.225]).view(1, 3, 1, 1)))370 371 # self.model.eval()372 for param in self.model.parameters():373 param.requires_grad = False374 375 def build_vgg_layers(self):376 vgg_pretrained_features = self.model.features377 self.features = []378 # feature_layers = [0, 3, 8, 17, 26, 35]379 feature_layers = [0, 8, 17, 26, 35]380 for i in range(len(feature_layers)-1): 381 module_layers = torch.nn.Sequential() 382 for j in range(feature_layers[i], feature_layers[i+1]):383 module_layers.add_module(str(j), vgg_pretrained_features[j])384 self.features.append(module_layers)385 self.features = torch.nn.ModuleList(self.features)386 387 def preprocess(self, x):388 x = (x + 1) / 2389 x = (x - self.RGB_mean) / self.RGB_std390 if x.shape[3] < 224:391 x = torch.nn.functional.interpolate(x, size=(224, 224), mode='bilinear', align_corners=False)392 return x393 394 def forward(self, x):395 x = self.preprocess(x)396 features = []397 for m in self.features:398 # print(m)399 x = m(x)400 features.append(x)401 return features 402 403def compute_sum(x, axis=None, keepdim=False):404 if not axis:405 axis = range(len(x.shape))406 for i in sorted(axis, reverse=True):407 x = torch.sum(x, dim=i, keepdim=keepdim)408 return x409def ToRGB(in_channel):410 return nn.Sequential(411 SpectralNorm(nn.Conv2d(in_channel,in_channel,3, 1, 1)),412 nn.LeakyReLU(0.2),413 SpectralNorm(nn.Conv2d(in_channel,3,3, 1, 1))414 )415 416def AttentionBlock(in_channel):417 return nn.Sequential(418 SpectralNorm(nn.Conv2d(in_channel, in_channel, 3, 1, 1)),419 nn.LeakyReLU(0.2),420 SpectralNorm(nn.Conv2d(in_channel, in_channel, 3, 1, 1))421 )422 423class UNetDictFace(nn.Module):424 def __init__(self, ngf=64, dictionary_path='./DictionaryCenter512'):425 super().__init__()426 427 self.part_sizes = np.array([80,80,50,110]) # size for 512428 self.feature_sizes = np.array([256,128,64,32])429 self.channel_sizes = np.array([128,256,512,512])430 Parts = ['left_eye','right_eye','nose','mouth']431 self.Dict_256 = {}432 self.Dict_128 = {}433 self.Dict_64 = {}434 self.Dict_32 = {}435 for j,i in enumerate(Parts):436 f_256 = torch.from_numpy(np.load(os.path.join(dictionary_path, '{}_256_center.npy'.format(i)), allow_pickle=True))437 438 f_256_reshape = f_256.reshape(f_256.size(0),self.channel_sizes[0],self.part_sizes[j]//2,self.part_sizes[j]//2)439 max_256 = torch.max(torch.sqrt(compute_sum(torch.pow(f_256_reshape, 2), axis=[1, 2, 3], keepdim=True)),torch.FloatTensor([1e-4]))440 self.Dict_256[i] = f_256_reshape #/ max_256441 442 f_128 = torch.from_numpy(np.load(os.path.join(dictionary_path, '{}_128_center.npy'.format(i)), allow_pickle=True))443 444 f_128_reshape = f_128.reshape(f_128.size(0),self.channel_sizes[1],self.part_sizes[j]//4,self.part_sizes[j]//4)445 max_128 = torch.max(torch.sqrt(compute_sum(torch.pow(f_128_reshape, 2), axis=[1, 2, 3], keepdim=True)),torch.FloatTensor([1e-4]))446 self.Dict_128[i] = f_128_reshape #/ max_128447 448 f_64 = torch.from_numpy(np.load(os.path.join(dictionary_path, '{}_64_center.npy'.format(i)), allow_pickle=True))449 450 f_64_reshape = f_64.reshape(f_64.size(0),self.channel_sizes[2],self.part_sizes[j]//8,self.part_sizes[j]//8)451 max_64 = torch.max(torch.sqrt(compute_sum(torch.pow(f_64_reshape, 2), axis=[1, 2, 3], keepdim=True)),torch.FloatTensor([1e-4]))452 self.Dict_64[i] = f_64_reshape #/ max_64453 454 f_32 = torch.from_numpy(np.load(os.path.join(dictionary_path, '{}_32_center.npy'.format(i)), allow_pickle=True))455 456 f_32_reshape = f_32.reshape(f_32.size(0),self.channel_sizes[3],self.part_sizes[j]//16,self.part_sizes[j]//16)457 max_32 = torch.max(torch.sqrt(compute_sum(torch.pow(f_32_reshape, 2), axis=[1, 2, 3], keepdim=True)),torch.FloatTensor([1e-4]))458 self.Dict_32[i] = f_32_reshape #/ max_32459 460 self.le_256 = AttentionBlock(128)461 self.le_128 = AttentionBlock(256)462 self.le_64 = AttentionBlock(512)463 self.le_32 = AttentionBlock(512)464 465 self.re_256 = AttentionBlock(128)466 self.re_128 = AttentionBlock(256)467 self.re_64 = AttentionBlock(512)468 self.re_32 = AttentionBlock(512)469 470 self.no_256 = AttentionBlock(128)471 self.no_128 = AttentionBlock(256)472 self.no_64 = AttentionBlock(512)473 self.no_32 = AttentionBlock(512)474 475 self.mo_256 = AttentionBlock(128)476 self.mo_128 = AttentionBlock(256)477 self.mo_64 = AttentionBlock(512)478 self.mo_32 = AttentionBlock(512)479 480 #norm481 self.VggExtract = VGGFeat()482 483 ######################484 self.MSDilate = MSDilateBlock(ngf*8, dilation = [4,3,2,1]) #485 486 self.up0 = StyledUpBlock(ngf*8,ngf*8)487 self.up1 = StyledUpBlock(ngf*8, ngf*4) #488 self.up2 = StyledUpBlock(ngf*4, ngf*2) #489 self.up3 = StyledUpBlock(ngf*2, ngf) #490 self.up4 = nn.Sequential( # 128491 # nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True),492 SpectralNorm(nn.Conv2d(ngf, ngf, 3, 1, 1)),493 # nn.BatchNorm2d(32),494 nn.LeakyReLU(0.2),495 UpResBlock(ngf),496 UpResBlock(ngf),497 # SpectralNorm(nn.Conv2d(ngf, 3, kernel_size=3, stride=1, padding=1)),498 nn.Conv2d(ngf, 3, kernel_size=3, stride=1, padding=1),499 nn.Tanh()500 )501 self.to_rgb0 = ToRGB(ngf*8)502 self.to_rgb1 = ToRGB(ngf*4)503 self.to_rgb2 = ToRGB(ngf*2)504 self.to_rgb3 = ToRGB(ngf*1)505 506 # for param in self.BlurInputConv.parameters():507 # param.requires_grad = False508 509 def forward(self,input, part_locations):510 511 VggFeatures = self.VggExtract(input)512 # for b in range(input.size(0)):513 b = 0514 UpdateVggFeatures = []515 for i, f_size in enumerate(self.feature_sizes):516 cur_feature = VggFeatures[i]517 update_feature = cur_feature.clone() #* 0518 cur_part_sizes = self.part_sizes // (512/f_size)519 520 dicts_feature = getattr(self, 'Dict_'+str(f_size))521 LE_Dict_feature = dicts_feature['left_eye'].to(input)522 RE_Dict_feature = dicts_feature['right_eye'].to(input)523 NO_Dict_feature = dicts_feature['nose'].to(input)524 MO_Dict_feature = dicts_feature['mouth'].to(input)525 526 le_location = (part_locations[0][b] // (512/f_size)).int()527 re_location = (part_locations[1][b] // (512/f_size)).int()528 no_location = (part_locations[2][b] // (512/f_size)).int()529 mo_location = (part_locations[3][b] // (512/f_size)).int()530 531 LE_feature = cur_feature[:,:,le_location[1]:le_location[3],le_location[0]:le_location[2]].clone()532 RE_feature = cur_feature[:,:,re_location[1]:re_location[3],re_location[0]:re_location[2]].clone()533 NO_feature = cur_feature[:,:,no_location[1]:no_location[3],no_location[0]:no_location[2]].clone()534 MO_feature = cur_feature[:,:,mo_location[1]:mo_location[3],mo_location[0]:mo_location[2]].clone()535 536 #resize537 LE_feature_resize = F.interpolate(LE_feature,(LE_Dict_feature.size(2),LE_Dict_feature.size(3)),mode='bilinear',align_corners=False)538 RE_feature_resize = F.interpolate(RE_feature,(RE_Dict_feature.size(2),RE_Dict_feature.size(3)),mode='bilinear',align_corners=False)539 NO_feature_resize = F.interpolate(NO_feature,(NO_Dict_feature.size(2),NO_Dict_feature.size(3)),mode='bilinear',align_corners=False)540 MO_feature_resize = F.interpolate(MO_feature,(MO_Dict_feature.size(2),MO_Dict_feature.size(3)),mode='bilinear',align_corners=False)541 542 LE_Dict_feature_norm = adaptive_instance_normalization_4D(LE_Dict_feature, LE_feature_resize)543 RE_Dict_feature_norm = adaptive_instance_normalization_4D(RE_Dict_feature, RE_feature_resize)544 NO_Dict_feature_norm = adaptive_instance_normalization_4D(NO_Dict_feature, NO_feature_resize)545 MO_Dict_feature_norm = adaptive_instance_normalization_4D(MO_Dict_feature, MO_feature_resize)546 547 LE_score = F.conv2d(LE_feature_resize, LE_Dict_feature_norm)548 549 LE_score = F.softmax(LE_score.view(-1),dim=0)550 LE_index = torch.argmax(LE_score)551 LE_Swap_feature = F.interpolate(LE_Dict_feature_norm[LE_index:LE_index+1], (LE_feature.size(2), LE_feature.size(3)))552 553 LE_Attention = getattr(self, 'le_'+str(f_size))(LE_Swap_feature-LE_feature)554 LE_Att_feature = LE_Attention * LE_Swap_feature555 556 557 RE_score = F.conv2d(RE_feature_resize, RE_Dict_feature_norm)558 RE_score = F.softmax(RE_score.view(-1),dim=0)559 RE_index = torch.argmax(RE_score)560 RE_Swap_feature = F.interpolate(RE_Dict_feature_norm[RE_index:RE_index+1], (RE_feature.size(2), RE_feature.size(3)))561 562 RE_Attention = getattr(self, 're_'+str(f_size))(RE_Swap_feature-RE_feature)563 RE_Att_feature = RE_Attention * RE_Swap_feature564 565 NO_score = F.conv2d(NO_feature_resize, NO_Dict_feature_norm)566 NO_score = F.softmax(NO_score.view(-1),dim=0)567 NO_index = torch.argmax(NO_score)568 NO_Swap_feature = F.interpolate(NO_Dict_feature_norm[NO_index:NO_index+1], (NO_feature.size(2), NO_feature.size(3)))569 570 NO_Attention = getattr(self, 'no_'+str(f_size))(NO_Swap_feature-NO_feature)571 NO_Att_feature = NO_Attention * NO_Swap_feature572 573 574 MO_score = F.conv2d(MO_feature_resize, MO_Dict_feature_norm)575 MO_score = F.softmax(MO_score.view(-1),dim=0)576 MO_index = torch.argmax(MO_score)577 MO_Swap_feature = F.interpolate(MO_Dict_feature_norm[MO_index:MO_index+1], (MO_feature.size(2), MO_feature.size(3)))578 579 MO_Attention = getattr(self, 'mo_'+str(f_size))(MO_Swap_feature-MO_feature)580 MO_Att_feature = MO_Attention * MO_Swap_feature581 582 update_feature[:,:,le_location[1]:le_location[3],le_location[0]:le_location[2]] = LE_Att_feature + LE_feature583 update_feature[:,:,re_location[1]:re_location[3],re_location[0]:re_location[2]] = RE_Att_feature + RE_feature584 update_feature[:,:,no_location[1]:no_location[3],no_location[0]:no_location[2]] = NO_Att_feature + NO_feature585 update_feature[:,:,mo_location[1]:mo_location[3],mo_location[0]:mo_location[2]] = MO_Att_feature + MO_feature586 587 UpdateVggFeatures.append(update_feature) 588 589 fea_vgg = self.MSDilate(VggFeatures[3])590 #new version591 fea_up0 = self.up0(fea_vgg, UpdateVggFeatures[3])592 # out1 = F.interpolate(fea_up0,(512,512))593 # out1 = self.to_rgb0(out1)594 595 fea_up1 = self.up1( fea_up0, UpdateVggFeatures[2]) #596 # out2 = F.interpolate(fea_up1,(512,512))597 # out2 = self.to_rgb1(out2)598 599 fea_up2 = self.up2(fea_up1, UpdateVggFeatures[1]) #600 # out3 = F.interpolate(fea_up2,(512,512))601 # out3 = self.to_rgb2(out3)602 603 fea_up3 = self.up3(fea_up2, UpdateVggFeatures[0]) #604 # out4 = F.interpolate(fea_up3,(512,512))605 # out4 = self.to_rgb3(out4)606 607 output = self.up4(fea_up3) #608 609 610 return output #+ out4 + out3 + out2 + out1611 #0 128 * 256 * 256612 #1 256 * 128 * 128613 #2 512 * 64 * 64614 #3 512 * 32 * 32615 616 617class UpResBlock(nn.Module):618 def __init__(self, dim, conv_layer = nn.Conv2d, norm_layer = nn.BatchNorm2d):619 super(UpResBlock, self).__init__()620 self.Model = nn.Sequential(621 # SpectralNorm(conv_layer(dim, dim, 3, 1, 1)),622 conv_layer(dim, dim, 3, 1, 1),623 # norm_layer(dim),624 nn.LeakyReLU(0.2,True),625 # SpectralNorm(conv_layer(dim, dim, 3, 1, 1)),626 conv_layer(dim, dim, 3, 1, 1),627 )628 def forward(self, x):629 out = x + self.Model(x)630 return out631 632class VggClassNet(nn.Module):633 def __init__(self, select_layer = ['0','5','10','19']):634 super(VggClassNet, self).__init__()635 self.select = select_layer636 self.vgg = models.vgg19(pretrained=True).features637 for param in self.parameters():638 param.requires_grad = False639 640 def forward(self, x):641 features = []642 for name, layer in self.vgg._modules.items():643 x = layer(x)644 if name in self.select:645 features.append(x)646 return features647 648 649if __name__ == '__main__':650 print('this is network')651 652 653 