ARTICLE DETAIL

建站实战干货

来自一线的建站与推广经验沉淀,每一条都经过真实交付验证。

基于深度学习的图像修复实战:从掩码生成到GAN与注意力机制

2026/9/12 0:29:38 拓冰建站 浏览量
基于深度学习的图像修复实战:从掩码生成到GAN与注意力机制 简介一套基于深度学习的图像修复系统Python实现与项目文档面向计算机专业毕业设计、课程设计以及需要完整实战项目的机器学习爱好者。项目采用卷积神经网络与对抗式训练策略可对图像划痕、噪点、局部遮挡等损伤进行智能补全代码结构完整、运行稳定初学编程者也能参照文档完成环境配置与执行。资源包共18个文件包含Python源码简单版与复杂版两个脚本、Markdown项目文档、png/jpg示例图片以及zip备份文件整体体积约5.61MB目录结构清晰便于按模块阅读。其中示例图片覆盖修复前后对比与中间结果有助于直观理解算法效果。已有84人学习浏览适合作为深度学习图像处理方向的课题参考。通过阅读源码与项目文档可以系统掌握图像修复的数据预处理、模型搭建、训练与推理全流程并可直接基于示例数据集动手实验或继续扩展。1. 基于深度学习的图像修复从补像素到理解语义图像修复Image Inpainting不是简单的“PS 去水印”它要解决的是一个 ill-posed 逆问题给定一块缺失或损坏的区域模型需要在没有唯一正确答案的前提下生成与周围环境在纹理、结构、光照甚至语义上一致的像素。传统方法用扩散或纹理合成只能处理窄条划痕一旦缺失区域变大结果就是一片模糊。深度学习的价值在于它把修复任务从“复制周围像素”提升到了“理解图像内容”的层面——网络需要先识别出这是一张人脸、一辆车或一片草地再根据语义先验补出合理的内容。这个标题里的“Python 实现”暗示了两层意思一是用 PyTorch 或 TensorFlow 搭模型并跑通训练推理二是项目需要形成可维护的代码结构包括数据加载、训练脚本、评估脚本和文档说明。适合的读者是已经跑过分类或分割任务、想转向生成式视觉任务的工程师以及需要把算法工程化落地的研究开发者。下文按照数据准备、模型结构、训练调优、推理加速这条路径展开全部基于 PyTorch 展开因为它对掩码操作的灵活性和生态成熟度在修复任务里最顺手。2. 数据与掩码图像修复的输入决定上限2.1 修复任务为什么必须自定义数据管线图像修复的数据集不能直接拿 ImageNet 就用因为模型需要一个“配对”的监督信号输入是带有掩码缺失的图像标签是原始完整图像。这意味着每个训练样本要经历三次处理读取原图、生成掩码、叠加生成输入。掩码的形状、大小、位置直接决定了任务难度如果掩码策略单一模型很快会过拟合到某种固定的缺失模式上。我一般会把数据管线分成三部分原始图像读取与增强、掩码生成策略、输入合成逻辑。其中掩码生成是核心它模拟的是真实场景里的划痕、遮挡、文字覆盖或大块损坏。真实场景里老照片破损往往是多条细线加几块斑块的组合而目标检测后的遮挡则是规则的矩形。掩码策略必须覆盖这两种分布模型才能泛化。2.2 掩码生成的不规则多边形算法OpenCV 的cv2.rectangle和cv2.line只能生成规则掩码实际效果很差。推荐用不规则多边形加随机宽度的曲线来模拟真实损坏。下面这段代码是训练时动态生成掩码的完整实现import cv2 import numpy as np def generate_mask(batch_size, img_size(256, 256), max_vertices10): masks [] for _ in range(batch_size): mask np.zeros((img_size[0], img_size[1], 1), dtypenp.uint8) num_polys np.random.randint(1, 4) for _ in range(num_polys): # 随机多边形模拟大块缺损 num_vertices np.random.randint(4, max_vertices 1) vertices [] cx, cy np.random.randint(0, img_size[0]), np.random.randint(0, img_size[1]) radius np.random.randint(20, 60) for i in range(num_vertices): angle 2 * np.pi * i / num_vertices np.random.uniform(-0.5, 0.5) r radius * np.random.uniform(0.3, 1.5) x int(cx r * np.cos(angle)) y int(cy r * np.sin(angle)) vertices.append([x, y]) cv2.fillPoly(mask, [np.array(vertices)], 255) # 随机曲线模拟划痕 num_curves np.random.randint(0, 3) for _ in range(num_curves): start (np.random.randint(0, img_size[0]), np.random.randint(0, img_size[1])) end (np.random.randint(0, img_size[0]), np.random.randint(0, img_size[1])) thickness np.random.randint(3, 12) cv2.line(mask, start, end, 255, thickness) # 在直线上叠加抖动模拟粗糙划痕 for t in np.linspace(0, 1, 20): pt (int(start[0] t * (end[0] - start[0]) np.random.randint(-5, 5)), int(start[1] t * (end[1] - start[1]) np.random.randint(-5, 5))) cv2.circle(mask, pt, thickness // 2, 255, -1) masks.append(mask) return torch.from_numpy(np.stack(masks)).float() / 255.0这段代码每次调用会生成batch_size个掩码每个掩码包含 1 到 3 个随机多边形和 0 到 2 条随机划痕。关键参数是max_vertices多边形顶点数和radius缺损半径前者控制形状复杂度后者控制缺失面积。划痕的厚度在 3 到 12 像素之间随机太细了模型容易用边缘插值糊弄过去太粗了训练难度过高导致早期 loss 不下降。在数据加载器里掩码生成必须放在每次__getitem__调用时动态执行而不是预先存好。原因有两个一是动态掩码相当于无限数据增强防止模型记住掩码位置二是同一次训练中模型可能先后见到同一张图的不同缺失位置语义理解更充分。如果使用多进程 DataLoader注意把num_workers设为 CPU 核心数的一半左右Python 的 GIL 会让高并发下的 numpy 操作反而变慢。2.3 训练集的输入合成与归一化陷阱输入合成逻辑看起来简单input image * (1 - mask)把掩码区域置零。但有一个常见陷阱mask必须是 0/1 浮点张量而 OpenCV 的掩码是 0/255 的 uint8直接用会得到全黑结果。另外很多修复论文会额外拼接一个 1 通道的掩码作为模型输入这样网络能明确知道哪里缺失。def collate_fn(batch, img_size(256, 256)): images, masks [], [] for img_path in batch: img cv2.imread(img_path, cv2.IMREAD_COLOR) img cv2.cvtColor(img, cv2.COLOR_BGR2RGB) img cv2.resize(img, img_size) img img.astype(np.float32) / 127.5 - 1.0 mask generate_mask(1, img_size).squeeze(0) images.append(img) masks.append(mask) images torch.from_numpy(np.stack(images)).permute(0, 3, 1, 2) masks torch.stack(masks) inputs images * (1 - masks) # 拼接掩码通道维度变为 [B, 4, H, W] inputs torch.cat([inputs, masks], dim1) return inputs, images, masks归一化用了(x - 127.5) / 127.5把像素映射到 [-1, 1]这是 GAN 类模型的标配。如果后续用纯 L1 损失也可以用 0-1 范围但统一到 [-1, 1] 的好处是生成器最后一层用 Tanh 时天然对齐。permute(0, 3, 1, 2)把 HWC 转为 PyTorch 默认的 CHW。拼接后的输入通道数是 4模型第一层卷积要对应调整很多人换模型时漏掉这一步导致 loading pretrained weights 时报 shape mismatch。3. 模型选型生成对抗与注意力结构的分界点3.1 为什么纯 CNN 不够用修复任务里单纯的卷积网络只能利用局部邻域信息。当缺失区域超过感受野时模型必须“脑补”内容。一个 256x256 输入如果缺失区域是 64x645 层 3x3 卷积的有效感受野只有大约 11x11远远覆盖不了。传统做法是堆深度或膨胀卷积但深层网络的梯度传播和显存开销很快成为瓶颈。生成对抗网络GAN在修复任务里几乎是标配因为它天然适合解决“没有唯一正确答案”的问题。生成器负责输出修复结果判别器学会区分“真实完整区域”和“修复区域”。对抗损失带来的压力让生成器必须输出高分辨率细节而不是平滑模糊的色块。我一般会在生成器里加入注意力机制因为修复任务对远程依赖的需求很明确当你要补一只眼睛另一只眼睛的信息可能在图像的另一侧。3.2 带门控卷积与注意力融合的生成器结构门控卷积Gated Convolution是修复任务里比普通卷积更可靠的选型。普通卷积对掩码区域和无掩码区域一视同仁门控卷积学习一个动态掩码让网络自己决定哪些位置的信息值得传播。实现如下class GatedConv2d(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) self.gate nn.Conv2d(in_ch, out_ch, kernel_size, stride, padding) def forward(self, x): features self.conv(x) gate torch.sigmoid(self.gate(x)) return features * gategate分支的 sigmoid 输出范围在 0 到 1 之间相当于一个软掩码。features * gate实现逐元素的门控。与普通卷积相比门控卷积在掩码区域上不会强制产生响应梯度可以更清晰地流向有效区域。代价是参数翻倍、推理变慢所以在浅层使用效果最好。在 U-Net 的编码器和解码器之间我会插入一个注意力融合模块。具体做法把编码器最深层特征图通过两个并行的 1x1 卷积分别生成 Q 和 K计算空间注意力图再用注意力加权的 V 特征与原始特征拼接class AttentionFusion(nn.Module): def __init__(self, in_ch): super().__init__() self.q_conv nn.Conv2d(in_ch, in_ch // 8, 1) self.k_conv nn.Conv2d(in_ch, in_ch // 8, 1) self.v_conv nn.Conv2d(in_ch, in_ch, 1) def forward(self, x): B, C, H, W x.shape Q self.q_conv(x).view(B, -1, H * W).permute(0, 2, 1) K self.k_conv(x).view(B, -1, H * W) V self.v_conv(x).view(B, -1, H * W) attn torch.softmax(Q K, dim-1) out (V attn.permute(0, 2, 1)).view(B, C, H, W) return x out注意力图attn是 H*W 规模的方阵256 分辨率下就是 65536x65536显存会直接爆炸。常见做法是下采样到 32x32 或 16x16 再算注意力或者用轴向注意力把二维注意力分解成行和列两次一维计算。大多数人忽略的是注意力必须配合残差连接使用否则网络退化为只关注全局而丢失局部纹理细节。3.3 判别器与损失组合的艺术判别器不需要太复杂我用的是 PatchGAN 结构输出一个 NxN 的矩阵而不是单个标量每个值对应输入图像的一个局部区域是否真实。PatchGAN 的好处是参数量小、训练稳定且能强制生成器关注高频细节。损失函数是修复任务里最影响结果的因素。对比一下常见组合损失项公式作用权重建议L1 重建损失x - x_hat/ (CHW)Perceptual 损失|VGG(x) - VGG(x_hat)|_2约束语义特征提升感知质量0.1对抗损失-log D(x_hat)提升细节真实感防模糊0.01风格损失|Gram(VGG(x)) - Gram(VGG(x_hat))|_F约束纹理一致性120权重设置有个经验规律感知损失权重不能超过 L1 的十分之一否则生成结果会偏向平滑。对抗损失权重要更低早期训练时甚至可以关掉等重建损失降到一定程度再打开。我在实际项目中会对掩码区域和非掩码区域分开计算 L1 损失掩码区域权重设为 6 倍这样网络把更多容量用在“真正需要修复”的地方。4. 训练策略与评估让模型真正学会修复4.1 分阶段训练是先粗糙再精细的实用策略修复模型的训练从零开始直接上完整损失容易崩常见做法是分阶段。第一阶段只用 L1 感知损失把生成器训到能补出大致结构第二阶段加入判别器和对抗损失微调细节纹理稳定性。这种课程学习策略让模型先学容易的再逐步增加难度收敛速度反而更快。一个典型的两阶段训练命令看起来像这样# 阶段1纯重建损失学习率1e-4训练30轮 python train.py --phase 1 --loss l1_perceptual --lr 1e-4 --epochs 30 --batch_size 16 # 阶段2加载阶段1权重加入对抗损失学习率降到5e-5 python train.py --phase 2 --loss all --lr 5e-5 --epochs 30 --batch_size 16 \ --pretrained checkpoints/phase1_last.pth训练脚本里我用argparse区分--phase阶段 1 的优化器只接收生成器参数阶段 2 才用 Adam 同时优化生成器和判别器。学习率策略用余弦退火而不固定不变因为修复任务后期需要一个逐渐收窄的步长来稳定对抗训练。批量大小建议 16 起步8 卡并行时每卡 2 个样本就够了。如果你只有单张 2080Ti 或更小显存的卡把批量降到 4同时把输入分辨率从 256 降到 192。低批量下可以用梯度累积模拟更大的 batchoptimizer.zero_grad()改成每 4 个 step 才执行一次这样等效 batch 变为原来的 4 倍只增加一点训练时间。4.2 训练过程最常见的四个崩坏信号与止损方法首先是 loss 变成 NaN。检查学习率是否超过 2e-4、判别器是否比生成器收敛快太多如果是后者降低判别器学习率或给生成器加谱归一化。其次是生成器 loss 降不下去卡在某个平台期这时候需要确认是否跳过了感知损失很多新手只用 L1 会发现图像始终是糊的。第三是判别器 loss 归零说明判别器太强了把对抗损失权重减半或增加判别器 dropout 就能舒缓。第四是训练后期图像出现棋盘格伪影这是转置卷积的固有缺陷把上采样层全部改成最近邻插值加 3x3 卷积能彻底消掉。4.3 指标不能只看 PSNR结构相似度与感知指标用 PSNR 和 SSIM 评估修复质量是业界标准但它们和人的感知并不总是对齐。我会上三个指标一起看其中 LPIPS 更接近人眼感受专门评估感知相似度。import lpips from skimage.metrics import structural_similarity as ssim def evaluate(model, dataloader, device): psnr_list, ssim_list, lpips_list [], [], [] lpips_fn lpips.LPIPS(netalex).to(device) model.eval() with torch.no_grad(): for inputs, targets, masks in dataloader: inputs, targets inputs.to(device), targets.to(device) outputs model(inputs) for i in range(targets.size(0)): # 原图与输出都要归一化到 [0, 1] 才能算 PSNR out (outputs[i].cpu().permute(1, 2, 0).numpy() 1) / 2 tgt (targets[i].cpu().permute(1, 2, 0).numpy() 1) / 2 psnr_list.append(10 * np.log10(1.0 / np.mean((out - tgt) ** 2))) ssim_list.append(ssim(out, tgt, channel_axis-1)) lpips_list.append(lpips_fn(outputs[i:i1], targets[i:i1]).item()) print(fPSNR: {np.mean(psnr_list):.2f} | SSIM: {np.mean(ssim_list):.4f} | LPIPS: {np.mean(lpips_list):.4f})评估时注意输出经过了 Tanh要先反归一化到 [0, 1] 再算 PSNR直接用 [-1, 1] 范围计算会得到虚低的数值。LPIPS 需要 0-1 的输入范围且模型本身内部有归一化逻辑不需要额外处理。一个实用经验LPIPS 在 0.05 以下肉眼已经很难挑出毛病SSIM 在 0.95 以上说明结构保持得不错。4.4 掩码区域的客观指标单独算全局 PSNR 会被大面积未损坏区域稀释。比如掩码只占全图的 5%那么即使修复全错全局指标也只掉不到 5%。所以我一般会额外计算掩码区域内的指标把评测集中在真正需要修复的地方def masked_psnr(outputs, targets, masks): diff (outputs - targets) ** 2 mse (diff * masks).sum() / (masks.sum() * outputs.shape[1]) return 10 * np.log10(1.0 / mse.item())注意masks的 shape 要广播对齐masks是 [B, 1, H, W]diff是 [B, C, H, W]这里除以outputs.shape[1]通道数是在做归一化因为masks.sum()是每个位置的 1 通道像素数之和再乘以通道数才是全部像素数。5. 推理优化与项目交付的隐藏成本5.1 把模型搬上生产环境的四个必要步骤训练结束后模型部署到服务端要过四道关。第一道是密钥量化FP32 权重转成 FP16 推理在 2080Ti 上速度提升约 40%精度几乎没有损失。用 PyTorch 自带 API 就能完成model.half()配合输入inputs.half()记得把所有输入数据都转 FP16否则会报 dtype mismatch。第二步是 TorchScript 或 ONNX 导出目的是脱离 Python 运行时、让 C 服务端加载更快。导出 ONNX 的基准命令python export_onnx.py --checkpoint best.pth --output inpainting.onnx \ --input_size 256 256 --opset 12opset 12很重要太低的版本不支持某些算子导出时动态报错太高的版本部分推理引擎用不起来。导出脚本里务必传入一个示例输入来 trace 计算图让torch.onnx.export执行一次前向。之后用onnxruntime-gpu加载推理吞吐可以比 PyTorch eager 模式再快 30% 左右。第三道是输入尺寸限制。如果你的推理服务接收任意尺寸图片直接送进模型会报错或速度骤降。常见做法是等比缩放短边到 256再居中裁剪到 256x256。要避免长边缩放因为人脸、建筑这类有强几何结构的图像非等比变形会让修复结果出现拉伸畸变。第四道是掩码预处理对齐训练逻辑。线上服务和训练脚本必须用同一套掩码生成和归一化代码最容易出 bug 的地方是线上图片走 OpenCV 读入是 BGR忘记转回 RGB结果所有颜色通道错位生成结果偏蓝。5.2 批量推理的缓存技巧实际部署中同一张图往往需要尝试不同的掩码区域。比如用户先擦除一个文字区域看一眼效果又调整一下掩码位置再看一次。这种情况下完整图像的特征提取是可以复用的。把编码器在无掩码输入上得到的特征缓存下来每次只有掩码变化时只重新跑解码器部分推理吞吐提升在 2 到 4 倍之间。class CachedInpaintingModel(nn.Module): def __init__(self, encoder, decoder): super().__init__() self.encoder encoder self.decoder decoder self.cache None def set_image(self, x): with torch.no_grad(): self.cache self.encoder(x) def forward(self, mask): # 用缓存的编码特征 新掩码直接解码 return self.decoder(self.cache, mask)上面的代码有个前提模型必须是在编码器和解码器之间直接传导特征的结构比如一个标准 U-Net且掩码只在输入侧拼接。如果模型内部有跨层连接依赖掩码信息缓存策略就不适用了这时可以退一步把掩码作为条件输入到解码器的每个上采样层。工程上更常见的方案是保存原始图片的中间特征到 Redis下次请求命中时直接读取等于把计算从服务端搬到了缓存。5.3 快速验证输出质量的视觉对照法修复模型的输出在定量指标之外始终需要人工视觉检查。我习惯在验证脚本中生成三列对照图原始图、掩码标记图掩码区域用红色半透明覆盖、修复结果。把三者横向拼接成一张对比图每轮训练结束自动保存几张到vis/epoch_xx.png这样训练过程中随时可以翻看效果。def save_visualization(images, masks, outputs, targets, save_path): vis [] for i in range(min(4, images.size(0))): orig denormalize(images[i, :3]) # 只取前3通道去掉掩码通道 mask_marked orig.copy() mask_marked[:, masks[i, 0] 0.5] [1.0, 0.0, 0.0] # 掩码红色标记 row np.hstack([orig, mask_marked, denormalize(outputs[i]), denormalize(targets[i])]) vis.append(row) cv2.imwrite(save_path, np.vstack(vis)[:, :, ::-1] * 255)mask_marked[:, masks[i, 0] 0.5] [1.0, 0.0, 0.0]这行把掩码区域整体置红从此一眼就能看出修复边缘是否生硬、是否出现色斑。检查时优先关注四个位置掩码边界处是否有白色光晕预测值和周围环境跳变太大、细纹理是否像“画糊了”纹理合成失败、大块缺失区域是否出现重复纹理模型偷懒复制远处图案、整体色调是否偏灰BatchNorm 统计量漂移。5.4 显存不够时可以使用切片推理如果线上环境只有 8G 显存而输入是 1024x1024 的高清老照片强行全图推理必然 OOM。常见做法是把图片切成 256x256 的 patch带 overlap 地推理再拼回来。关键在拼接边界处理无脑拼会留下明显的接缝。我一般让 patch 之间重叠 32 像素然后对重叠区域的输出按距离加权平均def blend_patches(patches, bboxs, out_shape): result np.zeros(out_shape, dtypenp.float32) weight np.zeros(out_shape, dtypenp.float32) for patch, (x1, y1, x2, y2) in zip(patches, bboxs): result[y1:y2, x1:x2] patch # 纵向线性权重越靠近patch中心权重越高 h, w y2 - y1, x2 - x1 w_row np.hanning(h)[:, None] np.hanning(w)[None, :] weight[y1:y2, x1:x2] w_row return result / weightnp.hanning(h)[:, None] np.hanning(w)[None, :]构造了一个二维汉宁窗patch 中心权重为 1边缘权重衰减到接近 0。多个 patch 在重叠区域相加分母weight做归一化接缝几乎看不见。这个技巧和处理大图时模型感受野不足的问题是两回事后者需要调整网络结构或者使用更大感受野的卷积切片只能解决显存限制。本文还有配套的精品资源点击获取