简介:基于Pytorch实现对偶生成对抗网络进行图像去雾的完整工程,包含源码、训练好的模型与文档说明。资源包共25个文件,以10个Python脚本为主,涵盖网络定义、训练流程、预测推理及参数解析等模块;另有5张JPG与6张PNG图像用于测试和效果展示,2个pkl格式的判别器模型可直接加载使用,并附带README文档说明项目结构。整个压缩包仅21.23MB,轻量易部署。目前已有144人学习下载,适合正在完成毕业设计、课程设计或期末大作业的计算机专业学生,也适合对图像去雾和GAN实战感兴趣的学习者。项目经导师指导并获高分评价,代码结构清晰、注释完善,从数据加载、模型构建到可视化展示均有对应脚本,配合预训练模型可快速复现去雾效果,降低入门门槛。若想深入理解对偶对抗网络的训练细节,也可借助源码和文档逐模块拆解,是一次完整且可直接运行的项目实践。
1. 对偶生成对抗网络去雾,为什么是“成对数据”的反面解法
图像去雾这个任务,绝大多数人第一反应是去搜集成对的“有雾图+清晰图”训练集,然后跑一个监督模型。真实工程场景里这套路最先翻车:你手里的有雾图和清晰图往往不是同一时刻、同一机位拍的,对齐误差直接让模型学到“重影”而不是“去雾”。所以看到“Pytorch实现对偶生成对抗网络”这个标题时,我关心的不是它能不能跑通,而是它怎么绕开成对数据的死结。对偶GAN的核心思路是让“去雾器”和“加雾器”互为逆映射,用循环一致性约束在没有像素级对齐数据的条件下也能训练。这篇文章从原理、推理、训练到踩坑,把这条路线完整拆开。
2. 先把原理立住:对偶GAN去雾的生成器、判别器与循环一致性
2.1 对偶是什么意思:去雾器与加雾器互为逆映射
“对偶”这个词在GAN框架里不是修饰,而是结构。常见的做法是同时训练两个生成器:G_A负责把有雾图像映射成清晰图像,G_B负责把清晰图像映射成有雾图像。二者不是独立训练,而是共享一个循环约束:一张有雾图像经过G_A去雾后再经过G_B加雾,应该回到原来的有雾图像;反过来也一样。这个往返误差就是循环一致性损失,它替代了“成对像素监督”。
我最初接触这个结构时有个误解,以为两个生成器都训练好之后只用G_A就行,G_B只是辅助训练。实际工程里G_B的价值不只在训练期。推理之前我会额外做一次验证:拿一张完全不同的雾图,先过G_B再过G_A,看能否还原出接近原始图像的版本。如果还原明显失败,说明训练阶段的循环一致性没有充分收敛,这时的G_A输出往往也不可靠。
对比一下CycleGAN,对偶GAN去雾的真正差异在生成器设计。去雾器不需要保持所有高频纹理,它更需要保持颜色恒常性,因此生成器主体用ResNet残差块是通用做法,但输入和输出侧的卷积层要额外注意归一化方式。在Pytorch里,我一般用InstanceNorm而不是BatchNorm,因为推理时的batch size经常是1,BatchNorm统计量在单样本时抖动很大。
2.2 为什么普通GAN做不了去雾:判别器的“审美”问题
普通GAN做去雾最大的问题是判别器只有“真假”一个标准,它判断的是“图像是不是清晰”,而不是“这张雾图的清晰版本长什么样”。于是生成器会找出判别器的盲区:把图像变得对比度很高、颜色过度饱和,判别器觉得“高清”就给过了,人眼一看是灾难。
对偶GAN的循环一致性损失正好补上了这个洞。生成器G_A想骗过判别器只解决了一半问题——G_A的输出还要能被G_B还原成原始雾图,这逼迫G_A必须保留真实场景的结构和颜色信息,不能随意“风格化”。Pytorch实现里这个损失非常好写,前向两次得到重建图后用L1距离约束,但权重的设置影响很大,这个放到第四节展开。
从选型角度说,如果你手上的数据是人工合成的成对雾图(比如用大气散射模型渲染的),那直接用监督学习更稳,没必要上对偶GAN。对偶GAN适合的是“只有一批雾图、没有对应的清晰图”的场景,或者两张图来源不同、像素级对齐做不到的场景。标题里强调“对偶生成对抗网络”而不是普通的DebazeGAN,说明作者默认了这种无成对数据的条件,这是理解整个项目源码的第一把钥匙。
2.3 判别器该看多细:PatchGAN的直觉与Pytorch里的实现
图像GAN的判别器不应该只看整张图的真假。去雾这种任务,雾的浓度在空间上是非均匀的,远处浓、近处淡,一个全局判别器很容易漏掉局部区域的残留雾。VGG感知损失能从特征层面缓解,但判别器结构上更直接的解法是用PatchGAN——判别器对图像分块做真假判断,每个patch独立评分。
具体做法是判别器不用全连接层收成单个数,而是输出一个N×N的特征图,每个像素位置对应原图一个感受野区域的真假概率。Pytorch里这几乎不需要额外代码:把最后一层卷积的输出通道设为1即可。patch的尺寸决定感受野,常见的是70×70,也就是每个输出像素看原图70×70的区域。patch太大趋近全局判别器,太小则容易被骗。
选择PatchGAN而不是全局判别器还有一层工程原因:显存。训练对偶GAN要同时塞下两组生成器和两组判别器,全局判别器一上来的分辨率要求会让batch size掉到个位数。PatchGAN输出张量小,反向传播的计算量也更可控。在Pytorch里训练时,我会把判别器的真实标签设置为0.9而不是1.0,这是GAN训练里的标签平滑技巧,能减弱判别器过拟合,稳定对抗过程。
3. 跑通最小推理路径:加载模型辨别雾图的三个关键参数
3.1 环境与依赖:把Pytorch基础框架装到能用GPU推理
标题给了“训练好的模型”,所以第一步是把环境配到能加载权重、做前向推理。这里先说结论:不要一上来就装最新版Pytorch,先看模型权重是用什么版本序列化保存的。torch.save在Pytorch 1.x和2.x之间的兼容性总体没问题,但个别算子(尤其是旧版本的自定义Module)在加载时可能报missing key。如果你手里的权重是以state_dict形式保存的,从Pytorch 1.8到2.x都能加载。
创建环境时我习惯用conda管理,命令行如下:
conda create -n defog python=3.9 conda activate defog pip install torch==2.1.0 torchvision==0.16.0 --index-url https://download.pytorch.org/whl/cu118 pip install numpy opencv-python pillow matplotlib这里刻意固定了torch版本而不是装latest,是为了让CUDA和cuDNN的版本对应关系可预期。cu118对应CUDA 11.8,如果你机器的驱动只支持CUDA 12.x,可以考虑把后面的cu118换成cu121。判断驱动支持情况很简单,在终端执行nvidia-smi看右上角的CUDA Version,这个数字是你驱动能支持的上限,torch的cu版本不能超过它。
安装完成后,验证Pytorch能否调用GPU是必须的一步,别跳过:
python -c "import torch; print(torch.__version__); print(torch.cuda.is_available()); print(torch.cuda.get_device_name(0) if torch.cuda.is_available() else 'CPU')"如果输出is_available为False,不要先怀疑显卡坏了,八成是torch的CUDA编译版本和驱动不匹配。换用CPU推理虽然能跑,但一张1080p图像的去雾前向在GPU上可能只要几十毫秒,CPU上会慢一到两个数量级。
3.2 加载训练好的模型:torch.load与模型结构定义的配套关系
Pytorch加载权重的标准姿势是state_dict方式,但很多人第一次加载失败是不知道“模型结构要先实例化才能load”。模型文件里存的只是张量字典,不包含网络结构定义,你必须用源码里的模型类先把网络搭出来,再把字典填进去。
import torch from model import DehazeGenerator # 来自项目源码 device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') model = DehazeGenerator(in_channels=3, out_channels=3, ngf=64) state_dict = torch.load('checkpoints/defog_gan.pth', map_location=device) if 'state_dict' in state_dict: state_dict = state_dict['state_dict'] model.load_state_dict(state_dict) model.to(device) model.eval()in_channels和out_channels都是3,对应RGB输入输出,ngf=64是生成器第一层卷积的通道数,这个值在训练时定了就不能改,否则权重形状对不上报size mismatch。map_location=device是把权重张量加载到CPU再搬运到GPU,避免在无GPU机器上直接死锁。torch.load会默认把张量加载到保存时的设备,如果权重是在GPU上保存的,而无GPU机器不去指定map_location,会直接报错。
加载时报错最典型的有三种。第一种是size mismatch,说明模型结构参数和训练时不一致;第二种是missing key(s),说明权重文件里缺了模块;第三种是unexpected key(s),说明权重是从另一个结构更复杂的模型里拷贝出来的,不是完全匹配。第三种情况可以尝试只加载前几层的参数来做迁移学习,但做推理就算了。
3.3 推理入口:前向一次与结果后处理
模型加载完成后,推理代码本身很短,但后处理决定了你看上去的自定义去雾效果是否正常。Pytorch默认的输入张量形状是(N, C, H, W),数值范围在0到1之间(训练时做了归一化),而图片读取出来通常是0到255的uint8格式,这个转换写错会让输出一片漆黑。
import cv2 import numpy as np import torchvision.transforms as T img = cv2.imread('hazy.jpg') img_rgb = cv2.cvtColor(img, cv2.COLOR_BGR2RGB).astype(np.float32) / 255.0 tensor = torch.from_numpy(img_rgb).permute(2, 0, 1).unsqueeze(0).to(device) with torch.no_grad(): output = model(tensor) # 输出范围在0~1附近 output = torch.clamp(output, 0.0, 1.0) out_img = output.squeeze(0).permute(1, 2, 0).cpu().numpy() out_img = (out_img * 255.0).astype(np.uint8) out_bgr = cv2.cvtColor(out_img, cv2.COLOR_RGB2BGR) cv2.imwrite('dehazed.jpg', out_bgr)torch.no_grad()是推理时的固定动作,目的是关闭梯度计算,省显存也提速。permute(2, 0, 1)把高度、宽度、通道的排列转换成通道、高度、宽度,这一步写漏了会得到“图像被旋转90度的花屏”。最后的clamp这一步不能省,训练时生成器的输出可能略超出0-1范围,直接转uint8会出现像素截断设置的偏移。
我一般在写完这段推理后,做一件事:把输出图像保存为16位PNG再观察直方图。如果直方图集中在某一小段区间且没有明显拉伸,说明模型权重可能出了问题,或者输入图像本身雾的浓度和训练集差异太大。一个能用的去雾模型,输出直方图应该比输入有更宽的动态范围。
4. 用自己的数据训练:数据组织、损失权重与训练循环怎么写
4.1 数据组织的两种路线:成对雾图与不成对雾图
训练对偶GAN之前,先认清你手上的数据形态,这决定了整个训练管线怎么写。第一条路线是有成对数据,比如你用大气散射模型人工合成了有雾图像和清晰图像对——这种条件下对偶结构依然可以用循环一致性做辅助监督,但主损失应该用像素级L1或L2。第二条路线是只有不成对的雾图,这是对偶GAN发挥真正价值的情形,G_A只在循环约束下学习去雾。
Pytorch实现时数据加载用torch.utils.data.Dataset,针对不成对数据,我的做法是两个文件夹各自打乱读取,不用保证样本索引对齐。
# dataset.py import os from torch.utils.data import Dataset from PIL import Image class UnpairedDataset(Dataset): def __init__(self, hazy_dir, clear_dir, transform=None): self.hazy_paths = sorted([os.path.join(hazy_dir, f) for f in os.listdir(hazy_dir)]) self.clear_paths = sorted([os.path.join(clear_dir, f) for f in os.listdir(clear_dir)]) self.transform = transform def __len__(self): return max(len(self.hazy_paths), len(self.clear_paths)) def __getitem__(self, idx): hazy_path = self.hazy_paths[idx % len(self.hazy_paths)] clear_path = self.clear_paths[idx % len(self.clear_paths)] hazy_img = Image.open(hazy_path).convert('RGB') clear_img = Image.open(clear_path).convert('RGB') if self.transform: hazy_img = self.transform(hazy_img) clear_img = self.transform(clear_img) return hazy_img, clear_img__len__取两个目录中较大的那个,并用取模让短的目录里的样本被重复使用。这样做的好处是每个epoch里两个域的样本数量都能被完整遍历,不会出现一个域被另一个域的数量“压制”的情况。数据增强上,随机水平翻转是必须的,随机crop建议只在图像尺寸较大的时候做——如果原图本身只有几百像素,再crop下去生成器的感受野都不够覆盖雾的分布。
4.2 损失函数配置:对抗损失、循环一致性损失和身份损失的权重
对偶GAN的总损失是我调参过程中花时间最多的部分。写法上它由三部分组成:生成器的对抗损失、循环一致性损失和身份损失。Pytorch里没有现成的组合模块,需要自己拼。
# train_step.py 摘要,只展示核心损失计算 import torch.nn as nn import torch.nn.functional as F criterion_gan = nn.MSELoss() # LSGAN用MSE替代BCE,训练更稳 criterion_cycle = nn.L1Loss() criterion_identity = nn.L1Loss() loss_gan_A = criterion_gan(disc_A(gen_A(hazy)), real_label) loss_gan_B = criterion_gan(disc_B(gen_B(clear)), real_label) recon_hazy = gen_B(gen_A(hazy)) recon_clear = gen_A(gen_B(clear)) loss_cycle = criterion_cycle(recon_hazy, hazy) + criterion_cycle(recon_clear, clear) loss_id_A = criterion_identity(gen_A(clear), clear) loss_id_B = criterion_identity(gen_B(hazy), hazy) lambda_cycle = 10.0 lambda_id = 0.5 * lambda_cycle loss_G = loss_gan_A + loss_gan_B + lambda_cycle * loss_cycle + lambda_id * (loss_id_A + loss_id_B)这里disc_A(gen_A(hazy))是让“去雾结果”尽量被判别器判为真实清晰图,使用的是LSGAN形式的MSE损失,比传统BCE在训练后期更不容易饱和。lambda_cycle=10.0是CycleGAN论文里验证过的默认值,我沿用至今;lambda_id设为0.5 * lambda_cycle,这是身份损失的关键——它强制生成器在输入已经是清晰图像时不要做过度修改,避免把原本清晰的图“加雾”来骗过循环约束。
身份损失权重太高的副作用也要注意。如果lambda_id超过lambda_cycle,生成器会变得保守,倾向于少改动输入,因为它发现“什么都不做”也能让身份损失为零。这个现象在雾图本身较薄时尤其明显,需要降权重。我自己的调参节奏是先固定lambda_cycle=10,然后跑50个step观察重建图的清晰度,再决定lambda_id往0.5倍方向调还是往0.1倍方向降。
4.3 训练循环与学习率节奏:Pytorch里怎么控制对抗训练的稳定性
对偶GAN训练不稳定是常态,Pytorch本身不会帮你解决这个问题,它只提供优化器和学习率调度的基础能力。我的习惯是生成器和判别器分开两个优化器,生成器学习率2e-4,判别器学习率降一半到1e-4,这样判别器不会“学得太快”反过来碾压生成器。
optimizer_G = torch.optim.Adam( list(gen_A.parameters()) + list(gen_B.parameters()), lr=2e-4, betas=(0.5, 0.999) ) optimizer_D_A = torch.optim.Adam(disc_A.parameters(), lr=1e-4, betas=(0.5, 0.999)) optimizer_D_B = torch.optim.Adam(disc_B.parameters(), lr=1e-4, betas=(0.5, 0.999)) # 学习率前100个epoch保持不变,后100个epoch线性衰减到0 def lambda_rule(epoch): return 1.0 - max(0, epoch - 100) / 100 scheduler_G = torch.optim.lr_scheduler.LambdaLR(optimizer_G, lr_lambda=lambda_rule)betas=(0.5, 0.999)是GAN训练的标准配置,Adam默认的betas=(0.9, 0.999)在GAN里容易产生震荡,0.5能显著压低动量影响。学习率线性衰减是从第101个epoch开始,这样的安排是为了让模型先在较高学习率下找到大致正确的映射方向,后期用低学习率细化纹理。
在判别器更新策略上有一个容易踩的坑:每步都更新判别器和生成器,会让二者陷入“互相追赶”的振荡。常见做法是判别器每步都更新,生成器只在判别器更新后更新,也就是说交替更新而不是同步更新。另外一个实用技巧是每5步把判别器的输入做一次随机擦除:对真实图像或生成的假图像随机遮挡一个小矩形区块,迫使判别器不能只靠局部高频纹理做判断。
5. 避坑记录:去雾模型最常见的四个翻车现场
5.1 现象:去雾结果色偏严重,天空变成灰色色块
这是我第一次跑通对偶GAN去雾后最先踩的坑。模型输出的去雾图整体偏灰暗,尤其天空区域从淡蓝色变成了灰色,像蒙了一层水泥色。原因在于循环一致性损失趋向于“保守重建”——生成器发现把输出图像整体亮度压低,可以让循环重建时的误差变小。判别器在这种低亮度图像上也可能给出模糊的真假判断,导致生成器向色偏方向滑落。
解决的办法是给循环一致性损失增加色彩约束的变体。我没有改损失公式,而是把训练数据里的清晰图像做了预处理:用白平衡算法把色温归一化到相近范围,并让生成器的输出在进入循环之前先经过一次RGB色彩分布对齐。这个预处理配合身份损失的约束,能让色偏问题明显缓解。如果在你的项目里这两种手段还不够,那就直接降低lambda_cycle从10到5,让对抗损失有更多的“发言权”来拉扯生成器不要过分保守。
5.2 现象:训练损失曲线下降正常,但推理时去雾效果几乎没有
训练过程中判别器和生成器的loss都有序下降,看起来非常健康,结果拿一张新雾图去测,输出的图像和输入几乎是同一个,去过雾等于没去。这个现象说到底是“模式坍缩”的变体:生成器发现一个几乎恒等映射也能骗过判别器。为什么能骗过?因为判别器只看到了训练集的清晰图像分布,它会认为这张“几乎没改动”的图虽然不理想,但也不至于太假。
原因在于身份损失的权重设得太大了,生成器从身份损失中获得的梯度远远大于从对抗损失和循环一致性损失中获得的梯度,于是它学到了“最小干预”策略。解决时我会做两个动作:第一,把lambda_id从0.5倍lambda_cycle降到0.1倍,观察生成器是否开始更积极地去除雾的纹理;第二,在训练集中加入少量人工雾图,并显式地把它们标为“有雾域”,人为拉大两个域的分布差异,逼迫生成器做出更多改变。
5.3 现象:训练时好时坏,同一个checkpoint不同epoch表现差距巨大
对偶GAN训练里的“好时坏时”和普通GAN一样常见,典型现象是第80个epoch保存的模型去雾效果不错,但第85个epoch的模型输出直接花掉。这说明对抗训练在后期进入了震荡区,生成器和判别器的loss在交替上升下降。很多人在这时选择骰子式地每隔几个epoch保存一次checkpoint,碰运气挑表现好的,这是可行的,但太被动。
我的做法是引入EMA(指数移动平均)来平滑生成器的权重。Pytorch没有像TensorFlow那样内置EMA,需要手动维护一份权重影子拷贝,每个step把当前生成器权重按一定比例合并进去。推理时用的是EMA版本而不是当前版本,这让生成器的行为更稳定,少受对抗振荡影响。EMA的衰减率我取0.999,推理时基本感觉不到和原版的差异,但稳定性提升是肉眼可见的。
5.4 现象:换数据集后效果骤降,训练集与测试集“长得很像”但结果差了
对偶GAN去雾有个隐藏条件:训练时加雾器G_B学到的“雾的分布”是训练集特有的。如果测试时遇到的是浓雾、夜间雾、带有特殊光照的雾图,G_B没有见过相似的加雾模式,G_A自然也就不知道怎么正确去雾。表现就是训练集上指标很好,换一批测试图就翻车。
这个坑没有完全根治的办法,只能从数据层面缩小差距。我一般会在训练集中混入多组不同浓度的合成雾图,并加一些随机色调偏移,扩大雾域的覆盖范围。另一个可选做法是推理时对输入做多尺度处理:把图像缩放到0.8倍、1.0倍、1.2倍分别去雾,再把结果融合回原尺寸,这种“测试时增强”能一定程度上抵抗域偏移带来的劣化。
6. 进阶验证:评价指标、L1调优与Pytorch转ONNX部署
模型训练完并保存好之后,下一步不是急着部署,而是先做定量评价。去雾任务和超分辨率一样,不能只用“看着顺眼”来验收。常用的客观指标是PSNR和SSIM,但这两个指标在无真实参考图时是用不了的。对于真实雾图,只能做无参考评价,比如计算图像的对比度、饱和度、暗通道先验的残差;对于人工合成雾图,则可以计算去雾结果与原清晰图的PSNR和SSIM。
Pytorch里算PSNR的代码我写成这样:
import torch import torch.nn.functional as F def psnr(img1, img2, max_val=1.0): mse = F.mse_loss(img1, img2) return 10 * torch.log10(max_val * max_val / mse) # 调用时输入应为[0,1]范围的张量 # score = psnr(output, clear_gt).item()如果PSNR正常但SSIM偏低,问题通常出在纹理保留上,去雾过度磨皮了。此时我会把循环一致性损失里的L1换成L1和SSIM损失的组合,SSIM损失需要额外安装pytorch_msssim包,权重设为0.3到0.5。Pytorch里L1正则和SSIM正则方向不一样——L1保住的是像素值,SSIM保住的是结构,两者平衡后才能让输出既清晰又保留细节。
验证通过后的部署环节,现在最常走的路线是转ONNX。Pytorch转ONNX的代码非常短,但有两个关键点要记得处理:一是输入尺寸要固定,ONNX导出时不支持动态高宽(可以设置dynamic_axes,但很多推理引擎对动态形状支持不友好);二是模型里的torch.nn.functional.interpolate在部分ONNX推理后端会导出成奇怪的算子,需要测试兼容性。基本导出代码如下:
model.eval() dummy_input = torch.randn(1, 3, 512, 512).to(device) torch.onnx.export( model, dummy_input, 'defog.onnx', input_names=['input'], output_names=['output'], opset_version=11, do_constant_folding=True )这里opset_version=11是我用的兼容性较好的版本,更高的opset支持更多算子但要求推理引擎更新。导出后一定用onnxruntime验证一遍:
import onnxruntime as ort import numpy as np sess = ort.InferenceSession('defog.onnx') input_data = np.random.randn(1, 3, 512, 512).astype(np.float32) output = sess.run(['output'], {'input': input_data})[0]如果onnxruntime的输出和Pytorch直接推理的输出最大差值超过1e-2,绝对不是正常浮动误差,需要检查模型里是否有自定义算子或训练专用的dropout层没关干净。
以上这套流程走完,方向基本就稳了。我在最初用对偶GAN做去雾时吃了不少亏,最大的教训是:不要迷信损失函数“更新鲜”就有效,先跑通一条基线,再在评价指标的引导下逐步加料。希望帮到你。
本文还有配套的精品资源,点击获取