☰
GAN图像修复实战指南:从原理到PyTorch源码解析
2026/9/28 1:45:53 网站建设 项目流程

简介:面向Python方向的毕业设计、课程作业与项目实战需求,这套基于深度生成对抗网络(GAN)的图像修复模型提供了完整可运行的源码与文档说明。项目覆盖生成器与判别器搭建、损失函数定义、DCGAN训练流程以及图像缺失区域的补全推理,适合计算机相关专业学生完成大作业或作为入门GAN项目的练习素材。压缩包共7个文件,以6个Python脚本和1个Markdown说明文档为主,脚本分别承担模型结构定义、训练流程与图像补全推理,说明文档帮助快速理解项目结构;整个压缩包仅12KB,结构紧凑,便于下载后直接阅读和本地调试。这套压缩包目前已有164人学习使用,经导师指导并评审通过,源码均经过本地编译调试,可稳定运行。下载后可获得一套包含模型定义、训练脚本、修复演示与说明文档的完整项目方案,对理解GAN在图像修复场景中的应用及Python工程组织方式均有参考价值。

1. 图像修复为什么非GAN不可:一个被小瑕疵逼出来的方向

你手上有一张祖辈的合影,中间一道折痕把脸劈成两半;或者一张商品图被杂物遮住一角,想卖给电商平台却怎么裁都裁不好。传统图像修复用的是扩散或纹理合成——从缺口边缘往内部“长”像素,对付细线、小污点没问题,一旦缺失区域超过几十像素,结果就是一片模糊的糊状物。深度生成对抗网络gan的出现让这个场景发生了质变:生成器能结合全局语义,把“这里应该是一只眼睛”这类高层信息补出来,而不是只盯着边缘做插值。带源码和文档说明的python实现,核心价值就是让你能快速拿到一条可修改、可重训的完整链路——数据处理、掩码生成、模型训练、推理评估,而不是陷在单篇论文的无尽复现里。这篇文章面向的是想真正把gan图像修复跑通、且愿意自己调参的人,我会按自己拆这类源码的习惯,把原理、代码和踩坑点一条条讲透。

2. 生成对抗网络图像修复的原理与模型选型:先看懂Loss再看源码

2.1 从Context Encoder到LaMa:修复模型的四代演变

最早把深度学习用于图像修复的代表是2016年的Context Encoder,思路很简单:一个基于AlexNet的编码器把残缺图压缩成特征向量,解码器再从这个向量重建整张图,配合GAN判别器让重建结果“看起来真实”。它的短板在于固定尺寸输入、不支持任意形状掩码,而且分辨率限制在128左右。

第二代是Partial Convolution系列(NVIDIA 2018),核心改动是卷积层只对有效像素进行计算,并把掩码作为输入通道一起更新。这样一来模型天然能处理任意形状、任意位置的掩码,生成的边缘过渡自然了很多。很多开源python源码都跑到这一步。

第三代是注意力机制的引入,代表是GCAN和EdgeConnect。GCAN生成两个分支——一个负责粗糙补全,一个负责精细纹理修复,然后靠上下文注意力把已知区域的纹理“搬运”到缺失区域。EdgeConnect更特殊,它把问题拆成两步:先修复边缘图,再在边缘引导下补全颜色纹理。

第四代就是近几年的快速模型(LaMa、Fire等)。LaMa用傅里叶卷积替换普通卷积,感受野直接拉满整个图像,加上一个像素级的高分辨率损失,能在50毫秒级出结果,且对大掩码的抗崩坏能力很强。如果你拿到源码不确定自己属于哪一代,先看卷积层和损失函数就能分清楚——修复网络的核心差异,几乎都体现在这两处。

2.2 生成器、判别器与损失函数怎么配:三种常见组合的取舍

GAN修复模型的训练本质上是一个多目标博弈:生成器想把残缺图补得让判别器信以为真,判别器则要精准识别哪里是补出来的。单拿图像重建的L2损失会得到模糊结果,单拿对抗损失又容易产生纹理爆炸。实际源码里常见的损失组合有三种。

