简介:基于 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和掩码M,M中值为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_holes=6, max_ratio=0.25): """生成不规则0/1掩码,1表示待修复区域""" h, w = shape mask = np.zeros((h, w), dtype=np.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, size=256): 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_size=3, stride=1, padding=1): super().__init__() self.conv = nn.Conv2d(in_ch, out_ch, kernel_size, stride, padding, bias=True) # 掩码卷积的权重恒为1,只用于统计有效像素个数 weight = torch.ones(in_ch, 1, kernel_size, kernel_size) self.mask_conv = nn.Conv2d(in_ch, 1, kernel_size, stride, padding, bias=False) self.mask_conv.weight = nn.Parameter(weight, requires_grad=False) 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(max=200.0) return out * scale, update_mask.clamp(0.0, 1.0)这段代码的关键逻辑在scale的计算:卷积窗口内有效像素数为update_mask,如果窗口完全被掩码覆盖,update_mask接近0,scale会被clamp限制在一个合理上限,避免输出爆炸。1e-8的加项是防除零。同时,更新后的掩码还要传给下一层,因为随着下采样,掩码会逐步变小,网络需要知道哪些区域仍然是可信任的已知像素。
整个生成器就是一组部分卷积块的下采样和上采样组合:下采样阶段把空间尺寸从256降到32,通道数从16逐步加到256;上采样阶段优先使用最近邻插值或像素重排,避免转置卷积导致的棋盘纹理伪影。
3.2 判别器设计:PatchGAN与掩码输入
判别器的作用是判断给定图像是真实还是修复出来的。常见做法是PatchGAN式的判别器,输出不是单个标量,而是一个N*N的矩阵,每个元素只负责判断图像一个局部区域的真伪。这样做的一个直接好处是:生成器没法只靠整体色调蒙混过关,每个局部块都要足够真实。
判别器输入需要拼接掩码通道,原因是如果不给判别器看掩码位置,它就只能靠边缘痕迹判断真假。拼接掩码后,判别器能聚焦到修复区域,评分更精准。这一点在图像修复模型里和普通图像生成的判别器有明显区别。
判别器网络的输入通道因此是4:RGB三通道加掩码通道。层结构可按风格化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, device='cuda') dummy_mask = torch.rand(1, 1, 256, 256, device='cuda').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_size=256, stride=224): h, w = img.shape[:2] out = np.zeros_like(img, dtype=np.float32) weight = np.zeros((h, w, 1), dtype=np.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下降曲线。这条曲线比单点修复效果更值得存档,因为它能标定出模型的鲁棒边界:曲线在某个面积点突然断崖式下跌,说明模型语义预测能力的阈值就在那里。后续再增加训练数据或调大模型容量时,对比这条曲线就能确认改动是否真的提升了修复能力,而不是只对测试集几张图有效。这个习惯对迭代模型版本非常有用,推荐长期沿用。
本文还有配套的精品资源,点击获取