guided-pix2pix 代码略解

2024-03-06 20:38
文章标签 代码 guided 略解 pix2pix

本文主要是介绍guided-pix2pix 代码略解,希望对大家解决编程问题提供一定的参考价值,需要的开发者们随着小编来一起学习吧!

《Guided Image-to-Image Translation with Bi-Directional Feature Transformation》

train.py

            model.set_input(data)model.optimize_parameters()

开始训练

models/guided_pix2pix_model.py

    def set_input(self, input):self.real_A = input['A'].to(self.device)self.real_B = input['B'].to(self.device)self.guide = input['guide'].to(self.device)

这里的guide是指的是引导的image/pose,real_A指的是作为输入的image,real_B指的是GT

    def forward(self):self.fake_B = self.netG(self.real_A, self.guide)
# load/define networks
self.netG = networks.define_G(input_nc=opt.input_nc, guide_nc=opt.guide_nc, output_nc=opt.output_nc, ngf=opt.ngf, netG=opt.netG, n_layers=opt.n_layers, norm=opt.norm, init_type=opt.init_type, init_gain=opt.init_gain, gpu_ids=self.gpu_ids)

models/networks.py

def define_G(input_nc, guide_nc, output_nc, ngf, netG, n_layers=8, n_downsampling=3, n_blocks=9, norm='batch', init_type='normal', init_gain=0.02, gpu_ids=[]):net = Nonenorm_layer = get_norm_layer(norm_type=norm)if netG == 'bFT_resnet':net = bFT_Resnet(input_nc, guide_nc, output_nc, ngf, norm_layer=norm_layer, n_blocks=n_blocks)elif netG == 'bFT_unet':net = bFT_Unet(input_nc, guide_nc, output_nc, n_layers, ngf, norm_layer=norm_layer)else:raise NotImplementedError('Generator model name [%s] is not recognized' % netG)net = init_net(net, init_type, init_gain, gpu_ids)return net

看一下bFT_resent

