ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

基于GAN与注意力机制的图像修复系统Python实现指南

基于GAN与注意力机制的图像修复系统Python实现指南 简介这是一套面向计算机相关专业毕业设计、课程设计及机器学习入门者的深度学习图像修复项目资料基于卷积神经网络与对抗式训练策略实现划痕修复、噪点消除和局部遮挡还原等图像缺失区域智能补全功能。资源包共18个文件约5.61MB以png与jpg示例图片、py算法源码、zbak备份文件、md说明文档及gitignore配置为主另含压缩包与项目说明目录结构清晰便于按模块查阅。项目经导师指导并获认可评定成绩达99分代码架构完整且经过验证初学者也能完成环境配置与运行。文档涵盖环境部署、模型训练到结果复现的逐步指南并附示例数据集与测试案例可帮助读者理解图像修复流程、掌握网络搭建与训练排错思路。目前已有84人学习适合需要完整案例锻炼实践能力的学习者参考。1. 图像修复系统到底在修什么从划痕老照片到深度学习补全很多人第一次听到「图像修复系统」脑子里浮现的是把模糊照片变清晰。实际上这是两件事超分辨率做的是放大图像修复Image Inpainting做的是填补——把图像里缺失、被遮挡、被划掉的那块区域用算法生成合理内容补回去。老照片的折痕、扫描件的污渍、截图里被水印盖住的区域、甚至人脸被口罩挡住的部分都是图像修复的典型场景。它和超分的核心区别在于超分输入是完整但低清的图修复输入是残缺的图模型必须「无中生有」地编出缺失像素。这个方向这几年从传统方法基于扩散、基于块匹配快速转向深度学习尤其是 GAN 和注意力机制成熟之后修复质量有了质变。标题里的「Python 实现」意味着整套东西可以用 Python 生态跑通PyTorch 做模型、OpenCV 做前后处理、NumPy 做掩码运算。适合谁有一定 Python 基础、想做一个能写进简历或落地到业务里的深度学习实战项目的人。下面我按「原理选型 → 数据与掩码 → 模型搭建 → 训练调参 → 避坑 → 进阶验证」的顺序把这条链路讲透。2. 图像修复的技术路线选型为什么我最终押注 GAN 注意力2.1 传统方法为什么在真实场景里翻车传统修复方法主要有两类。第一类是基于扩散的比如把缺失区域周围的像素往内推用热传导方程迭代填充。这类方法在窄划痕上效果还行但一旦缺失区域超过几十像素宽补出来的就是一片糊纹理完全丢失。第二类是基于块匹配的比如 PatchMatch从图像其他位置找相似块贴过来。它在纹理重复的场景砖墙、草地表现不错但遇到人脸、文字这种结构性强的内容就崩了——因为找不到语义上匹配的块。我早期用 OpenCV 的cv2.inpaint做过一个文档去水印的小工具TELEA 算法处理细线水印没问题但水印一旦加粗到覆盖文字笔画补出来的字就是鬼画符。这就是传统方法的边界它们只利用了低层像素统计没有语义理解能力。真实业务里的缺失区域往往又大又结构化传统方法基本没有落地价值。2.2 GAN 与注意力机制修复质量的分水岭深度学习修复的分水岭是 2018 年前后。核心思路变成用一个编码器把残缺图压成特征中间用生成器「脑补」缺失区域的特征再解码回图像同时用一个判别器判断补出来的图真不真。这就是 GAN 的对抗训练框架。生成器负责骗过判别器判别器负责识破两者博弈最终生成器学会生成语义合理的內容。但普通卷积有个致命问题卷积核的感受野是局部的当缺失区域很大时生成器看不到远处的有效信息。注意力机制解决了这个问题——它让缺失区域的每个像素都能「查询」全图其他位置的特征找到最相关的参考。比如补一只眼睛注意力可以从另一只完好的眼睛那里借信息。这就是为什么现在主流的修复模型如基于 contextual attention 的方案几乎都带注意力模块。选型结论如果你要做的是能处理大区域、结构化内容的修复系统GAN 注意力是当前性价比最高的路线。纯 CNN 回归L1/L2 loss补出来的图会偏模糊因为回归损失倾向于输出平均值纯 Transformer 效果好但训练成本高对个人项目不友好。GAN 的对抗损失能逼出高频纹理注意力能保证结构合理两者结合是甜点区。2.3 一个最小可跑的选型对照表方法适用缺失尺寸纹理质量结构合理性训练成本推荐场景扩散法 (cv2.inpaint) 10px差差无细划痕、快速预处理PatchMatch 50px中差无重复纹理背景纯 CNN 回归 64px中中低入门练手GAN 注意力 128px高高中人脸、通用场景Transformer 类任意高高高有算力预算时这张表是我自己踩过一圈后的经验值不是绝对标准。缺失尺寸指的是掩码区域的最大边长占图像短边的比例换算。实际选型还要看你的数据量和显卡。个人项目 8GB 显存GAN 注意力在 256×256 分辨率下是能训起来的。3. 数据与掩码修复系统的成败八成在这里3.1 掩码生成别用随机矩形糊弄自己很多人做修复项目掩码直接用随机矩形训练时 loss 降得很漂亮一上真实数据就废。原因是随机矩形和真实缺失的分布差太远。真实场景的缺失有几种典型形态不规则划痕、文字遮挡、块状污渍、以及自由笔刷涂抹。你的掩码生成必须模拟这些形态。我一般用两种掩码策略混合。第一种是规则掩码随机生成线条、矩形、椭圆模拟划痕和水印。第二种是不规则掩码用随机游走生成连通区域模拟污渍和涂抹。下面是一个不规则掩码生成的最小实现import numpy as np import cv2 def random_walk_mask(h, w, num_strokes8, max_vertex20): 用随机游走生成不规则掩码模拟真实涂抹/污渍 mask np.zeros((h, w), dtypenp.uint8) for _ in range(num_strokes): # 随机起点 x, y np.random.randint(0, w), np.random.randint(0, h) points [(x, y)] for _ in range(np.random.randint(5, max_vertex)): # 每步随机方向步长控制笔画粗细感 angle np.random.uniform(0, 2 * np.pi) step np.random.randint(10, 40) x int(np.clip(x step * np.cos(angle), 0, w - 1)) y int(np.clip(y step * np.sin(angle), 0, h - 1)) points.append((x, y)) # 用粗线连接线宽模拟涂抹宽度 for i in range(len(points) - 1): thickness np.random.randint(8, 25) cv2.line(mask, points[i], points[i1], 255, thickness) return mask # 参数说明 # num_strokes 控制笔画数量越多掩码越密建议 5~12 # max_vertex 控制单条笔画拐点上限越大越蜿蜒 # thickness 范围决定缺失区域宽度8~25 对应中等缺失这段代码的逻辑是每条笔画从一个随机点出发按随机角度和步长走若干步形成一条折线再用粗线画出来。这样生成的掩码边缘不规则、有连通性比随机矩形更接近真实。参数上num_strokes别设太大否则一张图被盖掉一半训练难度陡增thickness上限别超过图像短边的 1/8否则小分辨率下掩码会糊成一片。3.2 数据集准备与配对逻辑修复任务的数据是「完好图 掩码 → 残缺图」的配对。训练时你只需要完好的图掩码在线生成残缺图 完好图 × (1 - 掩码)。这样数据增强空间很大同一张图配不同掩码就是不同样本。常用数据集有 CelebA人脸、Places2场景、以及自己业务里收集的图。我一般会写一个 Dataset 类把「读图 → 随机裁剪 → 随机掩码 → 归一化」串起来。注意归一化要和模型输入匹配GAN 常用 [-1, 1] 而不是 [0, 1]因为生成器最后一层常用 tanh。这个细节不注意训练初期 loss 会异常大。import torch from torch.utils.data import Dataset import cv2 import numpy as np class InpaintDataset(Dataset): def __init__(self, img_paths, img_size256): self.img_paths img_paths self.img_size img_size def __len__(self): return len(self.img_paths) def __getitem__(self, idx): img cv2.imread(self.img_paths[idx]) img cv2.cvtColor(img, cv2.COLOR_BGR2RGB) img cv2.resize(img, (self.img_size, self.img_size)) # 归一化到 [-1, 1]匹配 tanh 输出 img img.astype(np.float32) / 127.5 - 1.0 img torch.from_numpy(img).permute(2, 0, 1) mask random_walk_mask(self.img_size, self.img_size) mask (mask 0).astype(np.float32) mask torch.from_numpy(mask).unsqueeze(0) # 残缺图掩码区域置 0 masked_img img * (1 - mask) return masked_img, mask, img # 输入、掩码、真值这里返回三个东西masked_img是模型输入mask用于计算只在缺失区域的 lossimg是监督真值。注意 loss 一定要只在掩码区域算否则模型会偷懒——因为完好区域本来就对整体 loss 会被拉低模型学不到补全能力。这是新手最常见的翻车点之一。4. 模型搭建生成器、判别器与损失函数怎么配4.1 生成器结构编码器-注意力-解码器生成器我一般用「下采样编码 → 中间注意力 → 上采样解码」的结构。编码器用几层 stride2 的卷积把 256×256 压到 64×64 甚至 32×32中间放注意力模块解码器用转置卷积或上采样卷积恢复分辨率。最后接一个 tanh 输出 [-1, 1]。注意力模块的核心是把特征图分成「已知区域」和「缺失区域」缺失区域的每个位置作为 query已知区域作为 key 和 value做一次注意力加权。这样缺失区域就能从已知区域借信息。下面是一个简化版上下文注意力的实现import torch import torch.nn as nn import torch.nn.functional as F class ContextualAttention(nn.Module): def __init__(self, channels): super().__init__() self.query nn.Conv2d(channels, channels // 8, 1) self.key nn.Conv2d(channels, channels // 8, 1) self.value nn.Conv2d(channels, channels, 1) def forward(self, x, mask): # x: [B, C, H, W], mask: [B, 1, H, W]1 表示缺失 B, C, H, W x.shape q self.query(x).view(B, -1, H * W).permute(0, 2, 1) # [B, HW, C/8] k self.key(x).view(B, -1, H * W) # [B, C/8, HW] v self.value(x).view(B, -1, H * W) # [B, C, HW] attn torch.bmm(q, k) / (C ** 0.5) # [B, HW, HW] attn F.softmax(attn, dim-1) out torch.bmm(v, attn.permute(0, 2, 1)) # [B, C, HW] out out.view(B, C, H, W) # 只在缺失区域用注意力结果已知区域保留原特征 return x * (1 - mask) out * mask参数上channels // 8是注意力头的压缩比经验值 8 比较稳压太狠信息丢失压太少显存吃紧。mask的语义要统一1 表示缺失。如果搞反了注意力会作用在已知区域训练直接不收敛。这个模块可以插在编码器和解码器之间也可以插在解码器的每一层。4.2 判别器与损失函数对抗损失 重构损失 感知损失判别器相对简单就是一个二分类 CNN输入图像输出真/假概率。但修复任务的判别器有个讲究最好用 PatchGAN即输出一个 N×N 的 patch 级真假图而不是单个标量。因为修复关注局部纹理patch 级判别能给出更细的梯度。损失函数是修复质量的关键。我一般用三项加权重构损失L1保证补出来的像素和真值接近权重最大通常 1.0。对抗损失逼出高频纹理权重 0.1 左右。感知损失perceptual loss用预训练 VGG 提取特征算 L1保证语义合理权重 0.05~0.1。import torch.nn as nn class InpaintLoss(nn.Module): def __init__(self, vgg): super().__init__() self.l1 nn.L1Loss() self.bce nn.BCEWithLogitsLoss() self.vgg vgg # 预训练 VGG冻结参数 def forward(self, pred, target, mask, disc_real, disc_fake): # 重构损失只在缺失区域算 loss_l1 self.l1(pred * mask, target * mask) / (mask.mean() 1e-8) # 对抗损失 loss_adv self.bce(disc_fake, torch.ones_like(disc_fake)) # 感知损失 feat_pred self.vgg(pred) feat_target self.vgg(target) loss_perc self.l1(feat_pred, feat_target) return loss_l1 0.1 * loss_adv 0.05 * loss_perc注意loss_l1除以了mask.mean()这是为了归一化——否则掩码越小 loss 越小不同样本的 loss 尺度不一致训练会不稳定。这个除法是我踩过坑之后加的不加的话小掩码样本几乎不贡献梯度。4.3 训练循环与关键超参训练循环没什么特别的但有几个超参必须调对。batch size 在 8GB 显存下 256×256 分辨率一般能到 8~16。学习率用 1e-4 起步用 Adam 优化器beta 设 (0.5, 0.999)——GAN 训练里 beta1 用 0.5 比默认 0.9 稳这是 DCGAN 以来的经验。生成器和判别器可以交替训练也可以同时更新我一般同时更新简单。optimizer_G torch.optim.Adam(gen.parameters(), lr1e-4, betas(0.5, 0.999)) optimizer_D torch.optim.Adam(disc.parameters(), lr1e-4, betas(0.5, 0.999)) for epoch in range(num_epochs): for masked, mask, real in dataloader: fake gen(masked, mask) # 判别器更新 pred_real disc(real) pred_fake disc(fake.detach()) loss_D 0.5 * (bce(pred_real, ones) bce(pred_fake, zeros)) optimizer_D.zero_grad(); loss_D.backward(); optimizer_D.step() # 生成器更新 pred_fake disc(fake) loss_G criterion(fake, real, mask, pred_real, pred_fake) optimizer_G.zero_grad(); loss_G.backward(); optimizer_G.step()判别器更新时对fake用了.detach()这是必须的否则梯度会回传到生成器打乱训练节奏。这个细节漏了训练会莫名其妙震荡。5. 避坑与排查那些让我重训三次的问题5.1 生成结果一片灰或模糊现象训练几十轮后补出来的区域是一团灰色或模糊色块没有纹理。原因通常是重构损失权重过大对抗损失没起作用模型退化成求平均。解决把对抗损失权重从 0.1 提到 0.2~0.5或者检查判别器是不是太强导致生成器梯度消失。也可以先单独训几轮重构损失 warm-up再开对抗。5.2 掩码区域出现明显边界接缝现象补全区域和周围已知区域之间有一条明显的线。原因有两个一是 loss 只在掩码内部算边界处没有约束二是生成器没有把已知区域的信息传过来。解决把掩码做一次膨胀dilate让 loss 覆盖边界一圈同时确认注意力模块的 mask 语义正确已知区域确实参与了 key/value。5.3 训练 loss 震荡不收敛现象生成器 loss 忽高忽低判别器 loss 趋近 0。原因判别器太强生成器学不动。解决降低判别器学习率比如生成器 1e-4判别器 2e-5或者给判别器加 label smoothing真标签用 0.9 而不是 1.0。这是 GAN 训练的经典问题不是你的代码写错了。5.4 显存溢出OOM现象训练到一半报 CUDA out of memory。原因注意力模块的注意力矩阵是 HW×HW256×256 输入下中间特征 64×64矩阵是 4096×4096batch 一大就爆。解决把注意力放在更低分辨率32×32的特征上或者用窗口注意力降低复杂度。我一般把注意力插在 32×32 那层显存和效果平衡最好。5.5 推理时结果和训练时不一致现象训练时看着还行推理单张图效果差很多。原因训练时用了 batch normalization推理时 batch size1统计量不匹配。解决推理前调model.eval()或者干脆把 BN 换成 instance norm——修复任务里 instance norm 通常比 BN 更稳因为不依赖 batch 统计。6. 进阶技巧用掩码膨胀和两阶段训练把 PSNR 再拉几个点基础版跑通之后想再提质量我一般上两个技巧。第一个是掩码膨胀训练训练时把掩码随机膨胀几像素让模型学会处理边界不确定的情况推理时反而更鲁棒。第二个是两阶段训练第一阶段只用重构损失训一个粗糙的生成器第二阶段加载这个权重加对抗损失和感知损失精调。这样比从头端到端训稳定得多收敛也快。验证方法上别只看 PSNR 和 SSIM。这两个指标对修复任务有偏差——它们偏好模糊结果因为模糊图和高清图的像素差反而小。我一般再加一个人眼评估固定几张测试图每个 epoch 存一次结果拼成网格图看趋势。下面这个保存网格的脚本我每个项目都会用import torchvision.utils as vutils def save_grid(masked, fake, real, epoch, save_path): # 拼成 [输入 | 生成 | 真值] 三列 grid torch.cat([masked, fake, real], dim0) vutils.save_image(grid, f{save_path}/epoch_{epoch}.png, nrowmasked.size(0), normalizeTrue, value_range(-1, 1))nrow设成 batch size这样每行是一组对比。value_range(-1, 1)对应前面的归一化不设的话图会发白。这个网格图比任何指标都直观哪一轮开始出纹理、哪一轮开始过拟合一眼就能看出来。最后说个习惯我每次改模型结构或损失权重都会先跑一个 500 步的小实验看 loss 曲线和网格图确认方向对了再开完整训练。直接上大训练翻车一次就是几小时后悔药没处买。图像修复这个方向数据掩码的质量比模型结构重要损失函数的配比比网络深度重要——把这两件事做扎实一个中等复杂度的 GAN 就能出能看的结果。希望帮到你。本文还有配套的精品资源点击获取
返回列表