资讯详情

基于GAN的图像修复:部分卷积与训练调优实战

📅 2026/9/12 14:15:03 | 华诺云谱 👁 阅读
基于GAN的图像修复:部分卷积与训练调优实战
简介基于 Python 的深度生成对抗网络 GAN 图像修复项目面向计算机相关专业毕业设计、期末大作业及深度学习实战练习人群也适合想了解图像补全原理的初学者参考。项目覆盖从数据集处理、生成器与判别器结构搭建、损失计算到图像补全推理的完整流程核心代码拆分为 utils、model、ops、train-dcgan、simple-distributions、complete 等模块各自职责清晰随附 README 文档对运行方式和设计思路进行说明便于快速上手与二次改造。压缩包共 7 个文件包含 6 个 Python 脚本和 1 个 Markdown 文档整体仅 12KB结构紧凑、无冗余数据适合直接阅读源码并结合文档理解 GAN 的关键实现。已有 164 人浏览学习代码经本地编译调试可正常运行项目设计评审得分 98 分属于难度适中、完成度较高的高分参考模板。对需要完成课程设计、期末大作业或毕业设计的学生来说这份资源既可提供完整项目骨架也能帮助梳理 GAN 图像修复的实验思路。1. 为什么图像修复要选GAN而不选传统插值与CV修补拿一张破损的老照片用OpenCV的inpaint函数跑一遍洞是填上了但放大看全是周边像素的模糊延拓表情和结构完全对不上。传统图像修复算法的前提是“缺失区域和已知区域的纹理统计一致”遇到大块遮挡或语义复杂区域就露馅。深度生成对抗网络GAN换了一个思路把修复过程当成条件生成问题生成器先预测“这个空缺最可能是什么”再画出对应纹理判别器负责对结果做真伪检验。基于Python实现这样一个图像修复模型训练系统只需要两个网络和几条经过平衡的损失项不需要任何手工特征介入这也是它能处理人脸五官修复、物体移除、划痕重建这批任务的主要原因。适合的读者是已经跑过基础PyTorch分类任务的工程师手头有显存8G以上的GPU想在生成任务上找一个完整可落地的训练链路。2. 图像修复的建模前提掩码策略与Python数据管线修复任务在数学上可以抽象成给定原图I和掩码MM中值为1的位置是缺损区0是保留区。送入模型的观察图是I_obs I * (1 - M)模型要生成完整图I_pred并且让I_pred在M区域与真实I在像素级、语义级和视觉真实度三个层面都对齐。这个建模方式决定了后续所有设计尤其是掩码怎么生成、怎么参与前向传播。第一个容易踩的坑是直接把I_obs拼上M丢给普通卷积网络。普通卷积会把缺损位置的0值当成“像素本身是黑的”在窗口滑动时这些无效像素照样参与加权求和结果就是修复区域旁边出现一圈明显的灰黑色污渍。常见做法是引入部分卷积Partial Convolution它的核心改动是每个卷积窗口先统计有效像素数量再按有效数量重新归一化输出无效位置的信息不会被当成特征吸收。实践中由NVIDIA提出的部分卷积层是修复GAN的常用基础设施实现成本不高但效果差异巨大。2.1 用Python批量生成不规则掩码训练GAN时掩码不能只用方形或圆形。方形掩码会让网络记住“从四边向中心补全”的先验换到真实破损照片时一塌糊涂。真实场景里的污渍、划痕、遮挡物边缘都是不规则的所以掩码生成器需要支持随机多边形、狭长裂缝、多块区域组合等形态。import cv2 import numpy as np def make_irregular_mask(shape, max_holes6, max_ratio0.25): 生成不规则0/1掩码1表示待修复区域 h, w shape mask np.zeros((h, w), dtypenp.float32) for _ in range(np.random.randint(1, max_holes 1)): cx np.random.randint(0, w) cy np.random.randint(0, h) radius np.random.randint(int(0.08 * h), int(max_ratio * h)) pts [] # 用锯齿多边形模拟真实破损边缘 for k in range(6): angle 2 * np.pi * k / 6 np.random.uniform(0.1, 0.5) r radius * np.random.uniform(0.4, 1.2) pts.append([int(cx r * np.cos(angle)), int(cy r * np.sin(angle))]) cv2.fillPoly(mask, [np.array(pts, np.int32)], 1.0) return mask逻辑说明每个孔洞由一个中心点和一个基准半径决定顶点在圆周附近随机抖动形成不规则的闭合多边形。cv2.fillPoly把多边形内部填成1外部保持0。如果把max_holes调大但max_ratio不变掩码会变成碎斑状更接近小面积多点污损把max_ratio提到0.4以上则变成大块遮挡对生成器的语义预测能力要求更高。训练时通常采用混合策略而不是固定单一掩码类型这样模型能适应更广的修复尺度掩码类型面积占比范围模拟场景训练建议不规则多边形10% ~ 30%物体移除、纸张破损主体占比60%细长划痕5% ~ 15%老照片划痕占20%防止过度平滑多块随机斑点10% ~ 20%污渍、喷溅占15%增强鲁棒性居中偏大遮挡30% ~ 45%水印、人物遮挡占10%提高上限初次实验建议把最大面积控制在30%以内。掩码面积过大时保留区域提供的上下文线索过少生成器基本靠猜训练初期很难收敛且判别器极易抓住明显的生成痕迹导致对抗损失提前饱和。2.2 归一化区间与数据加载配置修复模型的图像张量通常归一化到[-1, 1]这和生成器输出层使用tanh激活是配套的。tanh的输出本身就在[-1, 1]范围内如果输入用[0, 1]或[0, 255]生成器最后一步就要额外接一个裁剪或者缩放梯度传播路径变长且容易出现数值不平衡。from torch.utils.data import Dataset class InpaintingDataset(Dataset): def __init__(self, image_paths, size256): self.paths image_paths self.size size def __getitem__(self, idx): img cv2.imread(self.paths[idx]) img cv2.cvtColor(img, cv2.COLOR_BGR2RGB) img cv2.resize(img, (self.size, self.size)) img (img.astype(np.float32) / 127.5) - 1.0 mask make_irregular_mask((self.size, self.size)) masked img * (1 - mask[..., None]) return (img.transpose(2, 0, 1).copy(), masked.transpose(2, 0, 1).copy(), mask[None].copy()) def __len__(self): return len(self.paths)这里masked的计算用的是img * (1 - mask)缺损位置直接被置为0。部分卷积层会依据掩码识别哪些位置无效因此这种置0不会污染训练。但如果换成普通卷积生成器这个0值就会被当成黑色像素参与卷积计算属于常见的误用点。数据集预处理还需要注意一个细节水平翻转和随机裁剪必须图像与掩码同步进行。如果只翻转图像不翻转掩码相当于给网络引入了“破损总是偏左”的位置先验如果裁剪时图像和掩码的坐标错位则观察图和掩码失去了对应关系。最稳妥的做法是用同一个random种子控制所有变换。2.3 验证集掩码的独立性训练时掩码每次随机生成但验证集如果也每轮重新生成指标就无法横向比较。建议固定一个验证集掩码文件包含10到20组原图掩码对保存成npy文件。这样不同epoch之间的损失曲线才有可比性否则因为掩码随机性带来的波动会被误判成训练不稳定。3. 生成器、判别器与损失组合的设计图像修复模型的网络设计可以从三个角度拆开看生成器负责补全内容判别器负责评估真假损失函数负责把“补全”引导到语义正确的方向。三者互相制约任何一个选择都会直接影响最终的修复上限。3.1 生成器结构选型U-Net与部分卷积的结合生成器最常见的选择是U-Net结构编码器逐步压缩空间分辨率提取高层语义解码器逐步还原细节中间用跳跃连接把编码器特征直接传给解码器。跳跃连接在修复任务里尤其重要它把缺损区域边缘的局部纹理传入深层修复结果才能保留原始照片的细节风格。但标准U-Net配合普通卷积处理掩码输入时需要额外处理无效像素。一个工程上更稳妥的做法是把普通卷积替换为部分卷积。部分卷积的PyTorch实现核心是维护一个随网络下降而更新的掩码import torch import torch.nn as nn class PartialConv2d(nn.Module): def __init__(self, in_ch, out_ch, kernel_size3, stride1, padding1): super().__init__() self.conv nn.Conv2d(in_ch, out_ch, kernel_size, stride, padding, biasTrue) # 掩码卷积的权重恒为1只用于统计有效像素个数 weight torch.ones(in_ch, 1, kernel_size, kernel_size) self.mask_conv nn.Conv2d(in_ch, 1, kernel_size, stride, padding, biasFalse) self.mask_conv.weight nn.Parameter(weight, requires_gradFalse) def forward(self, x, mask): out self.conv(x * mask) with torch.no_grad(): update_mask self.mask_conv(mask) # 归一化有效像素越多输出越接近普通卷积 scale 1.0 / (update_mask 1e-8) scale scale.clamp(max200.0) return out * scale, update_mask.clamp(0.0, 1.0)这段代码的关键逻辑在scale的计算卷积窗口内有效像素数为update_mask如果窗口完全被掩码覆盖update_mask接近0scale会被clamp限制在一个合理上限避免输出爆炸。1e-8的加项是防除零。同时更新后的掩码还要传给下一层因为随着下采样掩码会逐步变小网络需要知道哪些区域仍然是可信任的已知像素。整个生成器就是一组部分卷积块的下采样和上采样组合下采样阶段把空间尺寸从256降到32通道数从16逐步加到256上采样阶段优先使用最近邻插值或像素重排避免转置卷积导致的棋盘纹理伪影。3.2 判别器设计PatchGAN与掩码输入判别器的作用是判断给定图像是真实还是修复出来的。常见做法是PatchGAN式的判别器输出不是单个标量而是一个N*N的矩阵每个元素只负责判断图像一个局部区域的真伪。这样做的一个直接好处是生成器没法只靠整体色调蒙混过关每个局部块都要足够真实。判别器输入需要拼接掩码通道原因是如果不给判别器看掩码位置它就只能靠边缘痕迹判断真假。拼接掩码后判别器能聚焦到修复区域评分更精准。这一点在图像修复模型里和普通图像生成的判别器有明显区别。判别器网络的输入通道因此是4RGB三通道加掩码通道。层结构可按风格化GAN的经典配置来搭稳定做法是每隔一层步长2下采样通道数从64逐步倍增输出层不接归一化直接用LeakyReLU激活。对抗损失的训练目标建议使用最小二乘形式L_D 0.5 * (D(real)^2 (D(fake) - 1)^2) L_G 0.5 * (D(fake) - 1)^2LSGAN形式的损失函数带来的梯度变化比原始GAN的交叉熵更平稳因为它在优化D(fake)趋近1的过程中不会出现梯度过早饱和。实际训练里判别器损失会周期性回弹这是正常现象不必因为某几个batch的剧烈波动就中断训练。3.3 损失函数组合重建、感知与对抗的平衡只用对抗损失训练生成器隐患是颜色偏移和细节纹理“自由发挥”整体看协调但和目标图像像素对不上只用L1损失训练结果会偏向模糊的像素均值因为L1的贝叶斯最优解就是中位数模糊。所以修复模型几乎都要把重建损失与对抗损失混合。def generator_loss(fake, real, mask, d_fake, perceptual_loss): # L1重建损失只计算掩码区域内的误差 l1 torch.abs(fake - real) * mask l1 l1.sum() / mask.sum() # 感知损失基于VGG特征 perc perceptual_loss(fake, real) # 对抗损失LSGAN形式 adv torch.mean((d_fake - 1) ** 2) total 1.0 * adv 30.0 * l1 0.5 * perc return total, {adv: adv.item(), l1: l1.item(), perc: perc.item()}损失权重参考范围如下表实际项目围绕这个基准做增减损失项权重范围作用权重过高时的副作用对抗损失0.5 ~ 2.0保证真实感纹理过度生成细节失真L1重建损失10 ~ 50保证全局结构结果平滑失去细节感知损失0.1 ~ 1.0保证语义一致高频纹理不丰富感知损失的意义是从高层特征约束“语义一致”而不是像素一致。具体实现通常用ImageNet预训练的VGG16取relu1_2到relu3_3之间的特征层做L1距离。特征需要归一化到ImageNet的标准范围否则预训练权重的统计分布不一致提取出的特征值偏差会导致感知损失数值异常偏大。4. GAN训练不稳定排查从迭代曲线到收敛状态训练GAN修复模型最让人头疼的就是训练曲线不能直接等同于修复质量。损失降得漂亮不代表边缘清楚损失震荡也不代表训练失败。实际排查逻辑可以分成三层数值异常、局部失败、全局不收敛。4.1 标准训练环路与优化器配置训练过程建议生成器和判别器交替更新且两者使用相同的优化器超参学习率从2e-4起步。优化器用Adam实际上已经是工程惯例其中betas参数的设置值得注意默认的(0.9, 0.999)在生成任务里会带来较严重的震荡常见做法是改成(0.5, 0.999)历史上这个配置在DCGAN和Pix2Pix等模型的训练中被反复验证过。import torch def train_one_epoch(gen, disc, dataloader, opt_g, opt_d, percep, device): for img, masked, mask in dataloader: img img.to(device) masked masked.to(device) mask mask.to(device) fake gen(masked, mask) # 判别器真样本标签1假样本标签0 d_real disc(img, mask) d_fake disc(fake.detach(), mask) loss_d 0.5 * (torch.mean(d_real ** 2) torch.mean((d_fake - 1) ** 2)) opt_d.zero_grad() loss_d.backward() opt_d.step() # 生成器目标是让判别器认为假样本是真的 d_fake_for_g disc(fake, mask) loss_adv torch.mean((d_fake_for_g - 1) ** 2) loss_l1 (torch.abs(fake - img) * mask).sum() / mask.sum() loss_perc percep(fake, img) loss_g loss_adv 30.0 * loss_l1 0.3 * loss_perc opt_g.zero_grad() loss_g.backward() opt_g.step()代码里的循环顺序是固定的先更新判别器再更新生成器。fake.detach()的作用是切断生成器梯度回传使判别器梯度只影响判别器自身参数。在loss_l1中要注意除的是mask.sum()不是整个图像的像素数否则掩码面积小时重建损失被稀释。生成器的梯度是四条路径的加和任意一个梯度过大都会污染另外几个可以在backward()之前对每个loss乘上权重也可以在step()前用torch.nn.utils.clip_grad_norm_对生成器参数做整体裁剪阈值经验值是1.0。训练循环每跑完一个epoch需要额外保存一个完整图例作为定性观察依据。从代码角度这个与模型权重同等重要的是把每轮的fake结果汇总拼接成一张对比图和上一个epoch放在一起训练中途就可以直观看到修复区域有没有逐步变清晰。4.2 模式坍缩与判别器过强是两种不同故障在排查训练问题时最重要的是分清“模式坍缩”和“判别器过强”这两种故障形态。模式坍缩的典型表现是生成器对所有遮挡区域输出同一个模板哪怕原图差异巨大。从训练日志看生成器损失持续下降判别器损失也能保持稳定但样例图几乎没有变化或者纹理区域一直是同一个色块。处理手段降低判别器学习率把它从2e-4降到5e-5给生成器更多追赶空间增大L1重建损失权重让像素级损失限制生成器的表达自由度在生成器中增加Dropout或者随机丢弃部分跳跃连接干扰生成器记忆训练样本。判别器过强的表现则是判别器损失极低生成器损失一路下不去样例图全是模糊残影。处理手段是反过来加大生成器的更新频率或者临时在判别器输入加入高斯噪声迫使判别器和真样本之间的分界线变得不那么严格。这里的噪声标准差建议从0.01起步。另外有一个常被忽略的外部因素是batch size。GAN训练对batch size非常敏感显存允许的情况下尽量用16以上。batch越小判别器在单个batch上看到的样本越少统计噪声越大训练陷入震荡的概率越高。如果显存被生成器大模型占满可以先用torch.cuda.amp混合精度训练释放显存而不是一味降低batch size。4.3 用张量指标判断修复质量的辅助手段python -c import torch from model import SimpleInpaintNet model SimpleInpaintNet() assert torch.cuda.is_available(), No GPU found model.cuda() # 构造固定形状的假输入验证前向传播通畅 dummy_img torch.randn(1, 3, 256, 256, devicecuda) dummy_mask torch.rand(1, 1, 256, 256, devicecuda).round() out model(dummy_img, dummy_mask) print(output shape:, out.shape) 在训练正式启动前跑通这个小脚本能提前暴露输入输出通道数不匹配、掩码类型错误这类问题训练到一半发现网络输出尺寸不对通常浪费的就不止半小时训练时间了。在训练中途要直接定量看修复区域的质量可以临时计算掩码区域的平均L1误差但更实用的指标是生成样本与真实样本之间的峰值信噪比或SSIM。因为SSIM对局部亮度、对比度和结构都做了建模比均方误差更能反映人眼对边缘纹理的感知。建议每50个epoch计算一次而不必每epoch都全量验证否则验证时间会超过训练时间。5. 验证图像修复效果与把模型用在大分辨率输入上5.1 掩码区域定向评估验证修复效果不能只看整张图的指标。因为原图中90%的像素没有缺失模型即使什么都不做只让剩余区域保持不变全图PSNR也能高达40以上。所以评估注意力要集中在掩码区域先根据掩码外接矩形裁剪局部区域再在该区域上计算SSIM或FID。SSIM适合单张图对比FID适合整个测试集与生成集分布对比。在不够充分的评估条件下SSIM是性价比最高的单图度量因为它本身包含了对局部窗口亮度、对比度和结构信息的综合比较比PSNR更接近人的感知判断。对清晰度要求较高的修复场景可以额外用拉普拉斯算子方差统计修复区域边缘的锐度。如果拉普拉斯方差偏低说明修复区域过度平滑这是L1损失权重过大时常见的结果此时应当适当降低L1的系数而不是无限堆更多层网络。5.2 大分辨率输入的分块推理策略训练时模型的输入固定为256或512推理时遇到几千像素的扫描件直接把全图缩放会丢失细节而且大图直接前向传播的显存占用也会超出限制。常见做法是把输入图分割成256×256的块块与块之间预留32像素重叠推理完成后在重叠区域做线性融合。融合权重按像素到块中心距离衰减越靠近块边界权重越低。拼接时再单独处理掩码每个块的掩码单独生成避免掩码跨越块边界导致修复语义割裂。def infer_large(gen, img, mask, block_size256, stride224): h, w img.shape[:2] out np.zeros_like(img, dtypenp.float32) weight np.zeros((h, w, 1), dtypenp.float32) for y in range(0, h - block_size, stride): for x in range(0, w - block_size, stride): block img[y:y block_size, x:x block_size] block_mask mask[y:y block_size, x:x block_size] # 前向推理结果写入out result gen(block, block_mask) out[y:y block_size, x:x block_size] result weight[y:y block_size, x:x block_size] 1.0 out out / np.maximum(weight, 1.0) return out分块还有一个额外好处同一个掩码可以在不同分块位置复用起到类似测试时数据增强的作用。如果多个分块给出的修复结果在重叠区域高度一致可以认为该处的修复是稳定的如果每次推理差异很大说明模型对这块区域的语义预测还不太确定这时候就算放大模型也无济于事应该增加训练数据多样性。5.3 实验记录技巧用掩码面积反向标定模型能力训练稳定后可以额外做一次面积梯度测试把同一张图的掩码面积从10%逐级增加到50%每级跑固定数次推理记录SSIM下降曲线。这条曲线比单点修复效果更值得存档因为它能标定出模型的鲁棒边界曲线在某个面积点突然断崖式下跌说明模型语义预测能力的阈值就在那里。后续再增加训练数据或调大模型容量时对比这条曲线就能确认改动是否真的提升了修复能力而不是只对测试集几张图有效。这个习惯对迭代模型版本非常有用推荐长期沿用。本文还有配套的精品资源点击获取
📝

华诺云谱内容团队

资深建站顾问 · 行业研究员

10年+企业数字化服务经验,专注智能建站、SEO优化与品牌营销,持续输出建站技巧、行业洞察与营销干货,已帮助5000+企业实现数字化增长。

你可能需要的服务

订阅华诺云谱资讯周报

每周一封,精选建站技巧、SEO与营销干货,直达邮箱。已有 8,000+ 企业主订阅,助你少走弯路。