Team Ai
Apppublic

GoodWin/Deep-Multi-scale

sourceHugging Faceupdated 5y agoView on Hugging Face
0likes
networks.py653 linesDownload Raw Back to models
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