组合生成器损失判别器损失适用场景训练成本
基础版L1 + 局部对抗掩码区域对抗小掩码、快速验证低
感知版L1 + Perceptual + 全局对抗全局对抗通用修复、人眼观感优先中
高保真版L1 + Perceptual + Style + 双判别器全局局部双判别器大面积缺失、老照片修复高

基础版适合练手,生成器直接吃“残缺图+掩码”,输出修复图,判别器只看掩码区域。感知版是最推荐的默认起点,L1保证结构,Perceptual(用VGG16中间层特征做比较)保住语义细节,全局对抗保证整体真实感。高保真版效果好但训练极不稳定,一般配合谱归一化和梯度惩罚才能收敛,新手容易被直接劝退。

2.3 修复模型效果验证的四个指标:PSNR、SSIM、L1、FID选哪几个

源码里如果带验证脚本,通常就是这四个指标。PSNR计算简单,数值越高失真越小,但它对结构信息不敏感,两个 visually 完全不同的图可能有接近的PSNR。SSIM则模拟人眼对亮度、对比度、结构的感知,修复任务里比PSNR更有参考价值。

L1是生成器损失函数里也常用的那个度量,直接算像素差绝对值,好处是梯度平稳。FID专门用来评估GAN生成分布与真实分布的距离,值越低越好。我的建议是发布或对比模型时同时报PSNR/SSIM/FID三个,单独报任何一个都容易被挑毛病——PSNR好但FID差的模型往往是“像素级逼近但风格不对”,这个组合问题在修复场景里经常翻车。

下面这段是源码里最常见的感知损失实现,我用PyTorch风格写,函数返回L1与VGG感知损失的和:

import torch import torch.nn.functional as F from torchvision import models class PerceptualLoss(torch.nn.Module): def __init__(self): super().__init__() # 取VGG16前三个block的特征,不参与梯度更新 vgg = models.vgg16(weights=models.VGG16_Weights.IMAGENET1K_V1) self.blocks = torch.nn.ModuleList([ vgg.features[:4], # conv1_2输出 vgg.features[4:9], # conv2_2输出 vgg.features[9:16], # conv3_3输出 ]) for p in self.blocks.parameters(): p.requires_grad = False def forward(self, pred, target): loss = 0.0 x_pred, x_target = pred, target for block in self.blocks: x_pred = block(x_pred) x_target = block(x_target) loss += F.l1_loss(x_pred, x_target) return loss

这段代码把VGG16拆成三段,对每一段输出的特征图都算L1距离。需要注意为什么不用VGG的最后一层:最后一层特征过于抽象,已经分辨不出像素级别的纹理差异,起不到引导生成器补细节的作用。前几层特征在修复任务里更关键。

另外,代码里requires_grad = False这一步别省。如果不冻结VGG,反向传播会更新VGG的参数,结果就是生成器学会了“用特征匹配欺骗特征提取器”,图像结构的真实一致性反而不保证。我自己踩过这个坑,花了一晚上才找到 loss 诡异下降但图越修越花的原因。

2.4 为什么大多数开源源码选PatchGAN做判别器

如果你翻开修复源码,发现它的判别器前向输出不是单个数而是某个矩阵,那多半是PatchGAN。它的原理是把图像切成多个小patch,对每个patch独立判断真假,最后对所有patch的判别结果取平均。

PatchGAN对高分辨率图像尤其友好。整图判别器只输出一个概率,很容易被大面积平滑区域骗过去,而局部细节的真实感反而被忽略。PatchGAN强制判别器关注局部纹理的真实性,这对图像修复来说正中要害——缺失区域补出来的纹理如果在局部尺度是合理的,整体观感就不会太差。后面第4章写判别器代码时会一起给出来,这里先明确一个选型结论:分支任务不明确的源码,优先用PatchGAN做判别器,踩坑率最低。

