GAN图像修复实战:从掩码设计到U-Net与PatchGAN训练优化

发布时间:2026/10/11 10:17:20
GAN图像修复实战:从掩码设计到U-Net与PatchGAN训练优化 简介面向毕业设计、课程设计与项目开发场景基于Python实现生成对抗网络GAN的破损图片修复项目核心包含cGAN模型、图像修复流程及SSIM/PSNR评估脚本适合具备一定深度学习基础的读者学习图像生成与修复技术。压缩包共65个文件涵盖47张png图片含修复前后对比、质量评估曲线、6个Python程序模型训练、图像修复、测试及结果打印、2个pb格式模型权重文件以及若干项目配置文件整体体积约2.92MB。目前已有86人浏览学习。源码经过严格测试可直接运行参考并支持在原有基础上扩展如替换数据集、调整网络结构等便于完成课题设计或深入理解对抗网络在图像补全中的应用逻辑。1. GAN 修复破损图片的 Python 源码先知道它解决什么再动手GAN 对抗网络修复破损图片这套 Python 源码解决的是图像修复image inpainting问题给定一张有遮挡、划痕或缺块的图让模型把缺失区域补全到肉眼难辨。修复没有唯一标准答案模型本质上是在上下文约束下做生成这正是 GAN 派上用场的地方。对做毕设或课设的学生来说它最大的价值是给了一条能跑通的路U-Net 生成器、PatchGAN 判别器、L1 对抗损失从数据到训练再到出图全部串起来。我拆完后判断是网络结构不复杂真正决定效果的是掩码设计与训练配置。同一套网络掩码合理、训练顺序讲究100 epoch 能出干净结果乱调权重500 epoch 依然模糊。适合需要用完整 GAN 源码支撑毕设课设、以及想在自有数据集上跑图像修复的开发者。2. 掩码与训练数据破损图靠“合成”掩码策略直接决定修复上限图像修复训练和分类、检测最大的区别在于你必须同时有“破损图”和“完整图”两张成对图片。现实里找不到大量同一张图的新旧对照所以训练数据必须自己合成。合成手段就是掩码mask一张和图像同尺寸的黑白蒙版白色区域代表待修复位置。拿到完整图后把白色区域的像素抹掉剩下的就是破损图。掩码长什么样、覆盖多大面积直接决定了模型学会的修复能力上限。2.1 训练样本为什么要人为制造合成破损的三种掩码策略真实世界的破损形态千奇百怪旧照片折痕、遮挡物、水印、涂鸦。图像修复没法像超分那样用固定下采样模拟退化所以常规做法是用掩码生成合成破损让网络在“已知破损形状”上学修复。三种掩码策略最常见掩码类型生成方式模拟场景训练难度中心掩码固定遮住图像中心一块矩形大面积内容缺失中随机矩形掩码随机位置、随机宽高遮挡物、贴纸低条纹 / 划痕掩码细长条可旋转角度老照片折痕、划损中毕设级源码一般先用中心掩码 随机矩形把流程跑通。真实应用里的破损大多是不规则形状free-form 掩码最接近现实但生成逻辑复杂、训练难度大属于进阶改动。要有一个预期规则掩码训练出来的模型对矩形缺损修复效果好碰到细小划痕或者不规则擦除效果会明显变差。这不是源码缺陷是训练分布决定的。掩码区域置 0 是最常见的做法源码里一般写成masked img * (1 - mask)。也有的实现会往掩码区域填随机噪声目的是让生成器一开始就意识到“这里没信息、别依赖输入内容”。两种方案都能跑通我只提醒一点如果填噪声噪声的分布要和数据分布大致同量级否则模型会把噪声当真实纹理学进去。我建议初期用置 0 方案简单可控。2.2 预处理参数尺寸、归一化、数据增强怎么设修复任务的预处理比分类任务敏感得多。训练尺寸选 256×256 是效果和显存的折中点128 太小修复细节糊512 显存压力陡增且训练时间翻几倍。归一化统一到 [-1, 1]因为生成器最后一层是 tanh输出范围必须和输入保持一致。数据增强方面随机水平翻转和颜色抖动亮度、饱和度微调是常用组合。这里有一个容易翻车的地方颜色抖动如果作用在“破损图”上而掩码区域当时是 0 值像素颜色抖动会把它变成非 0 噪声等于掩码信息被污染。我一般把颜色抖动限定在完整图上或者先做增强再应用掩码。参数推荐值说明训练尺寸256×256效果与显存折中可降到 192 缓解 OOM归一化[-1, 1]与生成器 tanh 输出对齐掩码类型随机矩形 中心区域混合覆盖常见缺损训练稳定掩码面积占比10% ~ 40%超过 50% 修复基本靠“幻觉”效果断崖增强水平翻转、亮度/饱和度抖动颜色抖动建议只作用于完整图2.3 数据模块实现Dataset、掩码生成、通道拼接细节源码里的数据加载器一般长这样import os, random import numpy as np from PIL import Image from torch.utils.data import Dataset class InpaintDataset(Dataset): def __init__(self, img_dir, img_size256, max_mask_ratio0.4): self.paths sorted(glob.glob(os.path.join(img_dir, *.jpg))) self.img_size img_size self.max_mask_ratio max_mask_ratio def __getitem__(self, idx): img Image.open(self.paths[idx]).convert(RGB) img img.resize((self.img_size, self.img_size)) img np.array(img, dtypenp.float32) / 127.5 - 1.0 # [-1,1] mask np.zeros((self.img_size, self.img_size), dtypenp.float32) h random.randint(self.img_size // 8, int(self.img_size * self.max_mask_ratio)) w random.randint(self.img_size // 8, int(self.img_size * self.max_mask_ratio)) y0 random.randint(0, self.img_size - h) x0 random.randint(0, self.img_size - w) mask[y0:y0h, x0:x0w] 1.0 mask torch.from_numpy(mask).unsqueeze(0) # [1, H, W] img torch.from_numpy(img.transpose(2, 0, 1)) # [3, H, W] masked img * (1.0 - mask) # 掩码区域置 0 return img, masked, mask # 完整图, 破损图, mask逻辑说明核心是“生成 mask → 应用 mask”。掩码宽高在两个随机区间里取值下限是img_size // 8保证最小遮挡不是一个点上限用max_mask_ratio控制防止遮挡面积失控。返回值里img是完整图Ground Truthmasked是喂给生成器的破损图mask会单独作为通道与破损图拼接。参数说明max_mask_ratio是修复任务最重要的旋钮。调大训练样本更难模型对严重缺损的泛化变强但训练不稳定调小模型学得轻松测试时遇到大缺损就露馅。我一般从 0.3 起步训练稳定后再加到 0.4。关于通道拼接生成器输入不是 3 通道而是 4 通道——破损图 3 通道加 mask 1 通道。只喂破损图的话模型不知道某个黑色区域是“原本就黑”还是“被抠掉”mask 通道是唯一的硬性先验。这里可以加一个可视化检查习惯def visualize_batch(gt, masked, mask, save_path): gt (gt.numpy().transpose(1, 2, 0) 1) / 2 masked (masked.numpy().transpose(1, 2, 0) 1) / 2 mask np.repeat(mask.numpy().transpose(1, 2, 0), 3, axis2) vis np.hstack([gt, masked, mask]) * 255 cv2.imwrite(save_path, cv2.cvtColor(vis.astype(np.uint8), cv2.COLOR_RGB2BGR))把完整图、破损图、掩码拼在一张图上导出每次改完数据管线先跑一遍再开训练。很多训练翻车其实在数据阶段就已经埋下了只是你没看数据长什么样。3. 生成器与判别器选型U-Net PatchGAN 为什么是稳定答案网络结构的选择直接决定修复质量上限。图像修复这个任务既要求生成结果在全局语义上说得通又要求局部纹理能骗过人眼。生成器用 U-Net、判别器用 PatchGAN是近几年开源项目里最稳定的组合不是唯一选择但它是新手最容易跑出效果的一套。3.1 生成器为什么选 U-Netskip connection 保细节普通 Encoder-Decoder 会把整张图逐级下采样压成特征向量再解码还原。这个过程里,浅层的位置、边缘、颜色信息大量丢失解码器只能凭高层语义“大概重建”。修复任务恰恰对局部细节极度敏感裂纹、纹理方向、颜色渐变都要保留所以必须让浅层信息绕开压缩瓶颈直接到达输出端。U-Net 的 skip connection 做的就是这件事编码器每一层输出经过跨层连接拼接给解码器对应层边缘纹理信息一路保送到最后的生成结果。规模上主流做法是 4 次下采样通道按 64 → 128 → 256 → 512 递增解码器对称回升。一个容易写错的地方是解码器每个上采样层输入通道必须翻倍因为要和编码器对应层做通道维拼接concat而不是相加。3.2 瓶颈层膨胀卷积修复需要“看到”远处的像素中心掩码训练时模型要从掩码周边推断中心内容。如果生成器感受野不够离掩码远的信息进不来生成区域只能靠近邻平滑外推结果就是一大块模糊。常见做法是在瓶颈层串联 4 个膨胀卷积dilation rate 分别取 2、4、8、16感受野指数级扩大。这是 GLCIC 那类经典工作的标准设计源码里基本都带。改造时注意膨胀卷积不改变特征图尺寸所以瓶颈层输入输出通道保持一致即可。3.3 判别器为什么用 PatchGAN稠密梯度比单一分数训练更稳普通判别器输出一个标量代表整张图真/假。生成器拿到的反馈只有“整张图不够真”但哪里假、哪个局部纹理不对信息量太少。PatchGAN 把判别器输出改成 N×N 矩阵每个元素负责判别原图上对应感受野比如 70×70 patch的真假生成器能收到稠密的、位置化的梯度信号。修复恰恰是局部纹理敏感的PatchGAN 对边缘和纹理的约束远强于普通判别器。实现上判别器用 Instance Norm 而不是 Batch Norm。原因很实际修复训练 batch 内每张图的掩码位置不同BN 的统计量会被 batch 里其他样本带偏Instance Norm 按单张图归一化更符合修复场景。3.4 核心结构定义与参数量参考class Generator(nn.Module): def __init__(self): super().__init__() self.enc1 conv_block(4, 64) # 输入 破损图3通道 mask1通道 self.enc2 down_block(64, 128) self.enc3 down_block(128, 256) self.enc4 down_block(256, 512) self.bottle dilated_stack(512) # dilation rate 2,4,8,16 self.dec4 up_block(512 512, 256) # concat enc4 - 1024 通道 self.dec3 up_block(256 256, 128) self.dec2 up_block(128 128, 64) self.dec1 up_block(64 64, 32) self.out nn.Sequential(nn.Conv2d(32, 3, 3, 1, 1), nn.Tanh()) def forward(self, masked, mask): x torch.cat([masked, mask], dim1) e1 self.enc1(x) e2 self.enc2(e1) e3 self.enc3(e2) e4 self.enc4(e3) b self.bottle(e4) d4 self.dec4(torch.cat([b, e4], dim1)) d3 self.dec3(torch.cat([d4, e3], dim1)) d2 self.dec2(torch.cat([d3, e2], dim1)) d1 self.dec1(torch.cat([d2, e1], dim1)) return self.out(d1)代码逻辑第一层接收 4 通道输入编码器逐级下采样瓶颈层膨胀卷积撑开感受野解码器上采样过程中每一级和对应编码器输出 concat所以输入通道是“上一级输出 编码器输出”。最后 tanh 把输出压回 [-1,1]。参数量参考模块结构参数量级生成器U-Net 4 层膨胀卷积约 35M ~ 45M判别器PatchGAN4 层卷积约 2M ~ 4M生成器比判别器大一个量级是正常的GAN 训练里判别器本来就不需要太大。如果判别器做得比生成器还深它学得太快生成器梯度消失后面会讲到这个坑。4. 训练循环与损失权重让生成器不乱编、判别器不抢戏的关键参数修复任务的训练难点不在网络定义而在三个损失的权重分配和两个网络交替更新的节奏。权重调不对生成器要么偷懒模糊输出要么完全放飞乱编纹理。这一章把损失函数、训练循环、优化器参数一次性讲清楚。4.1 损失组合L1 管轮廓、对抗管纹理、感知管语义三个损失按重要性排序L1 重构损失 对抗损失 感知损失。L1 是像素级逐点差异的绝对值之和保证修复结果和 Ground Truth 在整体结构上接近。修复任务里 L1 通常优于 L2L2 对大误差惩罚过重会把输出往均值方向推结果是模糊L1 对边界更宽容生成结果更锐利。对抗损失BCE让判别器区分“真实完整图”和“生成修复图”。单独用 L1 的模型输出发糊像开了磨皮单独用对抗损失的模型纹理看着很真但内容可能完全跑偏——把 GT 里的物体换成合理但错误的场景。两者必须同时存在。感知损失是把生成图和 GT 都过一遍 VGG16 前几层在特征空间算 L1约束高层语义一致。毕设源码不一定带不带就先跑 L1 对抗效果已经够用。损失计算对象推荐权重L1 重构掩码区域像素为主10对抗 BCEPatchGAN 输出1感知可选VGG16 relu1 ~ relu4 特征0.14.2 掩码区域加权 L1别让完整区域稀释监督信号训练输入里大部分区域是完好的模型只要把完好区域复制到输出loss 就已经很小掩码内部的修复误差被平均掉了。常见做法是给 L1 加一个像素级权重掩码内部权重 20 或更高外部保持 1。这样网络的注意力会被强制集中在待修复区域而不是偷懒抄完好区域。weight torch.ones_like(mask) 20.0 * mask # 掩码内权重 21, 外部 1 g_l1 (torch.abs(fake - gt) * weight).mean()逻辑说明mask是 0/1 张量乘 20 加 1 之后掩码区域每个像素的 L1 权重变成 21完整区域是 1。整个损失的平均值被掩码区域主导模型必须把修复区域练好才能让 loss 降下去。我试过不加权重直接全图 L1训练 60 epoch 后掩码边缘依然有明显的“补丁感”加权重后同样 epoch 数明显改善。4.3 训练循环先更新 D 再更新 G两个网络不能共用学习率GAN 训练每步的更新顺序是固定的先用当前生成器产出的假图更新判别器让 D 更会分辨真假再固定 D 更新生成器让 G 骗过 D。生成器更新时梯度只回流到 GD 的参数不能动。# gt 完整图, masked 破损图, mask 掩码 d_optim.zero_grad() pred_real D(gt) pred_fake D(G(masked, mask).detach()) # 假图梯度隔离 d_loss BCE(pred_real, ones) BCE(pred_fake, zeros) d_loss.backward() d_optim.step() g_optim.zero_grad() fake G(masked, mask) pred_fake D(fake) # 重新前向, 保留梯度 g_adv BCE(pred_fake, ones) # 生成器希望假图被判真 g_l1 torch.abs(fake - gt).mean() g_loss g_adv 10.0 * g_l1 g_loss.backward() g_optim.step()逻辑说明D 更新时传给它的假图必须 detach否则梯度会穿过 D 回流到 G导致 D 的这一步更新间接改了 G 的参数破坏对抗关系。G 更新时重新做一次前向让 D 对假图的打分携带梯度这个梯度只更新 G。显存不够时可以复用 D 更新那步的假图输出传入 D省一次生成器前向代价是 G 梯度在 D 更新阶段拿不到实际影响很小。参数说明两个优化器都建议用 Adam但学习率不能盲目相同。常见做法是 D 的 lr 设为 2e-4G 的 lr 设为 1e-4。更重要的是beta1要从 PyTorch 默认的 0.9 改成 0.5——默认动量累积历史梯度太长GAN 训练里容易震荡0.5 让动量“短视”收敛更稳。判别器增强可以通过梯度惩罚或在损失里加 D 的正则项但最常见、最不过度的是标签平滑真实标签从 1.0 改成 0.9。4.4 预训练与断点续训稳定性的最大保障修复 GAN 比普通 GAN 容易训练因为它有像素级监督信号。但如果一上来就让 D 参与,依然可能瞬间失衡。推荐两阶段顺序第一阶段只用 L1 损失训练生成器 30~50 epoch让 G 先把大面积缺失补出合理轮廓这个阶段 loss 稳定可预测第二阶段加入对抗损失让 D 从零学习什么是真实纹理。这个做法能规避最经典的翻车——开局 D 碾压 GG 梯度消失后续怎么调都救不回来。断点续训是必须的。源码里一般保存 G、D、优化器和当前 epoch加载续训时只恢复权重不够优化器的动量状态不对会影响后续收敛节奏torch.save({ epoch: epoch, G: G.state_dict(), D: D.state_dict(), G_optim: g_optim.state_dict(), D_optim: d_optim.state_dict(), best_l1: best_l1 }, ckpt/epoch{}.pth)加载时把optimizer.state_dict()也一并 load学习率调度从头接上。完整训练流程约 100 epoch 可以下线条件好可以跑到 150。5. 训练避坑与常见问题五个让修复效果翻车的典型场景5.1 判别器太强生成器梯度消失现象训练日志里 D loss 稳步下降并趋近 0G loss 不降甚至上涨生成图像几十个 epoch 几乎没变化。原因D 收敛速度远超 GG 生成的假图被 D 一眼识破BCE 给出的梯度极小G 学不到东西。解决三个手段按优先级试。先给真实标签做平滑1.0 改成 0.9再把 D 学习率降到 G 的一半最后把 D 的更新频率改成每 2 步更新一次。我实测标签平滑在修复任务里见效最快改一行代码就能明显缓解失衡。5.2 生成图出现棋盘格伪影现象修复区域有规律性亮暗格子尤其在边缘平滑区域明显。原因生成器解码器用了转置卷积ConvTranspose2d上采样时卷积核重叠产生不均匀响应这是转置卷积的经典问题。解决把转置卷积全部换成“上采样 普通卷积”的组合self.up nn.Sequential( nn.Upsample(scale_factor2, modebilinear, align_cornersFalse), nn.Conv2d(ch_in, ch_out, 3, 1, 1) )逻辑说明先用双线性插值把特征图放大两倍再做普通 3×3 卷积融合。没有转置卷积的重叠问题棋盘格直接消失代价是参数量略增。5.3 掩码边缘发灰、颜色断层现象修复区域和原图之间有明显接缝像贴了一块灰色补丁颜色过渡生硬。原因两个来源一个是 L1 权重太低模型在对抗损失压力下选择“平均色”敷衍边缘另一个是训练时掩码是硬边界增强管线里的颜色抖动或几何变换污染了掩码坐标导致模型学到的边缘永远是错位的。解决L1 权重从 1 提到 10 以上并给掩码区域加权同时在训练前把增强后的数据可视化一遍确认掩码边缘没有偏移。检查步骤我固定做用 2.3 那个visualize_batch函数把增强后的三张图拼出来用肉眼看一次。5.4 256×256 训练 OOM现象batch size 设置成 8程序启动后几秒就报 CUDA out of memory。原因生成器参数量大且训练循环里 G 前向计算了两次显存峰值高。解决batch 降到 2~4更新 D 时缓存 G 的 output 供 G 更新复用省一次前向还不行就把训练尺寸降到 192 或 128。128 训练的模型再上 256 测试会有精度损失但先跑通流程再升分辨率是性价比最高的路径。5.5 加载预训练权重报 missing key / size mismatch现象torch.load后model.load_state_dict报 unexpected keys 或 size mismatch训练没法续上。原因网络定义改过最常见是输入通道数对不上——之前用 3 通道没拼 mask训练后来改成 4 通道旧权重自然不匹配。也有的是同样的 checkpoint 被某个版本保存时带了额外 key。解决加载时用strictFalse先看哪些层对齐、哪些没对齐同时打印 checkpoint 的键名逐个核对。更省事的做法模型定义文件里加注释写死当前输入通道数和各层输出通道避免改结构后凭记忆猜配置。6. 推理验证与后续改进从出图到 PSNR/SSIM 怎么算再谈三处升级6.1 单图推理把训练好的模型接到真实破损图上训练完成后推理脚本比训练循环简单得多但输入拼接方式必须和训练完全一致否则出图直接不对。def infer(model, img_path, mask_path, out_path, img_size256): img cv2.imread(img_path) img cv2.cvtColor(img, cv2.COLOR_BGR2RGB) img cv2.resize(img, (img_size, img_size)).astype(np.float32) / 127.5 - 1.0 mask cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE) / 255.0 mask cv2.resize(mask, (img_size, img_size)).astype(np.float32) img_t torch.from_numpy(img.transpose(2, 0, 1)).unsqueeze(0) mask_t torch.from_numpy(mask).unsqueeze(0).unsqueeze(0) masked img_t * (1.0 - mask_t) inp torch.cat([masked, mask_t], dim1) # 4通道, 与训练一致 model.eval() with torch.no_grad(): out model(inp).squeeze(0).numpy().transpose(1, 2, 0) # [-1,1] img_np img_t.squeeze(0).numpy().transpose(1, 2, 0) # [-1,1] mask_np mask_t.squeeze(0).numpy()[..., None] # [H,W,1] composite out * mask_np img_np * (1 - mask_np) # 修复区域替换 composite (composite 1) / 2 * 255 cv2.imwrite(out_path, cv2.cvtColor(composite.astype(np.uint8), cv2.COLOR_RGB2BGR))关键点推理时掩码是你手动指定要修复的区域输出后把生成区域和原图非修复区域合成为一张图。硬边界拼接会有接缝要求高可以做一次泊松融合毕设展示阶段直接硬拼也够用。6.2 指标怎么算什么区间算合格有 GT 的测试集上用 skimage 一行计算from skimage.metrics import peak_signal_noise_ratio, structural_similarity psnr peak_signal_noise_ratio(gt, pred, data_range1.0) ssim structural_similarity(gt, pred, channel_axis2, data_range1.0)指标合格线说明PSNR 28 dB30 算优秀低于 25 意味着明显缺陷SSIM 0.900.95 人眼基本不可辨LPIPS 0.1感知相似度需要额外计算可选只看指标也不行。图像修复最终是给人看的训练 L1 不高不代表感知效果好每次评估同时扫几张生成图。6.3 时间充裕时最值得改的三处规则掩码换成 free-form 掩码更贴近真实划痕擦除生成器瓶颈的膨胀卷积换成 gated conv对不规则掩码更友好L1 换成感知损失或加入 SSIM 损失提升纹理自然度。三处改动独立不影响原有训练循环适合作为毕设的“创新点”方向。我在这套源码上翻过最大的车就是掩码边缘颜色断层查了一个通宵最后发现是增强管线里颜色抖动污染了掩码像素。从那以后我每跑一个新的图像修复项目都会强制走一遍“掩码可视化 增强后可视化”校验确认数据没问题再开训练。希望帮到你。本文还有配套的精品资源点击获取

关于本文作者

来自尧图内容编辑团队

尧图内容编辑团队 内容团队

尧图内容编辑团队

本文由尧图网络内容编辑团队执笔。团队由资深项目经理、前端工程师与设计师组成,所有内容均来自亲手交付的真实项目,先讲清问题、再给出可落地的解法。尧图深耕北京网站建设十年,服务过京华建材集团、智造科技等各行业客户,把一线经验沉淀为可复用的行业观察。

  • 十年建站经验,覆盖建材、制造、服务、文创等
  • 项目经理把关选题与事实准确性
  • 工程师与设计师联合撰写专业细节
  • 统一编辑规范,保证文风与排版一致
  • 每月复盘转化数据,迭代选题方向

延伸阅读

相关资讯与近期热门内容

深度阅读推荐

建站决策前值得细读的三篇

网站改版的5个关键决策
2024-08-12

网站改版的5个关键决策

什么时候该改版、改到什么程度、如何避免流量掉光,京华建材集团改版复盘给出答案。

获取专属建站方案

看完文章,把您的行业与预算告诉我们,免费获取一份量身定制的官网建设方案与报价。

立即免费咨询