class bFT_Resnet(nn.Module):def __init__(self, input_nc, guide_nc, output_nc, ngf=64, n_blocks=9, norm_layer=nn.BatchNorm2d,padding_type='reflect', bottleneck_depth=100):super(bFT_Resnet, self).__init__()self.activation = nn.ReLU(True)n_downsampling=3## inputpadding_in = [nn.ReflectionPad2d(3), nn.Conv2d(input_nc, ngf, kernel_size=7, padding=0)]self.padding_in = nn.Sequential(*padding_in)self.conv1 = nn.Conv2d(ngf, ngf * 2, kernel_size=3, stride=2, padding=1)self.conv2 = nn.Conv2d(ngf * 2, ngf * 4, kernel_size=3, stride=2, padding=1)self.conv3 = nn.Conv2d(ngf * 4, ngf * 8, kernel_size=3, stride=2, padding=1)## guidepadding_g = [nn.ReflectionPad2d(3), nn.Conv2d(guide_nc, ngf, kernel_size=7, padding=0)]self.padding_g = nn.Sequential(*padding_g)self.conv1_g = nn.Conv2d(ngf, ngf * 2, kernel_size=3, stride=2, padding=1)self.conv2_g = nn.Conv2d(ngf * 2, ngf * 4, kernel_size=3, stride=2, padding=1)self.conv3_g = nn.Conv2d(ngf * 4, ngf * 8, kernel_size=3, stride=2, padding=1)# bottleneck1self.bottleneck_alpha_1 = self.bottleneck_layer(ngf, bottleneck_depth)self.G_bottleneck_alpha_1 = self.bottleneck_layer(ngf, bottleneck_depth)self.bottleneck_beta_1 = self.bottleneck_layer(ngf, bottleneck_depth)self.G_bottleneck_beta_1 = self.bottleneck_layer(ngf, bottleneck_depth)# bottleneck2self.bottleneck_alpha_2 = self.bottleneck_layer(ngf*2, bottleneck_depth)self.G_bottleneck_alpha_2 = self.bottleneck_layer(ngf*2, bottleneck_depth)self.bottleneck_beta_2 = self.bottleneck_layer(ngf*2, bottleneck_depth)self.G_bottleneck_beta_2 = self.bottleneck_layer(ngf*2, bottleneck_depth)# bottleneck3self.bottleneck_alpha_3 = self.bottleneck_layer(ngf*4, bottleneck_depth)self.G_bottleneck_alpha_3 = self.bottleneck_layer(ngf*4, bottleneck_depth)self.bottleneck_beta_3 = self.bottleneck_layer(ngf*4, bottleneck_depth)self.G_bottleneck_beta_3 = self.bottleneck_layer(ngf*4, bottleneck_depth)# bottleneck4self.bottleneck_alpha_4 = self.bottleneck_layer(ngf*8, bottleneck_depth)self.G_bottleneck_alpha_4 = self.bottleneck_layer(ngf*8, bottleneck_depth)self.bottleneck_beta_4 = self.bottleneck_layer(ngf*8, bottleneck_depth)self.G_bottleneck_beta_4 = self.bottleneck_layer(ngf*8, bottleneck_depth)### 这些bottlenect_layer都是由1x1的卷积,激活层,1x1的卷积组成的,做从nc->nc的映射resnet = []mult = 2**n_downsamplingfor i in range(n_blocks):resnet += [ResnetBlock(ngf * mult, padding_type=padding_type, activation=self.activation, norm_layer=norm_layer)]self.resnet = nn.Sequential(*resnet)decoder = []for i in range(n_downsampling):mult = 2**(n_downsampling - i)decoder += [nn.ConvTranspose2d(ngf * mult, int(ngf * mult / 2), kernel_size=3, stride=2, padding=1, output_padding=1),norm_layer(int(ngf * mult / 2)), self.activation]self.pre_decoder = nn.Sequential(*decoder)self.decoder = nn.Sequential(*[nn.ReflectionPad2d(3), nn.Conv2d(ngf, output_nc, kernel_size=7, padding=0), nn.Tanh()])def bottleneck_layer(self, nc, bottleneck_depth):return nn.Sequential(*[nn.Conv2d(nc, bottleneck_depth, kernel_size=1), self.activation, nn.Conv2d(bottleneck_depth, nc, kernel_size=1)])def get_FiLM_param_(self, X, i, guide=False):x = X.clone()# bottleneckif guide:if (i==1):alpha_layer = self.G_bottleneck_alpha_1beta_layer = self.G_bottleneck_beta_1elif (i==2):alpha_layer = self.G_bottleneck_alpha_2beta_layer = self.G_bottleneck_beta_2elif (i==3):alpha_layer = self.G_bottleneck_alpha_3beta_layer = self.G_bottleneck_beta_3elif (i==4):alpha_layer = self.G_bottleneck_alpha_4beta_layer = self.G_bottleneck_beta_4else:if (i==1):alpha_layer = self.bottleneck_alpha_1beta_layer = self.bottleneck_beta_1elif (i==2):alpha_layer = self.bottleneck_alpha_2beta_layer = self.bottleneck_beta_2elif (i==3):alpha_layer = self.bottleneck_alpha_3beta_layer = self.bottleneck_beta_3elif (i==4):alpha_layer = self.bottleneck_alpha_4beta_layer = self.bottleneck_beta_4alpha = alpha_layer(x)beta = beta_layer(x)return alpha, betadef forward(self, input, guidance):input = self.padding_in(input)  guidance = self.padding_g(guidance)g_alpha1, g_beta1 = self.get_FiLM_param_(guidance, 1, guide=True)i_alpha1, i_beta1 = self.get_FiLM_param_(input, 1)guidance = affine_transformation(guidance, i_alpha1, i_beta1)input = affine_transformation(input, g_alpha1, g_beta1)input = self.activation(input)guidance = self.activation(guidance)input = self.conv1(input)guidance = self.conv1_g(guidance)g_alpha2, g_beta2 = self.get_FiLM_param_(guidance, 2, guide=True)i_alpha2, i_beta2 = self.get_FiLM_param_(input, 2)input = affine_transformation(input, g_alpha2, g_beta2)guidance = affine_transformation(guidance, i_alpha2, i_beta2)input = self.activation(input)guidance = self.activation(guidance)input = self.conv2(input)guidance = self.conv2_g(guidance)g_alpha3, g_beta3 = self.get_FiLM_param_(guidance, 3, guide=True)i_alpha3, i_beta3 = self.get_FiLM_param_(input, 3)input = affine_transformation(input, g_alpha3, g_beta3)guidance = affine_transformation(guidance, i_alpha3, i_beta3)input = self.activation(input)guidance = self.activation(guidance)input = self.conv3(input)guidance = self.conv3_g(guidance)g_alpha4, g_beta4 = self.get_FiLM_param_(guidance, 4, guide=True)# guidance在这一步之后就舍弃了input = affine_transformation(input, g_alpha4, g_beta4)input = self.activation(input)input = self.resnet(input)input = self.pre_decoder(input)output = self.decoder(input)return output

 

