拓冰建站拓冰建站
首页 / 资讯中心 / 正文

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

简介这是一份基于PyTorch的DeepFillv2门控卷积自由形式图像修复实现对应arXiv 1806.03589论文适合有一定深度学习基础、正在复现或扩展图像修复算法的开发者和研究人员。压缩包共91个文件体积约3.42MB内容围绕模型、训练、展示与配置四大部分展开11个Python脚本覆盖网络定义、损失函数、训练测试流程及数据预处理YAML文件分别提供CelebA、Places等数据集的训练参数JS/CSS/HTML与示例图片构成了可交互的前端演示界面便于直观对比修复效果。工程额外包含预训练权重、测试Notebook以及风格迁移示例无需从零训练即可体验门控卷积在自由掩码下的修复能力也能替换配置和数据集继续调优。目前已有460人学习浏览对理解DeepFillv2算法结构、训练细节和快速部署演示均有实用价值。1. 为什么 DeepFillv2 的门控卷积成了自由形式修复的事实标准自由形式图像修复跟“抠掉矩形再补一块”完全不同用户的涂鸦可能是细划痕、遮挡物轮廓甚至是不规则大洞。传统卷积面对这种任意形状掩码时会把空白区域的值不断卷进有效纹理越修越糊。DeepFillv2 的答案是门控卷积——让模型自己学会每个空间位置该不该放行特征而不是靠一套手工规则去更新掩码。这个仓库把论文1806.03589里的 coarse-to-refine 生成器、SN-PatchGAN 和感知损失用 PyTorch 重写了一遍并附带了测试 notebook、可交互的 app.py 和 CelebA/Places 训练配置。适合要复现论文、做修复对比实验或想给现有图像编辑流程加一个高质量 inpaint 模块的人。2. 门控卷积原理与 DeepFillv2 网络结构拆解2.1 从普通卷积到门控卷积掩码不再靠手写规则传播普通卷积在特征图上用固定权重扫过每个像素网络并不知道哪些位置对应原始缺失区域。之前的 partial convolution 做法是维护一张二进制 mask每次卷积后按照“只要卷积核覆盖到至少一个有效像素就把该位置标成有效”再对 feature 做重归一化。这套规则在浅层有效但到了深层mask 和真实语义边界会逐渐错位属于“人工固化规则”。门控卷积把这种规则换成可学习机制每个卷积位置同时计算 feature 分支和 gate 分支gate 分支经过 sigmoid 得到一个 0 到 1 的软开关。最后输出是 feature 与 gate 的逐元素乘积。模型可以针对某一类掩码、某一种纹理自由调整开关而不是只做二值判断。PyTorch 里最直接的自定义写法如下class GatedConv2d(nn.Conv2d): def __init__(self, in_channels, out_channels, kernel_size3, stride1, padding1, dilation1): super().__init__(in_channels, out_channels, kernel_size, stride, padding, dilation) self.gate_conv nn.Conv2d(in_channels, out_channels, kernel_size, stride, padding, dilation) nn.init.xavier_uniform_(self.gate_conv.weight) nn.init.zeros_(self.gate_conv.bias) def forward(self, x): feature super().forward(x) gate torch.sigmoid(self.gate_conv(x)) return feature * gate这段代码里有几个值得注意的地方super().forward(x)会走nn.Conv2d正常的初始化和卷积逻辑省掉自己写F.conv2d的繁琐参数传递gate_conv的卷积核尺寸、padding、dilation 必须和主分支完全一致否则空间对齐会出错sigmoid让每个门控值始终落在[0,1]训练时梯度不会像二值掩码那样断裂。实际仓库里的model/networks.py通常会对这两个分支分别初始化权重并且把 gate 分支的 bias 初始化为常数让模型一开始尽量全部放行避免收敛太慢。2.2 Coarse 到 Refinement 的两阶段生成器DeepFillv2 不是一锅端直接输出结果而是先粗后精。Coarse 网络负责“填空深度”输入[原图 掩码]先用一系列下采样和门控卷积把缺失区域填出整体结构输出一张粗糙但语义正确的结果。Refinement 网络再把它和原图重新拼接输入一个编码器-解码器结构用跳跃连接、密集连接块和全分辨率门控卷积把纹理细节补上。这个两阶段设计的关键在于粗细分工明确。Coarse 阶段不需要高频细节所以可以用步长卷积降分辨率扩大感受野Refinement 阶段则必须保持空间对齐常常在不上采样的前提下用膨胀卷积来聚合远处信息。仓库里的networks.py会把GatedConv2d封装成可重复堆叠的 block并在每个 block 后决定是否接上下采样或激活函数。有差异的版本还会把 Coarse 的输出作为 Refinement 输入的第二个分支和原图在通道维拼接而不是直接把空洞放到 Refinement 里这样梯度回传更直接。从张量形状上看假设输入是[B, 4, 512, 512]其中前三通道是 RGB最后一通道是掩码。Coarse 输出是[B, 3, 512, 512]Refinement 输入往往是[B, 6, 512, 512]或者[B, 3, 512, 512] mask拼接后的形状具体看实现。你看到某个训练脚本里有concat(coarse, image, mask, dim1)这类操作多半就是这样走的。2.3 SN-PatchGAN 判别器与损失函数训练修复模型只用像素级 L1 或 L2 会得到平滑模糊结果。DeepFillv2 引入了带谱归一化的 PatchGAN 判别器不对整张图打真/假标签而是对特征图上的若干局部 patch 分别打分。这个仓库叫它 SN-PatchGAN在model/losses.py里会同时算下面几项损失项计算内容作用L1 loss修复图与真图逐像素差保证整体颜色和亮度接近Perceptual lossVGG 中间层特征的 L1 距离约束语义和结构信息Style lossVGG 特征图的 Gram 矩阵距离约束纹理风格一致性Hinge GAN lossSN-PatchGAN 判别器输出逼真度抑制模糊感知损失和风格损失都要提取 VGG 特征所以训练时显存占用不小。常见做法是分别用relu1_2、relu2_2、relu3_2、relu4_2四层取每层的 L1 和 Gram 矩阵差异加权求和。GAN 输出使用 hinge loss配合谱归一化能有效避免判别器收敛过快导致生成器梯度消失。仓库里train.py对这几项损失会有一个权重字典比如perceptual0.01, style0.01, l11.0, gan0.1具体要按训练数据规模微调。3. 本地环境准备与预训练模型加载3.1 依赖安装与 PyTorch 版本注意点直接跑这个项目前建议先建一个干净的 Python 虚拟环境避免系统里已有的 PyTorch 或 OpenCV 版本互相干扰。把下载的deepfillv2-pytorch-master目录解压后在项目根目录创建虚拟环境并安装依赖cd deepfillv2-pytorch-master python -m venv venv source venv/bin/activate pip install torch torchvision --index-url https://download.pytorch.org/whl/cu121 pip install pyyaml opencv-python tqdm gradio安装时要注意 PyTorch 2.x 的torchvision和torch版本必须配套否则导入torchvision.models.vgg会直接报符号冲突。CUDA 版本不对时可以改用 CPU 版去掉--index-url直接pip install torch torchvision测试推理也能跑只是慢一些。如果机器显存小于 6GB建议把configs/train.yaml里的batch_size改成 1并在测试阶段固定输入尺寸避免掩码形状变化导致 cuDNN 反复自动调优。3.2 目录结构与配置文件这个仓库的配置采用 YAML 方式管理根目录和configs下有好几个.yaml文件它们的职责很不一样文件典型用途configs/train.yaml通用训练超参数包括学习率、批大小、损失权重configs/train-celeb.yamlCelebA 数据集路径和预处理参数configs/train-places.yamlPlaces2 数据集路径和预处理参数models.yaml模型结构参数如通道数、block 数量、掩码输入方式frontend/configsapp.py 的界面绑定和默认模型路径examples/inpaint样例图片和掩码图片模型结构和训练参数分开是合理的工程习惯。改通道数时只动models.yaml而调学习率只动train.yaml。如果你要修改归一化方式需要同时检查这两个文件里的mean和std字段并且和数据加载器保持一致。读取 YAML 的常见写法在utils/misc.py里都会有核心逻辑是import yaml with open(models.yaml, r, encodingutf-8) as f: model_config yaml.safe_load(f) generator_cfg model_config[generator] print(generator_cfg[conv_type]) # 应为 gated_conv print(generator_cfg[refinement_layers])代码中的yaml.safe_load避免加载任意 Python 对象这是配置文件的安全下限。models.yaml里的conv_type字段如果写成了gated或partial网络初始化会走不同分支千万不能只看文件目录结构。3.3 加载官方/重新实现的预训练权重跑通测试仓库的pretrained目录用来放权重。官方给出的权重通常是 TensorFlow checkpoint这个 PyTorch 重实现里还有一个networks_tf.py文件目的是把 TF 权重按名字映射到 PyTorch 的 state dict。简单做法是从仓库 release 或资源包附带的网盘下载.pth权重放到pretrained/后执行测试python test.py \ --model_path pretrained/places.pth \ --config configs/train.yaml \ --input examples/inpaint/input.png \ --mask examples/inpaint/mask.png \ --output results/result.png \ --device cuda:0参数说明--model_path是预训练权重路径--config指定模型结构配置--input和--mask分别是原图和掩码图片路径--output指定保存路径--device选择 CPU 或 GPU。如果你的版本没有test.py命令直接打开test.ipynb逐格执行前三格通常是在读配置、初始化网络和加载权重第四格开始跑单张测试。4. 自由形式修复实战从单图测试到批量推理4.1 掩码的生成与格式门控卷积能处理掩码但掩码本身的格式直接影响修复质量。最稳妥的掩码是单通道 PNG 或 BMP黑色表示保留像素白色表示待修复区域。不要直接用带抗锯齿的画笔存成 JPG边缘半透明像素会让门控拿不准该不该放行特征。训练时常用随机折线、随机圆和随机噪声点组合生成自由形式掩码测试阶段也可以手动生成import cv2 import numpy as np image cv2.imread(examples/inpaint/input.png, cv2.IMREAD_COLOR) h, w image.shape[:2] mask np.zeros((h, w), dtypenp.uint8) # 用三条折线模拟划痕 pts1 np.array([[100, 100], [200, 180], [300, 90]], dtypenp.int32) pts2 np.array([[400, 300], [350, 420], [500, 500]], dtypenp.int32) cv2.polylines(mask, [pts1, pts2], False, 255, thickness12) # 再加一个不规则椭圆 cv2.ellipse(mask, (300, 300), (50, 90), 30, 0, 360, 255, -1) cv2.imwrite(mask.png, mask)这里用polylines的False参数表示开放折线适合模拟真实笔划ellipse的厚度参数为-1表示实心填充。生成后务必预览一下掩码和原图的覆盖关系很多修复失败案例都是掩码边缘正好切在关键纹理边界上。如果你的目标是去除照片上的日期水印用矩形框也行但 DeepFillv2 的优势在于任意形状尽量模拟实际涂抹形态。4.2 批量推理与参数调整单张测试通过后批量处理时会发现不同掩码形状对参数敏感度不同。这里列一个参考表掩码类型典型场景建议设置细长划痕照片老旧损伤掩码膨胀 2-3 像素感知损失权重适当提高大面积物体遮挡物体删除关注 coarse 阶段质量批量时可先跑一次纯 L1 对比随机小块噪声异常像素清洗用中值滤波预清洗比修复更省时间半透明覆盖物字幕/贴纸需要把覆盖物区域按不透明度转为纯掩码批量推理时不要直接套训练配置。训练为了增强多样性会随机 resize 和 crop推理则建议关闭所有数据增强保证输入尺寸一致。如果发现修复结果出现重复纹理通常是把refinement_layers设置得过深可以在models.yaml里减少几个 block。用测试脚本跑一张结果后用cv2.absdiff对比原图与输出的边界区域确认是不是只有掩码内部被修改。4.3 用 app.py 交互式刷图修复仓库的app.py和frontend/configs提供交互式界面适合先试效果再落参数。启动方式一般是python app.py \ --checkpoint pretrained/celeba.pth \ --config models.yaml \ --port 7860启动后浏览器打开http://127.0.0.1:7860。界面里左侧放原图中间用画笔或橡皮工具绘制掩码区域右侧直接显示修复结果。这类工具通常基于 Gradio画笔粗细和 mask 显示为半透明红色提交后调用与test.py完全相同的推理函数。交互环节的价值在于快速确认“问题到底出在掩码生成还是修复网络”如果同一个掩码在 app.py 里结果正常但脚本批量处理不正常问题大概率出在你批量读取图像时丢了 alpha 通道或颜色顺序。5. 进阶用掩码膨胀与冻结模块提升修复边界质量5.1 掩码膨胀规则修复结果的拼接处偶尔会有一圈发暗或发白的边这是门控卷积在边缘处对大局感受野和局部纹理博弈产生的自然误差。一个不用重训模型就能明显改善的做法是推理时把传给网络的掩码先膨胀几像素得到宽松掩码网络在更大范围上做平滑过渡保存结果时再用原始掩码把原图抠回。这样网络有足够上下文去处理边界而最终输出的有效区域收敛到真正需要修复的范围。膨胀操作推荐用cv2.dilate核大小根据掩码线宽决定。kernel np.ones((3, 3), np.uint8) mask_dilated cv2.dilate(mask, kernel, iterations2) # mask_dilated 用于网络输入mask 用于输出融合迭代次数过多会明显增加计算量并且可能把旁边不相干的物体边缘也纳入修复区。如果只修细划痕kernel3, iterations1就够了大面积遮罩可以到 3 次不要超过 4 次。5.2 冻结训练与增量微调当你想在一个新场景上更贴合时不必从零训练。把预训练权重加载后冻结coarse网络只微调refinement网络能让模型更快适应同一类掩码。PyTorch 里给 optimizer 只传需要更新的参数即可注意冻结时把 BatchNorm 也切到 eval 模式否则计算图会额外保存中间统计量。5.3 显存不够时的 sliding-window 处理超大图修复时把整张图送进网络通常超过显存。换滑窗需要避免窗口边界出现拼接缝窗口之间重叠 16 像素修复后再用 alpha mask 做加权融合重叠区域权重按距离线性过渡。另一个技巧是把修复后的图转成 RGBA 存储alpha 通道就记录原始掩码后续无论做风格迁移还是再编辑都不用重复读掩码文件。本文还有配套的精品资源点击获取
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门