尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

PyTorch复现DeepFillv2:门控卷积与自由形式图像修复实战

PyTorch复现DeepFillv2:门控卷积与自由形式图像修复实战 简介DeepFillv2门控卷积自由形式图像修复的PyTorch重新实现资源包主要面向计算机视觉研究者、深度学习开发者以及需要复现论文效果或进行图像修复、风格迁移实验的读者。压缩包内含91个文件以Python脚本、Web前端JS/CSS/HTML、YAML/JSON配置、Markdown/TXT说明文档为主并带有notebook示例、图像素材等整体大小约3.42MB目录结构便于检索。目前已有460人学习下载。包内提供完整训练与测试流程、模型及损失函数定义、CelebA/Places等数据集配置示例还包含Web演示前端与预训练说明可帮助读者快速搭建环境、开展训练与推理并将门控卷积思路迁移到自己的项目或论文复现中显著降低上手门槛。1. 自由形式图像修复与门控卷积不只是修补矩形空洞常规图像修复假设掩码是规则的矩形区域而实际场景中的划痕、遮挡物、文字覆盖大多是任意形状的自由形式掩码。DeepFillv2论文提出用门控卷积替代普通卷积让网络在特征层面动态决定哪些像素参与修复解决了稀疏卷积和局部卷积的掩码泄漏问题。本文按论文路线用PyTorch重新实现门控卷积、粗到细生成器和SN-PatchGAN判别器并给出可直接落地的训练配置与掩码生成代码。适合想复现论文、把修复模型接到自有数据集的工程师也适合正在做图像编辑预处理、需要自由形式区域去除能力的算法团队。2. 门控卷积原理与DeepFillv2网络结构拆解2.1 为什么局部卷积处理不了自由形式掩码自由形式图像修复的直接思路是用掩码信息屏蔽无效像素。局部卷积是DeepFillv1的核心它对掩码区域做归一化并在每一层之后将掩码二值化为0/1只有有效区域参与卷积。问题是掩码一旦被卷积核涂抹二值化边界会产生不自然的阶梯效应而且更新规则是写死的网络无法针对不同语义内容调整对掩码的信任程度。门控卷积的核心改动是用一个可学习的sigmoid门控替换固定掩码更新规则。对于每个卷积层输入特征经过两个并行的卷积一个产生特征响应另一个产生门控系数最终输出是两者的逐通道元素级乘法。门控值在0到1之间连续分布网络通过训练自动学会将哪些区域视为有效、哪些区域作为边界过渡不再依赖手工掩码更新。从本质上看门控卷积在每一层引入了一个软注意力机制普通卷积对所有像素一视同仁局部卷积只区分有效和无效的二值状态而门控卷积能做到空间与通道维度上的自适应选择。这一特性正好贴合自由形式掩码的任意形状。掩码边界附近的特征需要被半保留浅层修复结果中物体边缘处的纹理连续性就是靠这种连续门控值维持的。论文中给出过一个直观现象经过门控卷积后浅层门控值会在掩码边缘形成渐变过渡带而深层门控值则与物体语义边界高度相关——这说明网络确实学会了按内容而非按掩码来决策。2.2 DeepFillv2生成器两阶段级联与门控卷积堆叠DeepFillv2的生成器沿用两阶段级联结构。粗网络接收被掩码遮蔽的RGB图像掩码区域的像素填充为255输入通道为3粗网络输出低层结构完整的粗略结果。细网络的输入是把掩码后图像、粗网络输出、原始掩码按通道拼接形成通道数为7的张量经过另一组编码器-解码器输出最终修复图。两阶段共享同一种门控卷积基本单元但粗网络只对整体结构负责细网络补充高频纹理细节。为什么必须分成两个阶段如果不分阶段单一解码器要从空洞里同时预测结构和纹理梯度信号在深层编码器中容易被噪声主导训练很不稳定。级联让粗网络在语义层先收敛细网络再去学习纹理修复这也是论文能在256×256分辨率下稳定训练的关键。细网络接收的输入通道数较多第一层门控卷积的参数量会明显上升实现时需要注意显存占用。2.3 判别器与训练目标SN-PatchGAN配合WGAN-GP判别器使用谱归一化的PatchGAN论文中称为SN-PatchGAN。谱归一化约束每层权重矩阵的最大奇异值让判别器的Lipschitz常数可控配合WGAN-GP的梯度惩罚项可以不用BatchNorm也能稳定训练。PatchGAN在输出特征图的每个位置上做真伪判别每个感受野是一个局部块这让判别器更关注纹理细节是否连贯而不是整图是否协调。自由形式掩码的面积和形状是变化的全局判别很难对齐不同尺度的信息PatchGAN的局部判别方式明显更适合该场景。提示复现时不要在生成器里加BatchNorm。门控卷积配合BatchNorm在小批量训练时统计量漂移明显实测用InstanceNorm或干脆不加归一化更稳定。3. 用PyTorch重新实现门控卷积与DeepFillv2生成器3.1 最小可跑的GatedConv2d模块门控卷积的PyTorch实现只需一个双分支卷积加一次逐元素相乘。下面的代码给出了不依赖任何第三方库的最小模块并支持通过use_sn开关控制是否启用谱归一化。import torch import torch.nn as nn import torch.nn.functional as F class GatedConv2d(nn.Module): def __init__(self, in_channels, out_channels, kernel_size, stride1, padding0, dilation1, use_snFalse): super().__init__() self.feature_conv nn.Conv2d(in_channels, out_channels, kernel_size, stride, padding, dilation) self.gate_conv nn.Conv2d(in_channels, out_channels, kernel_size, stride, padding, dilation) if use_sn: self.feature_conv nn.utils.spectral_norm(self.feature_conv) self.gate_conv nn.utils.spectral_norm(self.gate_conv) def forward(self, x): feature self.feature_conv(x) gate torch.sigmoid(self.gate_conv(x)) return feature * gate两条卷积分支的输入输出通道数完全一致。feature分支不做激活gate分支过sigmoid将值压缩到0到1之间二者相乘后结果的取值范围完全由feature分支决定。注意不要把ReLU放在feature分支后面再乘门控这样会让负值特征在门控为1时也无法表达修复结果会偏灰、缺乏暗部层次。同一篇论文里的局部卷积实现要维护一个不断更新的掩码张量而门控卷积不需要这是两者实现复杂度差异最大的地方。训练时如果显存紧张可以不先打开use_sn等判别器输出出现了明显的振荡再补上。3.2 搭建编码器-解码器骨架与掩码下采样生成器每阶段的核心是编码器-解码器。编码器用步长为2的门控卷积做下采样解码器使用双线性插值上采样再接门控卷积。下采样时掩码也需要同步缩放否则编码器深处分辨率缩小后掩码和特征图无法对齐。def downsample(self, x, mask, in_ch, out_ch): x self.gated_conv_down(x, in_ch, out_ch, kernel_size3, stride2, padding1) mask F.interpolate(mask, scale_factor0.5, modenearest) return x, mask def upsample(self, x, in_ch, out_ch): x F.interpolate(x, scale_factor2, modebilinear, align_cornersFalse) x self.gated_conv_up(x, in_ch, out_ch, kernel_size3, stride1, padding1) return x掩码下采样必须用nearest模式不能使用bilinear。因为掩码只有0和1两个值双线性插值会产生0.5这样的中间值门控卷积看到半透明的掩码会去修复原本不需要修复的像素尤其在掩码边缘会产生一圈虚影。3.3 通道数配置与细阶段输入按论文惯例编码器每下采样一次通道数翻倍解码器每上采样一次通道数减半。一个适合256×256输入的配置可以这样设定阶段输入通道输出通道分辨率变化模块堆叠编码器第1层3或732256→128GatedConv2d, stride2编码器第2层3264128→64GatedConv2d, stride2编码器第3层6412864→32GatedConv2d, stride2编码器第4层12825632→16GatedConv2d, stride2解码器第1层25612816→32Interpolate GatedConv2d解码器第2层1286432→64Interpolate GatedConv2d解码器第3层643264→128Interpolate GatedConv2d解码器第4层323128→256Interpolate GatedConv2d粗网络的输入通道是3即掩码填充后的原图细网络输入通道是7拼接方式为torch.cat([masked_image, coarse_result, mask], dim1)。生成器的最后输出建议接一个nn.Tanh激活将输出限制到-1到1区间与输入图像的归一化方式保持一致。4. 自由形式掩码生成与训练数据管道构建4.1 随机绘制任意形状掩码滑光标与椭圆笔刷自由形式掩码生成的核心是模拟用户在涂抹、刮擦、物体移除时产生的任意形状区域。最简单且最接近论文做法的是滑光标方式随机生成若干条折线段路径沿路径用大小可变的椭圆笔刷画出掩码区域。每个掩码的张数、折线长度、笔刷半径都从一定的范围内随机采样这样可以覆盖从细长划痕到大面积遮挡的各种形状。import numpy as np from scipy.ndimage import rotate def random_brush_mask(height, width, max_vertex12, max_brush24): mask np.zeros((height, width), dtypenp.uint8) num_strokes np.random.randint(1, 4) for _ in range(num_strokes): num_vertex np.random.randint(4, max_vertex 1) start_x np.random.randint(0, width - 1) start_y np.random.randint(0, height - 1) for _ in range(num_vertex): angle np.random.uniform(0, 2 * np.pi) dist np.random.uniform(0, 0.3 * max(height, width)) end_x np.clip(start_x dist * np.cos(angle), 0, width - 1) end_y np.clip(start_y dist * np.sin(angle), 0, height - 1) brush_radius np.random.uniform(2, max_brush) draw_line(mask, (start_y, start_x), (int(end_y), int(end_x)), int(brush_radius)) start_x, start_y int(end_x), int(end_y) return mask def draw_line(mask, start, end, radius): y1, x1 start y2, x2 end dist max(abs(x2 - x1), abs(y2 - y1)) for i in range(dist 1): t i / max(dist, 1) x int(x1 t * (x2 - x1)) y int(y1 t * (y2 - y1)) cv2.circle(mask, (x, y), radius, 1, -1)掩码面积比例需要严格控制。论文中训练时随机采样10%到40%的掩码面积比例这个比例既保证修复有难度又保证背景信息足够支撑生成器做推理。面积比例过小会让模型退化成几乎不做任何修复也能通过判别器面积比例过大会让生成器只能猜测颜色训练不出纹理。4.2 图像归一化与掩码注入方式图像修复训练不需要成对的ground truth之外的额外标注数据管道就是把原始图像作为监督信号。输入图像在送入生成器之前归一化到-1到1之间掩码则保持0和1的整数值。被掩码遮蔽的图像构造方式直接决定模型看到的空缺状态def apply_mask(image, mask): # image: [0, 1] float tensor, shape (C, H, W) # mask: 二值张量, shape (H, W), 1 表示需修复 masked image.clone() mask_bchw mask.unsqueeze(0).float() # (1, H, W) masked masked * (1 - mask_bchw) mask_bchw # 掩码区域填充为1.0白色 return masked掩码区域填充为白色只是论文中采用的其中一种注入方式。实践中还可以填充为随机噪声、数据集平均像素值甚至像素打乱结果。填充颜色的选择会轻微影响模型训练初期的收敛速度但最终修复效果差别不大因为门控卷积会学会忽略掩码区域的像素值重点提取掩码外区域的特征。4.3 DataLoader吞吐优化与补丁采样自由形式修复训练在256×256分辨率下单卡跑批大小8通常只能勉强支撑完整生成器加判别器。我一般会在数据管道里先随机裁剪512×512的大图再缩放到256×256送入网络这样既增加了样本多样性又避免直接加载超大原图浪费内存。对显存仍不足的情况可以先将图像降至128×128做粗网络预热训练待损失稳定后再切换回256×256微调。DataLoader中的num_workers在图像修复任务里的影响比较明显推荐设置为CPU核心数的一半左右。掩码生成运算量不小如果每次都在__getitem__里实时绘制会拖慢训练吞吐常见的做法是预生成一批掩码保存为npy格式训练时按索引直接读取减少CPU计算压力。5. 损失函数、训练超参设置与稳定性排查5.1 复合损失L1、感知损失与WGAN-GP的组合方式DeepFillv2训练损失是生成器损失与判别器损失的加权组合。生成器部分包括L1像素损失、VGG感知损失和对抗损失。L1损失保证生成结果与真实图像的逐像素距离最小感知损失约束特征空间上的语义一致性对抗损失在PatchGAN输出的每个位置上做WGAN-GP形式的最小二乘或最小绝对值优化。l1_loss F.l1_loss(coarse_out, gt_patch) * 1.2 l1_fine F.l1_loss(fine_out, gt_patch) * 1.2 perceptual vgg_loss(fine_out, gt_patch) * 0.05 wgan_gp d_loss(fine_out, gt_patch, mask) * 1.0 g_loss l1_loss l1_fine perceptual wgan_gp权重设置中L1损失的权重最高约1.2感知损失权重在0.05左右即可这是作者开源配置的大致区间。感知损失权重过大容易让修复区域纹理过于平滑权重过小则会在语义结构上出现断裂。对抗损失权重设为1.0即可不需要额外的平衡系数。5.2 学习率与迭代策略生成器和判别器使用相同的学习率训练Adam优化器的β1设为0.5、β2设为0.999是论文中的常见设置。生成器的学习率取0.0001判别器可以比生成器高一倍取0.0002两者交替更新。学习率过高时门控分支的sigmoid输出会迅速饱和到0或1导致门控失效学习率过低则掩码边界收敛极其缓慢往往需要数万次迭代才能看到门控值产生实际变化。WGAN-GP的梯度惩罚系数lambda设为10。每训练一个batch生成器之前先训练3个batch的判别器这个比例能有效避免判别器被骗过。实际训练中如果发现判别器损失降到0需要立即降低学习率并检查谱归一化是否在判别器每层都被启用。5.3 训练稳定性排查的三个常见现象门控卷积在训练初期常见的一个问题是生成器输出整体偏灰。原因通常是feature分支的初始权重让卷积输出集中在零附近门控值虽然接近0.5但乘积结果约等于原特征的一半。将feature分支的卷积权重按nn.init.kaiming_normal_初始化并把偏置置零可以在前几千步内缓解。第二个问题是修复区域出现棋盘格伪影。这种伪影大多来自双线性上采样后的3×3卷积配合转置卷积叠加。解决方案是将所有上采样都改为双线性插值加普通卷积不使用转置卷积。第三个问题是掩码边界出现一条明显的接缝线这通常是细网络输入拼接了掩码后未经归一化的掩码值0/1与图像特征量级差异过大造成的。在拼接前把掩码减去0.5即可将差异缩小。提示训练前先固定随机种子做两次相同配置的短训练对比损失曲线是否一致。门控卷积的初始化对结果影响较大保证实验可复现很重要。6. 模型评估、推理优化与代码打包分发6.1 用PSNR、SSIM与FID评估自由形式修复质量自由形式修复的评估不能只用一个指标。PSNR反映逐像素误差但自由形式掩码的面积比例不同直接对比不同掩码下的PSNR没有参考意义。我一般做法是固定一组测试掩码生成脚本保证所有对比模型使用完全相同的掩码与输入图像这样PSNR和SSIM才具有可比性。FID更关注生成分布与真实分布的差距对自由形式修复尤其关键建议在256×256分辨率下用3000张以上图像计算样本太少时FID方差很大。推理阶段要注意的一个细节是掩码在训练时经过了下采样与原始输入对齐推理时也要对掩码做同步的nearest缩放否则掩码与图像分辨率不一致会导致输出出现偏移。输入图像在送入模型前要确认归一化到-1到1掩码区域的值必须和训练时保持一致通常填充为1.0这样模型才能正确做出缺失区域判断。6.2 用torch.jit.script打包模型并与代码一起分发模型训练完成后常见的做法是导出为TorchScript格式方便在离线环境中直接加载。门控卷积模块包含两条独立的卷积分支TorchScript可以正常trace但要注意sigmoid在trace时会被内联导致脚本丢失部分调试信息。为了保留灵活性建议编写一个forward函数明确写出feature乘gate的操作并用torch.jit.script而非trace来导出。class InpaintModel(nn.Module): def forward(self, masked_img, mask): mask_scaled F.interpolate(mask, scale_factor0.5, modenearest) coarse self.coarse_net(masked_img, mask_scaled) fine_in torch.cat([masked_img, coarse, mask], dim1) return self.fine_net(fine_in, mask_scaled) scripted_model torch.jit.script(model) torch.jit.save(scripted_model, deepfillv2_gated.pt)代码分发时常见的做法是将训练脚本、掩码生成器、模型权重和README打包成一个zip压缩包。PyTorch权重文件本身较大建议在打包前清理临时日志文件。如果你接收到的zip压缩包在解压时报出error read zip archive或提示文件损坏先去检查压缩包是否下载完整在命令行用unzip -t做完整性测试能快速确认是网络传输问题还是文件本身就缺失了分卷压缩的某个part。6.3 压缩包验证与依赖锁定分发模型前用一条命令验证整个zip内的文件依赖是否齐全unzip -t deepfillv2_reimplementation.zipunzip -t只校验压缩包内每个文件的CRC是否正确不会检测Python import路径。更稳妥的做法是解压后在项目根目录执行python -c from gated_conv import GatedConv2d; print(ok)做导入冒烟测试同时检查requirements.txt中PyTorch版本是否与当前环境匹配。PyTorch 1.x与2.x的TorchScript兼容性存在差异如果模型在一个版本下script并且在另一个版本下加载可能会遇到无法加载的提示尽量保证训练环境和推理环境使用同一个小版本。将代码、权重、测试掩码和复现说明打包成zip并不意味着分发工作结束在解压后的全新Python环境中完整跑一遍推理脚本才能确认依赖没有遗漏。本文还有配套的精品资源点击获取
返回列表