ARTICLE DETAIL

资讯详情

深耕网站建设、视觉设计与SEO优化的一线实战洞察。

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

自由形式图像修复的门控卷积:DeepFillv2的PyTorch重实现 简介面向深度学习和计算机视觉研究者与图像修复方向初学者这是一份基于DeepFillv2论文arXiv:1806.03589的自由形式图像修复PyTorch重新实现可帮助解决门控卷积、部分卷积与注意力机制等模型的搭建、训练和二次开发问题。资源压缩包共91个文件、约3.42MB以11个Python源码为核心配合21个JS、14个CSS与18张PNG构建了可交互前端展示同时提供YAML/JSON配置、Markdown/Text说明和Jupyter Notebook示例。目前已有460人学习或下载具有一定参考热度。包内包含模型定义、多种损失函数、训练/测试脚本和app.py启动入口可按README配置CelebA、Places2等数据集参数直接运行整体目录划分清晰从前端演示到后端推理均有覆盖便于对照论文逐模块阅读快速开展自由形式的图像修复实验。1. 自由形式图像修复的门控卷积路线从“补方块”到“补任意形状”修复一个不规则的涂鸦和修复一个矩形空洞根本不是同一件事。传统深度修复模型在矩形 mask 上能够交出不错的答卷但一旦 mask 变成自由形式普通卷积会把缺失区域的像素和已知区域像素一样对待结果边缘出现色差、内部一片模糊。DeepFillv2 用门控卷积把“该不该相信这个位置的特征”变成网络自己学出来的逐像素权重让修复任意形状成为可能。即使你拿到手的只是一份 zip 源码包最核心的门控卷积也只需要几行代码就能看懂。在 PyTorch 里重新实现这条路线不需要魔改框架核心是把普通 Conv2d 拆成特征卷积分支和 sigmoid 门控分支两个分支做逐元素乘法。门控分支可以理解成一个由网络自动推导的软掩码不需要外部输入 mask 更新规则。相比 Partial Convolution它的实现更简单表达也更连续。下面会按照门控卷积原理、PyTorch 最小实现、自由形式 mask 生成、训练参数、推理技巧的顺序展开。适合做过图像编辑或 GAN 项目、想从论文到工程落地的从业者如果你只是调接口验证效果可以直接跳到第 5 章看参数。2. 门控卷积原理与 DeepFillv2 的设计取舍2.1 Partial Convolution 为什么在自由形式下受限在 DeepFillv2 之前Free-Form Inpainting 领域最常被拿来当基线的是 Partial Convolution。它把 mask 当成硬约束卷积只在有效像素上计算每次卷积后 mask 要跟着更新只要卷积核覆盖到至少一个有效像素该位置就被标记为有效。这在矩形空洞场景里很自然但自由形式 mask 往往又细又碎经过几次下采样小空洞很容易被周围有效像素“填满”mask 快速变成全 1后续层完全丢失缺失位置信息。第二个问题是硬掩码的梯度只在有效区域传播缺失区域本身没有特征提取梯度训练早期容易停在局部解。从工程角度说Partial Convolution 的实现需要额外维护一个 mask 分支每次前向都要做一次掩码更新判断推理时还要同步处理这个张量。自由形式修复数据集里的 mask 千奇百怪这种强约束会让模型泛化得很费力。这也是我转向门控卷积的直接原因代码更短表达更自由。两组卷积的差异可以用下面表格概括对比项Partial ConvolutionGated Convolution外部 mask 依赖是逐层更新否自动学习门控门控粒度二值硬掩码0~1 连续软掩码特征分支激活ReLU/ELUELU门控分支激活无Sigmoid自由形式 mask 表现小空洞易丢失可表达半遮挡2.2 门控卷积的数学形式特征卷积分支和门控分支的乘法门控卷积的关键是同时使用两个卷积核一个负责提取内容特征一个负责生成门控。两者的输入都是当前层的特征图但权重各自独立。前向计算可以用下面这段代码直观表示def gated_conv2d(x, conv_f, conv_g, activationtorch.nn.functional.elu): feature activation(conv_f(x)) gate torch.sigmoid(conv_g(x)) return feature * gate这个函数的含义是conv_f 的结果保留输入中的结构信息conv_g 的结果经过 sigmoid 后变成逐元素软开关。两个分支共享同一个输入特征图但不共享权重。门控分支能感知到当前像素周围是否存在空洞、颜色突变、边缘等线索从而决定特征分支的信息是否保留。这样 mask 不再作为额外输入而是在每一层被隐式地重新推导比手动更新 mask 更灵活。提示门控分支的初始偏差最好设成正数我习惯在模块构造函数里把 gate 卷积的 bias 初始化为 1.0。这样训练初期 gate 输出接近 1不会让网络中所有门都关闭导致梯度消失。2.3 DeepFillv2 在门控基础上做了哪些关键简化我理解的 DeepFillv2 核心结构是全卷积的编码-解码网络卷积层全部用门控卷积替换并配合空洞卷积扩大感受野。它有三个关键设计值得在重实现时保留第一是不再单独维护 mask 分支所有门控由网络自动学习这让网络可以在自由形式 mask 上泛化第二是生成器内部使用 ELU 激活特征分支门控分支用 Sigmoid梯度不容易在中途消失第三是判别器采用谱归一化 PatchGAN对抗损失写成 Hinge Loss而不是 WGAN 的 Wasserstein 损失省掉梯度惩罚的同时也让训练曲线更容易观察。这三个设计缺一不可。如果只把卷积换成门控而继续用原来的 GAN 损失判别器很容易比生成器强太多如果去掉谱归一化判别器输出尺度会震荡生成图片容易出现对比度过高的伪纹理。重实现时我一般把门控卷积层、谱归一化和 Hinge Loss 视为一个整体一次实现不在局部价值上抠太多。另外DeepFillv2 的生成器还有两种常见变体单阶段编码-解码和粗糙到精细两阶段。前者实现简单适合快速验证后者先在低分辨率预测结构再把粗略输出和原始输入拼起来细化。实际体验中两阶段在大的 mask 上的细节完整度明显更好但显存开销会增加三分之一左右。如果你只是复现论文的最小流程先从单阶段起步更省事。3. PyTorch 重实现门控卷积层、网络结构与损失函数3.1 最小门控卷积层实现这里直接给出一个可独立复用的门控卷积模块。它接受输入张量输出同大小的门控特征图可以像普通 Conv2d 一样放入任何网络主干中。import torch import torch.nn as nn import torch.nn.functional as F class GatedConv2d(nn.Module): def __init__(self, in_ch, out_ch, kernel_size3, stride1, padding1, dilation1, use_snTrue): super().__init__() self.conv_f nn.Conv2d(in_ch, out_ch, kernel_size, stride, padding, dilation, biasTrue) self.conv_g nn.Conv2d(in_ch, out_ch, kernel_size, stride, padding, dilation, biasTrue) if use_sn: self.conv_f nn.utils.spectral_norm(self.conv_f) self.conv_g nn.utils.spectral_norm(self.conv_g) nn.init.constant_(self.conv_g.bias, 1.0) # 门初始为开 self.act nn.ELU() def forward(self, x): feature self.act(self.conv_f(x)) gate torch.sigmoid(self.conv_g(x)) return feature * gate参数说明in_ch是输入通道数在修复模型的第一层通常是 4由 RGB 图像和单通道 mask 拼接而成。use_snTrue会同时给两个分支加谱归一化作用是限制卷积权重矩阵的最大奇异值让生成器输出数值范围更可控。conv_g的 bias 初始化为 1.0这样训练开始时门控接近全 1网络退化成一个普通卷积不会因为门全部关闭而无法优化。3.2 编码-解码主干下采样、空洞卷积和残差连接自由形式修复要求模型既能看到全局语义又要保留局部纹理。单靠下采样和上采样会丢失边缘高频信息。常见做法是编码器逐步下采样在瓶颈层叠加不同空洞率的门控卷积最后再利用 skip connection 逐级恢复空间尺寸。class DeepFillv2Generator(nn.Module): def __init__(self): super().__init__() self.enc1 GatedConv2d(4, 32, 3, 1, 1) self.enc2 GatedConv2d(32, 64, 3, 2, 1) self.enc3 GatedConv2d(64, 128, 3, 2, 1) self.dilated nn.ModuleList([ GatedConv2d(128, 128, 3, 1, d, dilationd) for d in [2, 4, 8, 16] ]) self.dec3 GatedConv2d(256, 64, 3, 1, 1) self.dec2 GatedConv2d(128, 32, 3, 1, 1) self.dec1 nn.Conv2d(64, 3, 3, 1, 1) def forward(self, img, mask): x torch.cat([img, mask], dim1) # 输入为4通道 e1 self.enc1(x) e2 self.enc2(e1) e3 self.enc3(e2) d e3 for layer in self.dilated: d layer(d) d # 空洞卷积之间带残差 d torch.cat([d, e3], dim1) d F.interpolate(d, scale_factor2, modebilinear, align_cornersFalse) d torch.cat([d, e2], dim1) d self.dec2(d) d F.interpolate(d, scale_factor2, modebilinear, align_cornersFalse) d torch.cat([d, e1], dim1) d self.dec3(d) out torch.tanh(self.dec1(d)) return out逻辑说明编码器使用 stride2 卷积做下采样不使用 BatchNorm。修复任务中输入通道数很小BatchNorm 在 batch 较小时容易引入噪声门控分支的 sigmoid 本身就提供了一定归一化作用。空洞卷积分支之间每个都做残差相加避免网络加深后梯度消失。上采样阶段使用F.interpolate配合 skip connection比转置卷积产生的棋盘格伪影少得多。最后一层用tanh把输出约束到 [-1,1]这与训练时图像归一化范围保持一致。3.3 损失函数L1重建损失与Hinge对抗损失的组合训练时生成器不仅要让判别器分不出真假还要直接让预测结果接近真实补全。L1 损失是逐像素回归项能约束颜色和空间结构对抗损失负责补充高频纹理细节。两者权重不同结果差异很大。下面给出两个基础函数def l1_loss_in_mask(pred, target, mask): diff torch.abs(pred - target) * mask return diff.sum() / (mask.sum() 1e-8) def hinge_d_loss(disc_real, disc_fake): real_loss torch.mean(torch.relu(1.0 - disc_real)) fake_loss torch.mean(torch.relu(1.0 disc_fake)) return real_loss fake_lossl1_loss_in_mask只在 mask 区域内计算损失而不是全图平均。这样做是为了让生成器把优化重点放在缺失区域避免“整体相似但洞内糊掉”的情况。hinge_d_loss中disc_real是判别器对真实补全图的输出disc_fake是对生成图的输出。真实样本希望输出大于 1生成样本希望输出小于 -1超出边界就不给梯度。这种基于边界的损失比 BCE 更容易保持生成器和判别器的优化平衡。完整训练时判别器损失就是hinge_d_loss(real, fake)生成器损失是-torch.mean(fake) lambda_l1 * l1_loss_in_mask(...)其中-torch.mean(fake)对应 Hinge Loss 下生成器的对抗目标。4. 自由形式 mask 生成与训练 Pipeline4.1 mask 的表示二值单通道还是 RGBA 四通道在 PyTorch 基础框架里我推荐把输入设计成原始图像[B,3,H,W]和二值 mask[B,1,H,W]forward 时在通道维拼接成 4 通道。有些工程喜欢把 mask 写进 RGBA 的 alpha 通道这只是数据载体不同模型内部处理方式一样。重点在于 mask 的取值语义要统一1 代表需要修复0 代表保留。如果从 PNG 读取时把透明部分当作 alpha注意把 0 映射成 1否则会把整个背景当成待修复区域。4.2 随机多边形 mask 生成器自由形式 mask 不是简单的矩形或圆形。实际训练中每张图要叠加多个随机多边形模拟真实世界中的划痕、污渍和遮挡物。生成器需要和图像尺寸绑定mask 太大会让模型失去上下文信息太小则学不到修复能力。下面是可用的随机多边形生成代码import numpy as np import cv2 def random_free_form_mask(h, w, max_parts8, max_vertex12, max_radius80): mask np.zeros((h, w), dtypenp.float32) for _ in range(np.random.randint(1, max_parts)): num_vertex np.random.randint(3, max_vertex) center_x np.random.randint(0, w) center_y np.random.randint(0, h) radius np.random.randint(20, max_radius) angle np.linspace(0, 2 * np.pi, num_vertex, endpointFalse) points [] for a in angle: r radius * (0.5 0.5 * np.random.rand()) points.append([ int(center_x r * np.cos(a)), int(center_y r * np.sin(a)) ]) cv2.fillPoly(mask, [np.array(points, dtypenp.int32)], 1.0) return mask逻辑说明max_parts控制一张图中最多出现的破损区域数量max_vertex控制多边形顶点数顶点越多形状越接近圆。每个顶点的半径在 0.5 到 1.0 倍之间随机缩放所以生成的是不规则多边形。cv2.fillPoly把多边形内部填充为 1。这里建议在数据增强阶段再叠加随机旋转、裁剪避免网络过拟合特定形状。mask 生成需要放在训练线程的预处理阶段不要预先存盘否则每个 epoch 看到的是同一组破损模式。4.3 训练循环与超参数配置训练循环中生成器和判别器要交替更新。我一般每个 iteration 先更新一次判别器再更新一次生成器判别器学习率略高于生成器。下面是一个简化的训练骨架for epoch in range(total_epochs): for img, mask in dataloader: img img.cuda() # [B,3,H,W] 已归一化到[-1,1] mask mask.cuda() # [B,1,H,W] 0/1 masked_img img * (1 - mask) # 把mask区域置零 fake generator(masked_img, mask) # 判别器 real_d discriminator(img) fake_d discriminator(fake.detach()) d_loss hinge_d_loss(real_d, fake_d) d_optimizer.zero_grad() d_loss.backward() d_optimizer.step() # 生成器 fake_d discriminator(fake) g_l1 l1_loss_in_mask(fake, img, mask) g_adv -torch.mean(fake_d) g_loss g_adv lambda_l1 * g_l1 g_optimizer.zero_grad() g_loss.backward() g_optimizer.step()代码中fake.detach()用于切断判别器反向传播到生成器的梯度。g_adv取-torch.mean(fake_d)对应 Hinge Loss 下生成器希望判别器对生成图片给出高置信输出。masked_img是在输入阶段把 mask 区域像素置零这样网络第一层能同时看到周围上下文和被破坏的痕迹。训练中我常用的初始超参数如下超参数推荐值说明输入分辨率256x256显存不够时先降为 224batch_size8单卡跑通再往大调优化器Adam(betas(0.0, 0.9))Hinge 损失下常用生成器学习率2e-4判别器学习率的一半判别器学习率4e-4判别器更新幅度略大lambda_l150重建损失权重5. 训练自由形式修复模型最值得调的 3 个参数5.1 谱归一化的实现位置和判别器学习率门控卷积中use_snTrue会同时对conv_f和conv_g做谱归一化。实际建议把谱归一化放在判别器上生成器的特征卷积分支不加因为生成器需要足够的容量去生成细节。但门控卷积的 gate 分支可以保留 SN约束这个分支的权重能让门控输出更平滑。判别器学习率决定对抗平衡。生成器学习率固定在 2e-4 时判别器学习率从 2e-4 调到 4e-4 是比较好的区间。观察生成图片的 mask 边缘如果边缘不断闪烁说明判别器太强把它降到 2e-4同时把 lambda_l1 提高。这个调参顺序比先改网络结构更快。5.2 L1 重建损失权重 λ从 10 到 100不要一开始给太大L1 权重直接决定生成器是优先处理颜色还是优先处理细节。权重太大网络会倾向于输出平滑的平均色来压低 L1放弃纹理权重太小结构会漂移。在 256x256 输入下λ10/50/100 的表现差异如下λ结构正确性纹理细节mask 边缘10容易歪丰富明显接缝50基本正确保留较好有点柔化100很正偏模糊干净但缺少高频我一般从 λ50 开始。如果发现整体模糊就降到 30如果发现 mask 边界有脏色就升到 100并增加一个感知损失来弥补 L1 导致的平滑问题。注意调整 λ 后不需要重跑整个训练可以在训练进行到一半时修改并继续损失会产生一个跳变随后会快速收敛到新平衡。5.3 PatchGAN 判别器输出尺寸与输入分辨率的匹配SN-PatchGAN 的判别器最后输出一个矩阵每个输出点对应原图一个 patch而不是整图一个标量。patch 的感受野要能覆盖自由形式 mask 的典型尺寸。如果 mask 平均半径是 30 像素而判别器输出 patch 感受野只有 16 像素那么判别器看不到足够上下文无法判断结构是否合理。排查方法是在 debug 时打印判别器输出形状x torch.randn(8, 3, 256, 256) out discriminator(x) print(out.shape) # 例如 [8,1,64,64]如果输入 256x256输出是 64x64那么每个输出点大约对应原图 32x32 感受野。想要覆盖更大 mask可以减少判别器的下采样次数或者在瓶颈层加空洞卷积。否则你会看到损失在下降但 mask 内部结构不对边缘又连续断裂。这个参数比学习率更隐蔽第一次复现时一定要先确认。6. 验证与推理mask 边缘羽化与评估指标6.1 用 PSNR/SSIM 评估不要只看眼睛自由形式修复效果评估除了肉眼观察我至少会算 mask 区域的 PSNR 和 SSIM。注意标准 PSNR 是全图计算当 mask 区域占比小时背景会把指标抬得很高掩盖洞内的糟糕结果。正确做法是只在 mask 区域内计算def psnr_mask(pred, target, mask): mse ((pred - target) * mask).pow(2).sum() / (mask.sum() 1e-8) return 20 * torch.log10(255.0 / torch.sqrt(mse 1e-8)).item()这里 pred 和 target 的取值范围是 [-1,1]严格来说需要先反归一化到 [0,255] 再算 PSNR。实际工程里为了方便可以先把 pred 和 target 乘 127.5 加 127.5再传入上面的函数。注意 mask 区域占比不能太小至少要占全图的 1% 到 10%否则指标方差很大一次随机 mask 就能让 PSNR 波动 2dB 以上。6.2 推理时对 mask 做边缘羽化的收益最后一个技巧是推理阶段的 mask 羽化。训练时使用二值 mask但推理时如果直接把二值 mask 拼输入生成区域和原图之间经常出现 1 像素的接缝。原因是模型输出的 mask 区域边缘像素没有和原图完全融合。常见做法是在推理时对 mask 做高斯羽化让网络看到软边缘最终合成时过渡自然。import cv2, torch, numpy as np def feather_mask(mask_tensor, ks7, sigma2.0): mask_np mask_tensor[0, 0].cpu().numpy() soft cv2.GaussianBlur(mask_np, (ks, ks), sigma) soft np.clip(soft, 0.0, 1.0) return torch.from_numpy(soft).float().unsqueeze(0).unsqueeze(0)mask_tensor 是二值 mask经过高斯模糊后变成 0 到 1 的软 mask。在推理 forward 中使用 soft mask 代替二值 mask门控卷积会看到边缘位置的门控值连续过渡从而减少接缝。注意 PSNR 和 SSIM 评估时仍然用硬 mask否则指标被人为提高。我一般把这个羽化函数放在推理脚本里作为演示和对比展示时的后处理开关不进入正式评估流程。2 到 3 像素的羽化宽度就能明显改善视觉接缝而不影响整体修复结构。本文还有配套的精品资源点击获取
返回列表