3. 把数据准备好:图像裁剪、掩码生成与训练集划分(源码里最容易忽略的三处)

3.1 目录组织与图片读入:先解决“图片太大进不了显存”

仓储类源码通常假设你已经有一批“正常图”,训练时随机挖掉一块再让模型修复,而不是真的准备了一堆带破损的图片。目录结构一般长这样:

data/ train/ 00001.jpg 00002.jpg val/ 00001.jpg

图片读入有一个常见问题:直接把原图塞进模型,而原图尺寸往往在1000像素以上。生成器Encoder虽然有下采样层,但显存开销依旧可能爆炸。安全做法是在Dataset类内部做一次中心裁剪或随机缩放:

import torch from torch.utils.data import Dataset from PIL import Image import torchvision.transforms as T class InpaintingDataset(Dataset): def __init__(self, root_dir, img_size=256): self.paths = sorted(glob.glob(root_dir + "/*.jpg")) self.img_size = img_size self.transform = T.Compose([ T.Resize((img_size, img_size)), T.ToTensor(), T.Normalize([0.5], [0.5]) # 单通道灰度图用;RGB图要三个通道分别写 ]) def __len__(self): return len(self.paths) def __getitem__(self, idx): img = Image.open(self.paths[idx]).convert("RGB") img = self.transform(img) return img

这里有两个参数值得细说。img_size=256是修复任务最常驻的平衡点:尺寸太小,模型学不到细节;尺寸太大,2333的显存根本不够。我自己训练时先跑256,模型稳定后再用同一个权重微调512。

Normalize([0.5], [0.5])这种写法只适合灰度图,RGB图要写Normalize([0.5, 0.5, 0.5], [0.5, 0.5, 0.5]),否则通道数不匹配会当场报错。更隐蔽的问题是normalize之后所有像素值都会变成[-1,1]范围,这对应生成器输出层的Tanh激活函数,后面推理时要把输出重新映射回[0,1]才能保存成图片。

3.2 矩形掩码与不规则掩码的生成:别一开始就上Brush掩码

掩码决定了训练模型的能力边界。只训练矩形掩码,模型遇到修长的不规则遮挡会手足无措;只训练不规则掩码,模型可能出现边缘收缩的问题。业界最常见的做法是两者混合。

源码里的掩码通常以二值numpy数组表示,1代表缺失区域,0代表已知区域。最简单的矩形掩码生成方式:

import numpy as np def make_rect_mask(img_size=256, max_ratio=0.35): mask = np.zeros((img_size, img_size), dtype=np.float32) h = np.random.randint(img_size * 0.1, img_size * max_ratio) w = np.random.randint(img_size * 0.1, img_size * max_ratio) y = np.random.randint(0, img_size - h) x = np.random.randint(0, img_size - w) mask[y:y+h, x:x+w] = 1.0 return mask

这个函数生成的遮罩占整图面积约10%~35%,覆盖位置随机。max_ratio不建议超过0.4,否则缺失面积太大,生成器只能靠猜,训练出来的结果会偏“幻觉”而非修复。

不规则掩码的经典生成方式是模拟多段线条,类似画笔涂抹的效果:

def make_brush_mask(img_size=256, max_brush=8): mask = np.zeros((img_size, img_size), dtype=np.float32) for _ in range(np.random.randint(2, max_brush)): # 随机起点、随机方向、随机长度 y0, x0 = np.random.randint(0, img_size, size=2) angle = np.random.uniform(0, 2 * np.pi) length = np.random.randint(20, 120) width = np.random.randint(2, 12) # 沿直线方向画粗线 for t in range(length): y = int(y0 + t * np.sin(angle)) x = int(x0 + t * np.cos(angle)) if 0 <= y < img_size and 0 <= x < img_size: cv2.circle(mask, (x, y), width, 1.0, thickness=-1) return mask

这段代码用OpenCV的circle函数模拟画笔,每一笔都是一个随机方向、随机宽度的线段。max_brush控制笔画数,笔画越多掩码越碎、面积越大。

训练时我一般让每张图有50%概率使用矩形掩码、50%概率使用不规则掩码。注意掩码必须和图片在同一设备上、同一尺寸,后续要concat成生成器的输入,后面会用到。掩码生成也千万别固定随机种子,否则每个epoch挖的洞都一样,模型会过拟合到特定掩码位置。

3.3 训练集/验证集划分与归一化细节:mean/std不要随手写0.5

仓库里如果已经给了划分脚本,看看有没有用随机种子打乱。常见的坑是划分时用了带排序的文件名,比如000001.jpg到001000.jpg,按前800/后200切,结果恰好像白天和黑夜的照片集中在不同区间,训练集里全是白天、验证集全是黑夜,验证loss一直震荡。

更稳妥的做法是在读DataLoader之前显式做一次随机split:

from sklearn.model_selection import train_test_split paths = glob.glob("data/raw/*.jpg") train_paths, val_paths = train_test_split(paths, test_size=0.1, random_state=42)

random_state=42是为了让每次复现划分一致。源码里如果没带这一步,我就会补上,不然没法公平地对比不同训练轮次的模型效果。

归一化系数不要无脑照抄。ImageNet上常用的[0.485, 0.456, 0.406]是针对自然图像的统计均值,适合一般照片修复;如果数据集是医学影像或特殊的灰度图,就得自己统计。源码里如果看到写死的均值,建议先跑一小批数据核对一下,我遇到过用ImageNet均值训练工业零件图,整个色调偏蓝的问题,排查了很久才发现是归一化没对齐。

4. 从零跑通训练:最小可复现的生成器、判别器与训练循环

4.1 生成器:UNet换掉普通卷积,修复边界的正确姿势

修复任务里几乎看不到纯Encoder-Decoder,因为修复图需要保留原始图像的纹理细节,而UNet的skip connection能把编码器的高分辨率特征直接传到解码器,避免细节丢失。但普通UNet有它的坑——缺失区域的卷积特征会被“洞”污染,训练慢且边缘修复看着像糊了一层雾。

我常用的是UNet + 掩码更新的做法:每一层卷积之后,特征图再乘以一个持续下采样的掩码版本,让模型永远不把洞内无效特征继续向前传递。

import torch.nn as nn class Generator(nn.Module): def __init__(self, in_ch=4, base_ch=64): super().__init__() # 输入是“RGB图+掩码”四通道 self.enc1 = nn.Sequential( nn.Conv2d(in_ch, base_ch, 4, 2, 1), nn.LeakyReLU(0.2, inplace=True)) self.enc2 = nn.Sequential( nn.Conv2d(base_ch, base_ch*2, 4, 2, 1), nn.BatchNorm2d(base_ch*2), nn.LeakyReLU(0.2, inplace=True)) self.dec1 = nn.Sequential( nn.ConvTranspose2d(base_ch*2, base_ch, 4, 2, 1), nn.BatchNorm2d(base_ch), nn.ReLU(inplace=True)) self.dec2 = nn.Sequential( nn.ConvTranspose2d(base_ch*2, 3, 4, 2, 1), nn.Tanh()) # 输出范围[-1,1] self.up_mask = nn.Upsample(scale_factor=2, mode='nearest') def forward(self, img, mask): x = torch.cat([img, mask], dim=1) e1 = self.enc1(x) e2 = self.enc2(e1) d1 = self.dec1(e2) d1 = d1 * self.up_mask(mask) # 只保留已知区域特征 d2 = self.dec2(torch.cat([d1, e1], dim=1)) return d2

这段代码里in_ch=4是关键:输入不是三通道RGB,而是把掩码作为额外通道拼进去。生成器看到掩码之后才知道要修复哪个位置。

mask在下采样之后尺寸会逐步减半,所以在dec1之后我用up_mask把掩码放大回当前分辨率,再乘到特征图上。这个操作的目的是确保编码器学到的已知区域特征不被缩小后的掩码误伤,它本质上是Partial Convolution思想的简化版。

生成器输出用Tanh是因为数据归一化到[-1,1],如果归一化时用的是ORIGINAL范围,这里就要换成Sigmoid。这是源码里常见的“隐式约定”,改数据归一化时最容易忘改这里,不配对会导致loss起步就极高而且长时间降不下去。

4.2 判别器:PatchGAN为什么比整图判别更适合修复场景

整图判别器给的是单一布尔判断,容易忽略局部纹理的细节。PatchGAN 的核心是把输入切成多个小patch,分别判断真假——我一般用70x70的receptive field版本,性能好、内存开销适度。

下面的判别器输入是“修复后的整图”或“原图”,输出还是3D特征图,最后一维的每个点代表该位置的patch真伪:

class PatchDiscriminator(nn.Module): def __init__(self, in_ch=3, base_ch=64): super().__init__() self.model = nn.Sequential( nn.Conv2d(in_ch, base_ch, 4, 2, 1), nn.LeakyReLU(0.2, inplace=True), nn.Conv2d(base_ch, base_ch*2, 4, 2, 1), nn.BatchNorm2d(base_ch*2), nn.LeakyReLU(0.2, inplace=True), nn.Conv2d(base_ch*2, base_ch*4, 4, 2, 1), nn.BatchNorm2d(base_ch*4), nn.LeakyReLU(0.2, inplace=True), nn.Conv2d(base_ch*4, 1, 4, 1, 1) ) def forward(self, x): return self.model(x)

最后那个卷积的stride=1(不是2)是为了保留空间尺寸,输出是 N×1×H×W 的patch评分图,每个位置的值代表该区域是真图的置信度。实际loss计算用BCEWithLogitsLoss作用在这个评分图上。

这里BatchNorm2d如果放在判别器最后一层前,可能导致训练初期震荡厉害。我在调参时会把最后一个block的BN去掉,换成普通的Conv+LeakyReLU,稳定感会好很多。这种小幅改动不影响整体架构,但确实是我自己对比下来明显有效的细节。

4.3 训练循环与超参数:lr=1e-4、BS=4起手,再谈优化

训练GAN修复模型和训练普通分类网络感觉完全是两回事——要同时watch多个loss,且稍有不慎就mode collapse。给一个最精简的循环模板:

import torch import torch.optim as optim from torch.autograd import Variable G = Generator().cuda() D = PatchDiscriminator().cuda() optG = optim.Adam(G.parameters(), lr=1e-4, betas=(0.5, 0.999)) optD = optim.Adam(D.parameters(), lr=4e-4, betas=(0.5, 0.999)) criterion = torch.nn.BCEWithLogitsLoss() for epoch in range(200): for batch in dataloader: real_img = batch.cuda() mask = make_rect_mask().unsqueeze(0).cuda() masked_img = real_img * (1 - mask) # 训练判别器 fake_img = G(masked_img, mask) fake_img = fake_img * mask + masked_img * (1 - mask) # 已知区域回贴原图 pred_fake = D(fake_img.detach()) pred_real = D(real_img) lossD = criterion(pred_fake, torch.zeros_like(pred_fake)) + \ criterion(pred_real, torch.ones_like(pred_real)) optD.zero_grad() lossD.backward() optD.step() # 训练生成器 pred_fake = D(fake_img) lossG_adv = criterion(pred_fake, torch.ones_like(pred_fake)) lossG_l1 = torch.nn.L1Loss()(fake_img, real_img) lossG = lossG_adv + 100 * lossG_l1 optG.zero_grad() lossG.backward() optG.step()

四个参数需要特别说明。lr=1e-4是生成器的学习率,判别器用4e-4,这符合“判别器带得更快”的经验——判别器如果跟不上生成器的进化速度,对抗训练就会退化成一个自娱自乐的循环。betas=(0.5, 0.999)是DCGAN系列论文里的标配设定,默认Adam的betas=(0.9, 0.999)在GAN训练里容易导致梯度震荡,影响收敛稳定性。

fake_img = fake_img * mask + masked_img * (1 - mask)这一行是修复任务和其他生成任务的关键区别——已知区域的像素是有标准答案的,必须原样保留,只要让生成器专注补缺失区域。100 * lossG_l1里的权重需要根据你的数据调整,太小会让修复区域模糊,太大会让对抗训练名存实亡。我自己测试的经验是50~150这个区间,边调边看生成区域纹理,浮在表面上比对loss数值更可靠。

4.4 断点续训与日志:省一晚上的后悔药

训练跑到第40轮机器断电,如果没做断点续训,前面三天的算力就彻底打水漂。标准做法是把Generator、Discriminator、optimizer状态、当前epoch和随机种子一起存下来:

torch.save({ 'g': G.state_dict(), 'd': D.state_dict(), 'opt_g': optG.state_dict(), 'opt_d': optD.state_dict(), 'epoch': epoch, 'rng_state': torch.get_rng_state(), }, "checkpoint.pth")

恢复时倒序加载,把每个对象注入对应模型即可:

ckpt = torch.load("checkpoint.pth") G.load_state_dict(ckpt['g']) D.load_state_dict(ckpt['d']) optG.load_state_dict(ckpt['opt_g']) optD.load_state_dict(ckpt['opt_d']) epoch_start = ckpt['epoch'] + 1

有两个细节容易被忽略。其一是torch.get_rng_state最好也保存,否则后续数据增强的随机性无法复现,同一份源码两次训练出来可能差很远。其二是加载optimizer状态后建议少调一次学习率,因为Adam的动量信息是累计在state里的,直接改lr会让之前momentum的状态和新的lr不匹配,模型可能要花几百步才能“缓过来”。

5. 避坑清单:GAN图像修复最容易翻车的5个问题

5.1 棋盘伪影:生成器输出像蒙了一层网格

现象:修复区域有明显的格子状纹理,尤其在渐变平滑的地方肉眼可辨。

原因:转置卷积的kernel size不是偶数导致重叠不平整,尤其是kernel=4、stride=2这种配置最容易出现交错叠影。

解决:把生成器里所有ConvTranspose2d换成“Upsample + Conv2d”组合。上采样的插值方式用nearest,然后接一个普通卷积做特征学习,这样能彻底消除棋盘伪影。如果不想改动网络结构,另一种妥协是用kernel=3、stride=1的分步上采样。

5.2 修复区域一片发灰:判别器太强而生成器跟不上

现象:loss看起来在下降,但修复区域颜色趋同,像打了马赛克的灰色块。

原因:生成器面对过于强大的判别器时选择了保守策略——输出接近均值的结果,因为这种输出的平均损失最小,但判别器能轻易识别成fake,于是进入“生成器摆烂、对抗loss震荡”的恶性循环。

解决:先调低判别器学习率到生成器的1/4(比如G=2e-4,D=5e-5)。同时把lossG_l1的权重提高,用强L1把生成器“拽”回清晰解附近。还有一种快速验证手段:把BatchNorm换成InstanceNorm,某些数据集上BN会使颜色偏移,换成IN之后发灰问题直接消失。

5.3 掩码边缘生硬:修复区域和原图像是两张图拼起来的

现象:修复大体合理,但边缘一圈有明显割裂,放大看有白边或黑边。

原因:模型只在掩码内生成,训练时mask是一个硬0/1边界,没有对边缘做平滑过渡。

解决:训练时对mask随机做一个小半径的膨胀和腐蚀(让掩码边缘抖动),或者对mask做高斯模糊后参与计算。更简单的方式是给掩码边缘区域额外加一个边缘loss:先提取原图的边缘,再对修复图的边缘求L1距离。前提是你源码里有边缘特征提取的模块,EdgeConnect那套结构天然带这个模块,普通UNet可以先用模糊平滑mask顶住。

5.4 Generator Loss居高不下,D Loss却快速归零

现象:训练刚开始几百步,判别器loss就接近0且不回升,生成器loss却越走越高。

原因:判别器学得太快,直接把生成样本和真实样本的特征空间完全分开了,生成器接收到的是几乎为0的更新梯度,动弹不得。

解决:用label smoothing。把真实样本的标签从1改成0.9~0.95,让判别器即使“猜对了”也得不到满分,保留一定的梯度信息给生成器。另一个选择是固定判别器不训练几百步,让生成器先单独用L1损失收敛一波,再开启对抗训练。这个“预热”操作见效很快,强烈建议第一次跑项目时安排上。

5.5 颜色整体偏移:修复区域色相偏紫或偏蓝

现象:语义结构修复对了,但颜色明显和环境不符,人眼一眼发现异常。

原因:最常见的是模型用了近似特征匹配,却在不同色彩空间之间做变换时发生了浮点误差。也可能是归一化均值和数据本身不匹配,尤其是灰度图用了三通道均值归一化。

解决:先在推理端检查归一化和反归一化是否配对;如果确认无误,把判别器输入从RGB改成RGB+YUV(把YUV各通道concat进去),颜色一致性会显著提升。这个技巧在部分老照片修复源码里是默认选项,但普通UNet源码里很少见,属于性价比很高的增强。

6. 让修复结果再上一个台阶:掩码迭代与图像金字塔推理

很多人在第4章训练跑通之后就觉得完事了,但直接拿训练好的模型去修一张1K分辨率的大图,往往会发现修复结果发虚、细节糊。原因很简单:模型训练在256分辨率下,上采样到1K之后高频信息根本不存在。这里有一个不需要重新训练的技巧——把推理过程改成分阶段迭代。

def iterative_inpaint(model, img, mask, steps=4): # img: 1x3xHxW, 数值范围[-1,1] cur_img = img.clone() cur_mask = mask.clone() for _ in range(steps): # 对图像做下采样修复,再上采样回来 with torch.no_grad(): out = model(cur_img, cur_mask) # 只替换掩码中心区域(边缘留一点缓冲) refined_mask = torch.nn.functional.max_pool2d( cur_mask, kernel_size=3, stride=1, padding=1) cur_img = out * refined_mask + cur_img * (1 - refined_mask) cur_mask = refined_mask return cur_img

原理是通过多次迭代、每次只修一小圈的方式,逐层收缩掩码。第一次迭代负责修复大轮廓,之后的迭代误差梯度会被慢慢消减,纹理细节越到后面越清晰。max_pool2d在这里的作用是把掩码“膨胀”一圈,让每一次修复都从已知区域延伸进去。这种方法对大掩码的修复效果提升比较明显,唯一的代价就是推理时间按步数线性增长。

另外一个常用技巧是图像金字塔推理:把整图缩放到512、再缩到256,分别跑一次修复,把256的修复结果上采样回到512作为初始参考,最后在512上再精修一次。这种做法在处理高清老照片时比直接跑原图更稳定,因为模型在高分辨率上只要补细节,不需要重新构想全局结构。

这些技巧的共同点是它们都没改动训练代码,只是改变了推理的输入输出交互方式。如果你在源码基础上想继续做提升,这是性价比最高的一步。最后分享一个我的个人习惯:每次训练完都会留一批固定的“难例图”(大掩码、纹理复杂、光照异常),只在这批图上评估模型更新效果,单纯盯loss变化容易自欺欺人。忠实复现这条链路,少则三五天,多则两周,就能把GAN图像修复的每个环节摸到实感。希望这条路线对你有帮助,也欢迎把你跑通过程中踩到的新坑记下来——每多一个人记录,后面的复现者就能少加一次班。

本文还有配套的精品资源,点击获取

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询