简介:基于 Python 实现的生成对抗网络(GAN)图像修复项目,面向计算机相关专业正在完成课程设计、毕业设计或需要项目实战练习的学习者,也适合对深度生成模型感兴趣的入门者。项目从数据分布建模、生成器与判别器搭建到图像缺失区域补全,覆盖了完整实现流程,难度适中,适合对照源码逐步理解 GAN 的训练机制与图像修复思路。压缩包共 7 个文件,包含 6 个 Python 脚本和 1 个 Markdown 文档;Python 文件分别承担模型结构定义、工具函数封装、训练流程与修复推理等环节,Markdown 文档用于说明整体设计与使用方法,包体仅 12KB,结构紧凑、便于快速查看。资源已在平台获得 164 人次学习下载,源码经过本地编译调试,可正常运行,并配有文档说明,能够帮助使用者快速跑通实验并掌握 GAN 图像修复的基本实现方式。
1. 为什么图像修复要选GAN:不是补洞,是补语义
老照片断痕、监控遮挡、旧电影划痕,这些场景落到技术上都是同一件事:图像修复。早年常用的做法是拿周围像素插值或找相似纹理块粘贴,对付纯色背景还行,一旦缺口落在人脸、车牌这类结构密集的区域,补出来的东西一眼假。GAN 出现后,思路换了——不是把像素算出来,而是把语义生成出来。基于 Python 的 GAN 图像修复模型,核心就是生成器补全缺失区域、判别器监督结果是“像真的”,两段对抗,最终得到自然且结构合理的修复图。这篇写给拿到源码后想改模型、调参数的人,把原理、训练代码和调参经验一次讲透。
2. 生成器、判别器与三种损失:GAN图像修复模型的结构选型
2.1 图像修复问题的本质:补像素还是补语义
动手改代码前,先把问题形式化。图像修复的输入是一张部分缺失图 x_masked 和一张掩码 M,掩码里为 1 的区域表示缺失像素,输出是完整的预测图。如果缺失区是天空、草地这类重复纹理,传统插值算法勉强能应付,它统计邻域像素颜色,平滑地补上。但现实里的破损很少这么友好,更多是物体边缘横穿缺失区,比如合影中一张脸被划掉一半,或者车牌上贴了一条胶带。邻域像素统计给不出任何有意义的线索,必须靠先验知识猜“这里应该有什么”。
GAN 带来的变化,是把“猜”变成一个可训练的过程。生成器 G 根据可见区域预测缺失内容,判别器 D 负责判断 G 的输出与真实图像能否区分。训练收敛后,G 学到的就不只是像素的统计相关性,而是“这类场景下缺失区域在语义上应该长什么样”。这也是为什么基于 Python 实现 GAN 图像修复模型时,网络结构、损失权重、训练节奏的选择,会直接决定你是得到一张边缘自然的修复图,还是得到一片涂抹痕迹明显的水彩。
2.2 生成器选型:U-Net、部分卷积与注意力怎么选
生成器结构,主流有三类,选型先看你的掩码形态。
纯 U-Net:编码器逐层下采样提取语义,解码器通过跳跃连接把多尺度特征拼回来。优点是结构简单、参数适中、训练稳定,适合掩码面积占比不高的场景;缺点是缺失区域很大时,容易产生模糊伪影。
部分卷积(Partial Convolution):卷积只在有效像素上进行,每一步同步更新掩码,把有效区逐步收缩。这类结构对规则掩码,比如矩形涂鸦、字幕遮挡,效果比普通 U-Net 高一个档次,但实现复杂,训练时要额外维护动态掩码。
注意力增强:在解码器特征层加 self-attention,让生成区域能参考远处未见区域的特征。适合遮挡物移除,比如人站在树前这种重复纹理场景。
我的经验路径是:先上纯 U-Net 把基线跑通,确认损失和视觉效果趋势正常,再根据失败案例决定加不加模块。不要一上来就堆 PartialConv 加注意力,调试成本会高到怀疑人生。常见做法是把掩码作为额外通道拼到输入里,让网络从一开始就知道哪里不能看。
2.3 判别器选型:PatchGAN加谱归一化的理由
判别器的作用是逼生成器输出更真实的结果,但它本身不能太强,否则生成器会被碾压。推荐用 PatchGAN:不对整张图只输出一个真假概率,而是输出一个 N×N 特征图,每个位置对输入图像的一个局部区域做判断。PatchGAN 对纹理细节的敏感度比全局判别器高,参数少,训练稳定。
更关键的是加谱归一化。实现层面很简单:
import torch.nn as nn def build_discriminator(in_channels=3): """ PatchGAN判别器:输出 28x28 的真伪图,每个值代表局部区域的真假。 中间层全部做谱归一化,约束判别器权重,稳定GAN训练。 """ return nn.Sequential( nn.utils.spectral_norm(nn.Conv2d(in_channels, 64, 4, stride=2, padding=1)), nn.LeakyReLU(0.2, inplace=True), nn.utils.spectral_norm(nn.Conv2d(64, 128, 4, stride=2, padding=1)), nn.LeakyReLU(0.2, inplace=True), nn.utils.spectral_norm(nn.Conv2d(128, 256, 4, stride=2, padding=1)), nn.LeakyReLU(0.2, inplace=True), nn.Conv2d(256, 1, 4, stride=1, padding=1), )这段代码里有两个要点。一是中间卷积层做谱归一化、输出层不做,既稳定训练又不影响最终判断。二是输出通道为 1、不接 Sigmoid,损失函数用 BCEWithLogitsLoss 在内部处理数值稳定,不需要手动加 Sigmoid。我对比过加不加谱归一化的训练曲线:没加的版本大约 3000 步左右 D loss 显著上升,G loss 振荡,生成图变成彩色噪点;加了之后曲线全程平稳。新手最容易翻车的地方就是判别器太强,谱归一化能挡掉一大半问题。
2.4 三种损失的配比:重建、感知与对抗缺一不可
源码里的 GAN 修复模型通常不只对抗损失,实践里最少叠加三种。L1 重建损失给模型提供强梯度,逼生成器把颜色和边缘位置放对,但它带来的结果是模糊但位置准确的图。感知损失用预训练 VGG16 提取特征,比较特征图差异,让修复区域在高层语义上一致,这是从“模糊”变“清晰”的关键。对抗损失让生成图像符合数据集的整体分布,负责最后一块“像不像”。
三个损失的配比,我用的基准是 L1 给 1.0,感知给 0.1,对抗给 0.01。对抗权重给到 0.1,生成器会被带偏去追求夸张的纹理细节,忽略真实结构;给到 0.001,对抗约束又几乎失效,效果和纯回归差不多。合理范围我建议在 0.005 到 0.05 之间做小范围搜索。
3. 在本地把修复模型跑起来:PyTorch环境、掩码数据与训练主循环
3.1 环境准备:Python版本、VSCode配置与依赖安装
源码落地第一步是环境。Python 版本建议 3.8 以上,Windows 安装时勾选 Add to PATH,Linux 下用系统源装。编辑器用 VSCode 配 Python 插件,或者 PyCharm,都不影响模型本身。国内安装可以配镜像源,不然 torch 全家桶下载会比较煎熬。
依赖文件 requirements.txt 长这样:
torch>=2.0.0 torchvision>=0.15.0 opencv-python>=4.6.0 numpy>=1.24.0 tqdm>=4.65.0 tensorboard>=2.13.0安装命令就是常规的 pip install -r requirements.txt。PyTorch 优先装 CUDA 版,哪怕只有一块 6GB 显存的旧卡,训练 256x256 图像修复也够了;没有 NVIDIA GPU,CPU 版也能跑通流程,只是要把 batch size 和迭代次数调小。
3.2 训练数据构造:不规则掩码生成函数
图像修复训练不需要人工标注,常规图像数据集就行,Places365、CelebA,或者自己收集的图片都可以。每次取一张原图,随机生成掩码,把原图和掩码相乘得到缺图。掩码怎么生成,直接决定模型能力边界。我写了一个不规则掩码函数:
import cv2 import numpy as np def random_mask(batch, channels, height, width, max_ratio=0.4): """ 生成不规则掩码,模拟划痕、污渍遮挡。 返回0-1的mask,mask=1表示缺失区域。 """ mask = np.zeros((batch, channels, height, width), dtype=np.float32) for i in range(batch): brush_count = np.random.randint(3, 8) for _ in range(brush_count): y, x = np.random.randint(0, height), np.random.randint(0, width) angle = np.random.randint(0, 360) length = np.random.randint(height // 6, height // 3) brush_width = np.random.randint(6, 16) cv2.ellipse(mask[i], (x, y), (length, brush_width), angle, 0, 360, 1, -1) return mask参数逻辑不复杂:brush_count 是每张样本上的随机笔刷数,决定掩码复杂度;length 控制缺损尺寸;brush_width 控制划痕粗细。注意 max_ratio 只是个经验参考,训练时缺损面积别超过 40%,否则生成器要凭空想象大半张图,难度陡增、效果也差。
这里有个常见误区要说明:训练数据里矩形掩码太多,模型容易把修复问题简化成“边缘外扩”,一遇到真实的不规则破损就失效。训练用不规则掩码当主力,验证再用规则掩码看边界效果。
3.3 生成器实现:掩码输入通道与U-Net跳跃连接
我在生成器的编码器入口做了一次改动:把掩码作为第 4 个通道拼进去。这样生成器从一开始就知道哪里不能看,哪里必须补。代码简化如下:
import torch import torch.nn as nn class ConvBlock(nn.Module): def __init__(self, in_ch, out_ch): super().__init__() self.block = nn.Sequential( nn.Conv2d(in_ch, out_ch, 3, padding=1), nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True), nn.Conv2d(out_ch, out_ch, 3, padding=1), nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True), ) def forward(self, x): return self.block(x) class UNetGenerator(nn.Module): def __init__(self, in_channels=4, out_channels=3): super().__init__() # 编码器逐级下采样:256 -> 128 -> 64 -> 32 self.enc1 = ConvBlock(in_channels, 64) self.enc2 = ConvBlock(64, 128) self.enc3 = ConvBlock(128, 256) self.pool = nn.MaxPool2d(2) self.bottleneck = ConvBlock(256, 512) # 解码器输入要将编码器的跳跃连接拼进来 self.dec3 = ConvBlock(512 + 256, 256) self.dec2 = ConvBlock(256 + 128, 128) self.dec1 = ConvBlock(128 + 64, 64) self.out = nn.Sequential( nn.Conv2d(64, out_channels, 3, padding=1), nn.Tanh() ) def forward(self, x): e1 = self.enc1(x) e2 = self.enc2(self.pool(e1)) e3 = self.enc3(self.pool(e2)) b = self.bottleneck(self.pool(e3)) # 解码,逐层拼接编码器对应层的特征 d3 = self.dec3(torch.cat([self._upsample(b), e3], dim=1)) d2 = self.dec2(torch.cat([self._upsample(d3), e2], dim=1)) d1 = self.dec1(torch.cat([self._upsample(d2), e1], dim=1)) return self.out(d1) def _upsample(self, x): return nn.functional.interpolate(x, scale_factor=2, mode='nearest')设计点有三个。第一,跳跃连接不可省,没有 skip 连接,编码器里的浅层边缘信息在深层已经丢光,修复区边界会模糊。第二,输出层用 TanH 把输出压到 [-1,1],这对应输入的归一化范围;如果输入是 [0,1],输出层要改成 Sigmoid。第三,上采样用 nearest 插值而不是转置卷积,前者没有可学习参数,不容易出现棋盘伪影。
3.4 训练主循环:判别器、生成器交替更新
训练修复模型的主循环,本质是交替优化。每步先更新判别器,再更新生成器。
import torch import torch.nn.functional as F def train_one_epoch(g, d, loader, opt_g, opt_d): """训练一个epoch,返回平均生成器损失""" total_loss = 0.0 for step, (real_img, _) in enumerate(loader): real_img = real_img.to(device) mask = torch.tensor(random_mask(real_img.size(0), 1, H, W)).to(device) masked_img = real_img * (1 - mask) g_input = torch.cat([masked_img, mask], dim=1) fake_img = g(g_input) # 更新判别器:真实图判真,生成图判假 d.zero_grad() real_pred = d(real_img) fake_pred = d(fake_img.detach()) # detach关键,阻断梯度回传到生成器 d_loss = F.binary_cross_entropy_with_logits(real_pred, torch.ones_like(real_pred)) + \ F.binary_cross_entropy_with_logits(fake_pred, torch.zeros_like(fake_pred)) d_loss.backward() opt_d.step() # 更新生成器:对抗损失 + L1重建 + 感知损失 g.zero_grad() fake_pred = d(fake_img) adv_loss = F.binary_cross_entropy_with_logits(fake_pred, torch.ones_like(fake_pred)) rec_loss = F.l1_loss(fake_img, real_img) g_loss = rec_loss + 0.1 * perceptual_loss(fake_img, real_img) + 0.01 * adv_loss g_loss.backward() opt_g.step() total_loss += g_loss.item() return total_loss / len(loader)关键点有三个:一是更新判别器时 fake_img 必须调 detach(),否则梯度会顺带更新生成器,破坏交替逻辑;二是生成器三个损失合并成标量后一次 backward,不能分开 backward 三次,否则梯度叠加方式不对;三是真实图和掩码要放同一设备,不然报 device mismatch。
训练节奏上,我倾向判别器与生成器同步更新。如果发现 d_loss 长期趋近 0,说明判别器太强,生成器学不到梯度,改成分隔更新——判别器每更新 1 次,生成器更新 2 次。
3.5 结果落盘与项目文档说明怎么补
训练过程每个 epoch 结束后,保存检查点:
# 生成器与判别器权重分开存,方便单独加载推理 torch.save({ 'generator': g.state_dict(), 'discriminator': d.state_dict(), 'optimizer_g': opt_g.state_dict(), 'epoch': epoch, }, f'checkpoints/g_{epoch:04d}.pth')标题里带了“源码+文档说明”,很多人忽略文档价值。一个修复模型项目,文档至少要有四块内容:环境依赖与安装命令、数据集目录结构、训练与评估命令、训练日志与调参记录。README 里放一张已测过的训练配置表,写清输入尺寸、batch size、学习率、损失权重、总步数,并注明在什么数据集上取得过什么 PSNR/SSIM。这样接手源码的人不用花两周重新猜参数,这是“有文档”和“能交接”的分界线。
4. 决定修复效果的6个关键参数:怎么调、为什么这么调
4.1 输入分辨率和batch size:显存、感受野与细节的平衡
输入分辨率直接决定修复细节上限。256x256 是起步配置,它训练速度快、显存占用低,但发丝、纹理这类细节会糊。512x512 效果明显提升,代价是训练时间变长、显存占用翻倍。我的做法是在 256x256 上把结构跑通、损失确认正常,再升到 512。
batch size 方面,256 分辨率下 6GB 显存可以设 8,升到 512 建议降到 2–4。显存不够有个折中:输出分辨率保持 512,编码器输入降到 256,只在解码器最后上采样回 512,细节保留比纯 256 略好。
4.2 损失权重:1.0 : 0.1 : 0.01 的调整方向
三个损失权重是最值得花时间的参数。基准配比 1.0 : 0.1 : 0.01,在这个基础上微调。边缘断痕明显、修复区与周围结构不连贯,调大感知损失到 0.3,让高层语义逼近。颜色断层、边界发暗,调大 L1 权重,但别超过 3.0,太大纹理就平均了。纹理细节丰富但位置不对,把对抗权重从 0.01 提到 0.05,能显著增加细节真实感,但过头会出现彩色噪点。
不建议同时改两个以上权重,否则效果归因不清楚。更靠谱的办法是固定其他参数,只动一个权重,在验证集上算 PSNR/SSIM,用数据决定。
4.3 判别器更新频率:1:1改1:2救命的场景
GAN 训练是博弈,判别器更新太频繁生成器跟不上,更新太少判别器又学不到真假差异。默认 1:1,即每步两个网络都更新。但很常见的情况是 d_loss 很快降到 0 而 g_loss 还在高位,这时候果断改成判别器每更新 1 次、生成器更新 2 次。源码里这个逻辑只要在训练循环里把判别器更新放到 if step % 2 == 0 条件下,成本极低,却常能把濒临崩溃的训练救回来。
4.4 掩码面积占比:训练分布要覆盖目标场景
掩码面积占比是决定泛化性的重要参数。很多源码默认用很小的随机掩码,训练损失降得漂亮,测试时遇到大块遮挡立刻失效。我的习惯是训练阶段掩码面积在 10% 到 40% 之间浮动,每个 batch 混入 20% 左右的大掩码样本。损失曲线会略有抖动,但换来的是真实场景下更稳的修复效果。
记住一个原则:测试场景的掩码面积只会比训练更大,不会更小,训练覆盖范围应当略大于目标场景。
4.5 优化器参数:Adam的betas不是默认值
优化器用 Adam 没问题,但两个参数必须改。betas 建议设成 (0.5, 0.999),GAN 论文里大多数是这个配置,因为 beta1 默认 0.9 会带入动量惯性,容易在博弈中震荡。初始学习率 G 和 D 都设为 2e-4,损失在很长区间不下降时,降到 4e-5 再训练几千步,通常能解锁更低损失平台。
学习率调度上,先用固定学习率跑完 2/3 预算,后面 1/3 用余弦退火或线性衰减,避免步长太大把已经收敛的生成器推离稳定点。
4.6 训练轮数与早停判断:不能只看训练loss
修复模型不能只看训练损失。每若干轮在验证集做一次修复推理,算 PSNR 和 SSIM。保存 checkpoint 时,不要以为最后一个 epoch 一定最好,要用验证指标决定。这里有个常见误会:验证损失一直下降不代表生成质量上升,因为重构损失在压缩纹理细节时会牺牲感知真实度。除了数值指标,每轮固定保存几张可视化对比图,人眼扫一眼,再决定要不要停。这是最土也最可靠的方法。
5. 避坑指南:图像修复模型常见的6个翻车场景
5.1 现象:训练半天,loss振荡,修复图变成色块
原因:最常见是对抗权重过大,生成器只顾骗判别器,放弃结构重建;另一个可能是学习率太高,参数在最优值附近反复横跳。
解决:先把对抗权重从 0.01 减半到 0.005,再把学习率从 2e-4 降到 1e-4,重启训练观察前 2000 步曲线。如果依然振荡,检查 batch size,小于 4 时梯度噪声太大,加大 batch。
5.2 现象:修复区域与周围颜色断层,边界生硬
原因:生成器在边界处没有把可见区域信息融入,感受野够不到边界外侧的纹理;另一种可能是输入归一化范围与输出激活不匹配。
解决:确认输入图像归一化到 [-1,1],生成器输出用 TanH;输入是 [0,1],输出用 Sigmoid。其次在损失函数里加边界权重:把掩码边缘外扩 5 个像素,该区域 L1 损失乘 2 或 3,强制生成器把边界过渡处理好。
5.3 现象:生成器把整张图重画了一遍,可见区域也被改了
原因:这是修复模型很隐蔽的陷阱。生成器发现与其专注缺失区域,不如把整张图统一生成,L1 损失也可能较小,但可见区域细节被改掉,结果不自洽。
解决:推理阶段做像素回填,最后一步执行 output = fake_img * mask + original * (1 - mask),保证可见区域不被改动。训练阶段也在输出前做同样操作,让生成器专心处理缺失区域。
5.4 现象:d_loss无限趋近0,g_loss不降,模型罢工
原因:判别器能力太强,生成器的梯度被完全压制,学不到有效信号。
解决:改成生成器每 2 步更新一次(2:1 比重)。进一步缓解,把判别器换小,减少通道数或下采样层数;往判别器输入加噪声,或者用 label smoothing 让判别器别太自信。
5.5 现象:训练集修复效果很好,验证集一塌糊涂
原因:过拟合。模型记住了训练集特定的掩码形状和图像分布,换到新的形状、光照、内容就失效。
解决:对比训练集和验证集 PSNR,差值超过 4–5dB,就要加数据增强:随机旋转、翻转、颜色扰动,同时增大掩码形状随机范围。还不行就减少模型参数,U-Net 通道基数从 64 减到 48。
5.6 现象:大块区域修复尚可,细小物体修复全是糊的
原因:小物体在深层特征图里连几个像素都占不到,下采样后信息几乎丢失。
解决:一是输入分辨率从 256 升到 384 或 512;二是感知损失的特征层选低一点,别只取 VGG 最后一层,多取 relu2_2、relu3_3 这些中低层特征;三是在解码器浅层加一个辅助 L1 损失,让梯度尽早回传到低层网络。
6. 从源码变成可用模型:评估指标、导出部署与业务改造
6.1 三个指标一次看懂:PSNR、SSIM、FID怎么测
训练结束必须量化验收。PSNR 反映像素误差,对噪声敏感;SSIM 侧重结构相似。两个一起看,一个看整体对错,一个看局部失真。FID 衡量生成分布与真实分布的差异,适合比较生成质量,需要预训练 InceptionV3 特征,用 pytorch-fid 包一行命令出结果。可视化对比时,把原图、掩码图、修复图、真实图四张横向拼在一起,保存成一张大图抽查。我每次调参后先盯这批图,再去看数值,视觉审查比盯着 loss 曲线管用得多。
6.2 推理导出:ONNX与BatchNorm冻结的坑
线上部署可以把 PyTorch 模型导出为 ONNX,用 ONNX Runtime 跑推理。这里有一个我踩过的坑:生成器里有 BatchNorm 层,导出前必须切到 eval 模式并固定 BN 的统计量。如果漏了这步,导出后的模型推理结果在边界上会有颜色跳跃,因为 BN 层还在用当前 batch 的统计量。正确的导出姿势是先把 model.eval() 调到推理模式,再调 torch.onnx.export。
6.3 业务接入:基线先行、逐项放开
想把这套方案接进自己的业务,第一步是收集业务数据构造不规则掩码,跑一轮 256x256 基线,记录 PSNR 和 SSIM。不要上来就改模型结构,先调损失权重把基线稳住,再决定加不加注意力或部分卷积。老照片修复大概率要 512 分辨率加部分卷积;监控遮挡物移除要加大掩码面积,必要时引入视频帧的时间一致性约束。文档里留下基线配置、回合曲线和失败案例,这些资料才是源码快速落到业务的关键。
拿我自己来说,踩过最大的坑是一次性把所有 trick 全用上,结果模型严重过拟合,花了两周才调回基线。现在我的习惯是每次只改一个变量,先立住基线再逐步放开。希望帮到你。
本文还有配套的精品资源,点击获取