FEATURED · 精选文章

PyTorch实现对偶GAN图像去雾:从原理到工程实战

发布时间 / 2026/9/8 23:55:42
来源 / 创域科博编辑部
栏目 / 资讯中心
PyTorch实现对偶GAN图像去雾:从原理到工程实战 简介基于PyTorch实现图像去雾的对偶生成对抗网络是一个包含完整Python源码、项目说明及详细代码注释的毕业设计项目。项目针对雾气导致图像对比度下降、细节丢失等问题利用生成器与判别器相互对抗的方式恢复清晰无雾图像适合计算机相关专业学生用于毕业设计、课程设计或深度学习实战练习。压缩包共26个文件其中10个py文件涵盖生成器、判别器、训练脚本及预测脚本2个pkl文件为预训练模型权重6个png与5个jpg分别提供测试图像和去雾效果图另有Markdown说明文档便于快速上手整体压缩包大小约21.23MB。目前已有67人浏览学习。项目源自经导师指导并通过的高分毕业设计所有代码均经过严格调试确保可运行。通过学习该项目读者既能掌握生成对抗网络的对偶训练机制也能了解暗通道先验等图像去雾知识并快速搭建自己的去雾模型是兼具学术参考价值与工程实用性的优质资源。 雾天拍出来的照片灰蒙蒙一片对比度和色彩全被压住了这对目标检测、语义分割这类下游视觉任务来说几乎是灾难。我这次折腾的项目就是用PyTorch搭一个基于对偶生成对抗网络DualGAN的图像去雾模型不依赖传统物理模型去估计透射率和大天气光而是让网络直接学习“有雾图→无雾图”的端到端映射整套代码包含完整的网络结构、训练脚本和详细注释。如果你正在做图像复原、图像翻译或者想搞清楚对偶GAN的循环一致性损失在实际任务里怎么落地这篇东西应该能帮你省不少时间。我最初以为去雾和普通图像增强差不多真正动手才发现坑不少成对的有雾/无雾训练数据很难拿、直接套用普通GAN又容易出现颜色偏移和伪影、训练过程动不动就不收敛。这篇文章会把我的整体设计思路、网络核心结构、关键代码实现以及训练调参时踩过的坑全部整理出来偏向工程实操可以直接照着复现。1. 项目背景与整体设计思路1.1 为什么选择对偶GAN做图像去雾图像去雾的主流路线大致分两类。一类是传统物理模型方法最典型的是暗通道先验通过估计大气光和透射率来反演清晰图像这类方法在某些场景下效果稳定但遇到天空、白色物体这类不符合暗通道假设的区域时容易出现色块和光晕。另一类是深度学习方法早期用CNN去回归透射率本质还是在物理模型框架里打转后来GAN被引入直接做有雾到无雾的图像翻译跳过了中间物理量的估计效果上限更高。我选择对偶GAN的核心原因是它解决了训练数据的问题。真实场景里想采集严格配对的同一场景有雾和无雾图像难度极高一般只能用合成数据。普通有监督GAN要求输入输出成对数据集制作成本大。而对偶GAN用循环一致性损失只要求两个域的图像集合不要求像素级配对这大大放宽了数据限制。雾天图像和无雾图像可以分别从不同来源收集网络自己学习两个域之间的双向映射同时保证转换后的图像能再转回来结构信息不会在翻译过程中丢失。1.2 项目架构与文件组织整个项目的代码组织比较清晰我分了几个模块models.py定义生成器和判别器网络结构。dataset.py加载有雾/无雾图像数据集做归一化和随机裁剪。train.py训练主流程包含对偶训练、损失计算和模型保存。inference.py加载训练好的模型对单张图片或文件夹去雾。这样做的好处是每个模块职责单一调试时只改对应文件就行。我在训练脚本里把整个对偶循环都做了详细注释关键张量的维度变化也标了出来方便理解每一步在做什么。1.3 环境准备与依赖安装PyTorch环境这一块确实容易卡住新手尤其是GPU版本的安装。我这次用的是Python 3.10 PyTorch 2.0 CUDA 11.8的组合训练一张256x256的图片单卡显存占用大概在4GB左右普通消费级显卡都能跑。安装时可以直接用官方提供的pip命令pip install torch torchvision --index-url https://download.pytorch.org/whl/cu118如果下载速度慢可以换阿里云或清华的镜像源或者使用离线whl包安装。CPU版本也能运行只是训练速度会慢好几倍前期验证代码逻辑是够用的。2. 网络结构与损失函数设计2.1 生成器U-Net结构生成器我选用了U-Net结构而非单纯的编码器-解码器。原因在于图像去雾属于像素级翻译任务需要输出的图像在边缘、纹理等高频细节上和输入保持一致。U-Net在下采样提取高层语义特征的同时通过跳跃连接把浅层的细节特征直接传递到解码器对应层相当于给解码器开了一路“直达通道”恢复出来的图像纹理更清晰。我的实现里下采样采用4层卷积每层卷积步长为2实现空间尺寸减半通道数从64逐层倍增到512。归一化层使用了InstanceNorm而不是BatchNorm这是图像翻译任务里一个很重要的细节。BatchNorm在小batch size下统计量不稳定而且会引入batch内图像的相互影响而InstanceNorm对单张图像的通道做归一化更符合单图风格转换的场景。解码器部分使用转置卷积上采样通道数逐层减半每一层都先和编码器对应层的输出做channel维度的拼接再接卷积和ReLU。最后一层用Tanh激活把输出限制在[-1,1]区间和输入图像的归一化范围保持一致。2.2 判别器PatchGAN判别器我采用了PatchGAN。它和传统GAN判别器输出一个标量真/假不同PatchGAN输出的是一张N×N的特征图特征图上每个像素对应输入图像的一个局部patch分别判断该patch是否为真实图像。以70×70的PatchGAN为例它的有效感受野是70×70强制判别器关注局部纹理和结构而不是只看整张图的整体分布。这个设计对去雾任务尤其关键。有雾和无雾图像在全局色调上可能相近差异主要体现在局部细节和纹理清晰度上PatchGAN能更精细地监督这种局部差异避免生成器通过改变整体颜色蒙混过关。判别器网络我采用3层卷积输入是有雾图A和生成图或真实图拼接成的6通道张量输出patch特征图的每个像素值代表对应patch的真假置信度。训练时使用最小二乘损失LSGAN代替标准二分类交叉熵它的梯度更平滑训练也更稳定生成的图像质量比普通GAN更高。2.3 对偶一致性损失与整体目标函数整个对偶GAN的损失函数由三部分组成。对抗损失让生成器学会生成以假乱真的目标域图像判别器学会区分真假。这部分定义了两个生成器G_AB有雾→无雾和G_BA无雾→有雾各自的对抗损失。循环一致性损失是对偶GAN的核心。它的直觉是如果把一张有雾图A先转成无雾图再通过反向生成器转回来得到的重建图应该和原始图片A尽可能接近。这个约束保证了转换过程保留了原图的结构和内容防止生成器随意发挥。另外我还加了一项身份损失把无雾图直接喂给G_AB要求输出还是无雾图本身。这个损失的作用是保持颜色和色调稳定防止生成器在去雾过程中引入不必要的颜色偏移。整体目标函数就是这三项损失的加权和循环一致性损失权重取10身份损失权重取5对抗损失权重取1。3. 核心代码实现与详细注释3.1 数据加载与预处理数据部分我写了一个继承自torch.utils.data.Dataset的类分别加载有雾图文件夹和无雾图文件夹通过索引对应关系配对。如果两个文件夹数量不一致用取模的方式循环配对虽然这样会产生一些不对齐的配对但对偶GAN恰好不要求严格配对所以不影响训练。预处理关键是随机裁剪到256×256、随机水平翻转、归一化到[-1,1]。我特别把归一化放在最后一步先做几何增强再做数值归一化这样避免翻转时数值统计出错。完整的数据加载代码如下class DehazeDataset(Dataset): def __init__(self, hazy_dir, clean_dir, transformNone): self.hazy_paths glob.glob(os.path.join(hazy_dir, *.png)) self.clean_paths glob.glob(os.path.join(clean_dir, *.png)) self.transform transform def __len__(self): return max(len(self.hazy_paths), len(self.clean_paths)) def __getitem__(self, idx): hazy_img Image.open(self.hazy_paths[idx % len(self.hazy_paths)]).convert(RGB) clean_img Image.open(self.clean_paths[idx % len(self.clean_paths)]).convert(RGB) # 随机水平翻转 if torch.rand(1) 0.5: hazy_img hazy_img.transpose(Image.FLIP_LEFT_RIGHT) clean_img clean_img.transpose(Image.FLIP_LEFT_RIGHT) hazy_tensor self._to_tensor(hazy_img) # 归一化到[-1,1] clean_tensor self._to_tensor(clean_img) return hazy_tensor, clean_tensor3.2 生成器与判别器代码解析生成器我封装了一个UNetBlock类每个下采样和上采样块都作为独立模块代码可读性更高。关键的跳跃连接实现如下class UNetGenerator(nn.Module): def __init__(self, in_channels3, out_channels3, ngf64): super().__init__() # 编码器逐步下采样通道数依次变为64、128、256、512 self.down1 self._block(in_channels, ngf, normFalse) # 256x256 self.down2 self._block(ngf, ngf*2) # 128x128 self.down3 self._block(ngf*2, ngf*4) # 64x64 self.down4 self._block(ngf*4, ngf*8) # 32x32 # 解码器转置卷积上采样与编码器输出拼接 self.up1 self._up_block(ngf*8 ngf*4, ngf*4) # 64x64 self.up2 self._up_block(ngf*4 ngf*2, ngf*2) # 128x128 self.up3 self._up_block(ngf*2 ngf, ngf) # 256x256 self.up4 nn.Sequential( nn.ConvTranspose2d(ngf in_channels, out_channels, kernel_size4, stride2, padding1), nn.Tanh() ) def forward(self, x): d1 self.down1(x) d2 self.down2(d1) d3 self.down3(d2) d4 self.down4(d3) u1 self.up1(torch.cat([d4, d3], dim1)) u2 self.up2(torch.cat([u1, d2], dim1)) u3 self.up3(torch.cat([u2, d1], dim1)) out self.up4(torch.cat([u3, x], dim1)) return out判别器使用PatchGAN输入是6通道的拼接图输出是16×16的patch特征图。这里有一个细节需要注意网络最后一层卷积不使用激活函数因为LSGAN的判别器输出需要的是原始分数而不是经过sigmoid的概率值训练时直接和0/1目标计算均方误差。class PatchDiscriminator(nn.Module): def __init__(self, in_channels3, ndf64): super().__init__() self.model nn.Sequential( nn.Conv2d(in_channels*2, ndf, 4, 2, 1), nn.LeakyReLU(0.2, inplaceTrue), nn.Conv2d(ndf, ndf*2, 4, 2, 1), nn.InstanceNorm2d(ndf*2), nn.LeakyReLU(0.2, inplaceTrue), nn.Conv2d(ndf*2, ndf*4, 4, 2, 1), nn.InstanceNorm2d(ndf*4), nn.LeakyReLU(0.2, inplaceTrue), nn.Conv2d(ndf*4, 1, 4, 1, 1) # 输出16x16特征图 )3.3 训练循环对偶行为的关键实现训练循环是整个项目最核心的部分。每一步迭代需要依次完成四个网络的更新两个生成器和两个判别器并且要正确构建循环一致性损失的计算图。我定义了两个优化器分别管理AB域和BA域的生成器与判别器参数这样当一个生成器更新时不会误更新另一个生成器的梯度。具体流程是先固定所有生成器更新判别器D_B和D_A然后固定判别器更新生成器G_AB和G_BA。更新生成器时同时计算对抗损失、循环一致性损失和身份损失累加起来做反向传播。for epoch in range(num_epochs): for i, (hazy, clean) in enumerate(train_loader): hazy hazy.to(device) clean clean.to(device) # 前向计算 fake_clean netG_AB(hazy) # 有雾 - 去雾 fake_hazy netG_BA(clean) # 无雾 - 增雾 rec_hazy netG_BA(fake_clean) # 重建有雾图 rec_clean netG_AB(fake_hazy) # 重建无雾图 # 更新判别器D_B判断真实无雾图与生成无雾图 pred_real netD_B(clean, hazy) pred_fake netD_B(fake_clean.detach(), hazy) loss_D_B 0.5 * (F.mse_loss(pred_real, real_label) F.mse_loss(pred_fake, fake_label)) # 更新判别器D_A判断真实有雾图与生成有雾图 pred_real netD_A(hazy, clean) pred_fake netD_A(fake_hazy.detach(), clean) loss_D_A 0.5 * (F.mse_loss(pred_real, real_label) F.mse_loss(pred_fake, fake_label)) # 更新生成器 loss_cycle (L1_loss(rec_hazy, hazy) L1_loss(rec_clean, clean)) * 10.0 loss_identity (L1_loss(netG_AB(clean), clean) L1_loss(netG_BA(hazy), hazy)) * 5.0 loss_gan_AB F.mse_loss(netD_B(fake_clean, hazy), real_label) loss_gan_BA F.mse_loss(netD_A(fake_hazy, clean), real_label) loss_G loss_gan_AB loss_gan_BA loss_cycle loss_identity optimizer_G.zero_grad() loss_G.backward() optimizer_G.step()这里有几个容易出错的地方。第一计算判别器损失时传入判别器的生成图像需要使用detach()切断梯度否则梯度会反向传播到生成器导致判别器和生成器同时更新训练过程会非常不稳定。第二循环一致性损失中的L1损失比L2损失效果更好L1损失对异常像素不那么敏感重建图像更锐利不会出现L2容易产生的模糊问题。4. 训练实践与去雾效果4.1 数据集选择与合成策略数据集方面我使用了公开的RESIDE数据集的子集里面包含合成有雾图像和对应的清晰图像。如果找不到现成的配对数据也可以用NYU Depth V2深度数据集自己合成雾图。合成方法很简单把深度图归一化后作为透射率t在随机选取的大气光A下用大气散射模型I Jt A(1-t)先生成有雾图其中texp(-beta*d)beta在[0.5,1.5]之间随机取值模拟不同浓度的雾。这种合成策略成本低能灵活控制雾的浓度而且可以大量生成训练样本。但要注意合成雾和真实雾之间存在域差距训练出来的模型在真实雾天图像上效果会打折扣。所以训练时可以适当加入少量真实雾天图像做微调或者使用风格迁移的思路让模型适应真实雾的分布。训练时batch size建议设置在4到8之间。PatchGAN在小batch size下也能稳定训练我实测batch size为4时256×256分辨率下RTX 3060显卡能跑到每秒2.5次迭代训练100个epoch大约需要4到6小时。4.2 超参数配置与调优经验学习率设置上我使用了Adam优化器初始学习率2e-4beta1取0.5beta2取0.999。beta1取0.5而不是默认的0.9是GAN训练的惯例因为0.5能更快地遗忘历史梯度减少训练震荡。前50个epoch保持学习率不变后50个epoch线性衰减到0这种做法能让模型早期快速探索后期精细收敛。标签平滑也是我比较推荐的操作。把真实标签从1替换成0.9到1之间的随机值把假标签从0替换成0到0.1之间的随机值可以降低判别器过度自信避免生成器梯度消失。这在提升图像质量方面的效果非常明显模型不容易出现模式坍塌。还有一个容易被忽略的点是图像分辨率。如果显卡显存不够不要硬撑256×256可以先用128×128训练训练结束后再用256×256微调几十个epoch。逐步增大分辨率的策略能让模型先学全局结构再学细节。4.3 去雾结果评估训练完成后我一般会同时看客观指标和主观效果。客观指标主要看PSNR和SSIMPSNR关注像素级的重建误差SSIM关注结构相似性。对偶GAN在同分辨率情况下PSNR能达到22到25dBSSIM在0.88到0.93之间虽然比不上专门的物理模型方法但图像观感更自然没有明显的颜色畸变和光晕。在做主观评估时我发现一个有趣的现象网络在薄雾区域的去雾效果最好在浓雾区域会有一定程度的过曝或细节丢失这和训练数据中浓雾样本占比少有关。后来我通过增加浓雾样本的采样权重把这个现象缓解了不少。所以如果训练效果在某类场景下不理想优先检查训练数据分布是否均衡。5. 常见问题与排查技巧实录5.1 训练不收敛或模式坍塌这是最难排查的问题之一。表现是损失震荡剧烈或者生成器只输出单一色调的图像。我遇到过一次模式坍塌原因是两个判别器的学习率设置比生成器高判别器训练过快导致生成器梯度消失。解决方法是降低判别器学习率或者每训练一次判别器训练两次生成器让两边保持均衡。另外一个常见原因是初始化不当。我建议对生成器最后几层做正态分布初始化均值0标准差0.02而不是用默认的均匀分布。经验上合理的初始化能显著加快收敛速度。5.2 去雾后图像颜色偏移或出现伪影颜色偏移经常出在身份损失权重太小或者训练数据里干净图像本身色调分布不均衡的情况下。如果模型去雾后偏绿先检查无雾训练集是不是以绿色植物场景为主是的话就要做颜色空间的数据增强比如随机调整HSV通道让模型不依赖特定色调。伪影则大多和PatchGAN的感受野设置有关。70×70的PatchGAN对高频纹理敏感但如果伪影是块状或条状的可以把PatchGAN从70×70换成140×140扩大判别器的感知范围强制生成器在更大尺度上保持一致性。代价是显存占用增加两种方案权衡一下就行。5.3 显存不足与训练速度慢显存不足最有效的解决办法是减小batch size和图像裁切尺寸其次是使用梯度累积等效增大batch size而不增加显存。速度慢的话优先确认PyTorch是否真的用上了GPU训练过程中可以用nvidia-smi查看显卡占用。还有一个小技巧是在DataLoader中设置num_workers大于0并开启pin_memoryTrue数据加载瓶颈对整体速度的影响在数据量大时非常大。最后再分享一个节省时间的经验训练过程中每5个epoch就把生成器的输出图保存到本地文件夹肉眼观察去雾效果的演变。损失曲线只能反映数值变化很多问题从图像上就能直接看出来比盯着loss省钱省力得多。本文还有配套的精品资源点击获取
RELATED — 相关阅读

相关资讯

LATEST — 最新资讯

最新发布

TODAY — 本日精选

新闻

WEEKLY — 本周精选

新闻

MONTHLY — 本月精选

新闻