这篇关于guided-pix2pix 代码略解的文章就介绍到这儿,希望我们推荐的文章对编程师们有所帮助!



http://www.chinasem.cn/article/781284

相关文章

Python实例题之pygame开发打飞机游戏实例代码

《Python实例题之pygame开发打飞机游戏实例代码》对于python的学习者,能够写出一个飞机大战的程序代码,是不是感觉到非常的开心,:本文主要介绍Python实例题之pygame开发打飞机... 目录题目pygame-aircraft-game使用 Pygame 开发的打飞机游戏脚本代码解释初始化部

Java中Map.Entry()含义及方法使用代码

《Java中Map.Entry()含义及方法使用代码》:本文主要介绍Java中Map.Entry()含义及方法使用的相关资料,Map.Entry是Java中Map的静态内部接口,用于表示键值对,其... 目录前言 Map.Entry作用核心方法常见使用场景1. 遍历 Map 的所有键值对2. 直接修改 Ma

深入解析 Java Future 类及代码示例

《深入解析JavaFuture类及代码示例》JavaFuture是java.util.concurrent包中用于表示异步计算结果的核心接口,下面给大家介绍JavaFuture类及实例代码,感兴... 目录一、Future 类概述二、核心工作机制代码示例执行流程2. 状态机模型3. 核心方法解析行为总结:三

python获取cmd环境变量值的实现代码

《python获取cmd环境变量值的实现代码》:本文主要介绍在Python中获取命令行(cmd)环境变量的值,可以使用标准库中的os模块,需要的朋友可以参考下... 前言全局说明在执行py过程中,总要使用到系统环境变量一、说明1.1 环境:Windows 11 家庭版 24H2 26100.4061

pandas实现数据concat拼接的示例代码

《pandas实现数据concat拼接的示例代码》pandas.concat用于合并DataFrame或Series,本文主要介绍了pandas实现数据concat拼接的示例代码,具有一定的参考价值,... 目录语法示例:使用pandas.concat合并数据默认的concat:参数axis=0,join=

C#代码实现解析WTGPS和BD数据

《C#代码实现解析WTGPS和BD数据》在现代的导航与定位应用中,准确解析GPS和北斗(BD)等卫星定位数据至关重要,本文将使用C#语言实现解析WTGPS和BD数据,需要的可以了解下... 目录一、代码结构概览1. 核心解析方法2. 位置信息解析3. 经纬度转换方法4. 日期和时间戳解析5. 辅助方法二、L

Python使用Code2flow将代码转化为流程图的操作教程

《Python使用Code2flow将代码转化为流程图的操作教程》Code2flow是一款开源工具,能够将代码自动转换为流程图,该工具对于代码审查、调试和理解大型代码库非常有用,在这篇博客中,我们将深... 目录引言1nVflRA、为什么选择 Code2flow?2、安装 Code2flow3、基本功能演示

IIS 7.0 及更高版本中的 FTP 状态代码

《IIS7.0及更高版本中的FTP状态代码》本文介绍IIS7.0中的FTP状态代码,方便大家在使用iis中发现ftp的问题... 简介尝试使用 FTP 访问运行 Internet Information Services (IIS) 7.0 或更高版本的服务器上的内容时,IIS 将返回指示响应状态的数字代

MySQL 添加索引5种方式示例详解(实用sql代码)

《MySQL添加索引5种方式示例详解(实用sql代码)》在MySQL数据库中添加索引可以帮助提高查询性能,尤其是在数据量大的表中,下面给大家分享MySQL添加索引5种方式示例详解(实用sql代码),... 在mysql数据库中添加索引可以帮助提高查询性能,尤其是在数据量大的表中。索引可以在创建表时定义,也可

使用C#删除Excel表格中的重复行数据的代码详解

《使用C#删除Excel表格中的重复行数据的代码详解》重复行是指在Excel表格中完全相同的多行数据,删除这些重复行至关重要,因为它们不仅会干扰数据分析,还可能导致错误的决策和结论,所以本文给大家介绍... 目录简介使用工具C# 删除Excel工作表中的重复行语法工作原理实现代码C# 删除指定Excel单元