news 2026/9/8 23:55:35

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

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
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, transform=None): 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_tensor

3.2 生成器与判别器代码解析

生成器我封装了一个UNetBlock类,每个下采样和上采样块都作为独立模块,代码可读性更高。关键的跳跃连接实现如下:

class UNetGenerator(nn.Module): def __init__(self, in_channels=3, out_channels=3, ngf=64): super().__init__() # 编码器:逐步下采样,通道数依次变为64、128、256、512 self.down1 = self._block(in_channels, ngf, norm=False) # 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_size=4, stride=2, padding=1), 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], dim=1)) u2 = self.up2(torch.cat([u1, d2], dim=1)) u3 = self.up3(torch.cat([u2, d1], dim=1)) out = self.up4(torch.cat([u3, x], dim=1)) return out

判别器使用PatchGAN,输入是6通道的拼接图,输出是16×16的patch特征图。这里有一个细节需要注意,网络最后一层卷积不使用激活函数,因为LSGAN的判别器输出需要的是原始分数,而不是经过sigmoid的概率值,训练时直接和0/1目标计算均方误差。

class PatchDiscriminator(nn.Module): def __init__(self, in_channels=3, ndf=64): super().__init__() self.model = nn.Sequential( nn.Conv2d(in_channels*2, ndf, 4, 2, 1), nn.LeakyReLU(0.2, inplace=True), nn.Conv2d(ndf, ndf*2, 4, 2, 1), nn.InstanceNorm2d(ndf*2), nn.LeakyReLU(0.2, inplace=True), nn.Conv2d(ndf*2, ndf*4, 4, 2, 1), nn.InstanceNorm2d(ndf*4), nn.LeakyReLU(0.2, inplace=True), 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)先生成有雾图,其中t=exp(-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-4,beta1取0.5,beta2取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和SSIM,PSNR关注像素级的重建误差,SSIM关注结构相似性。对偶GAN在同分辨率情况下PSNR能达到22到25dB,SSIM在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_memory=True,数据加载瓶颈对整体速度的影响在数据量大时非常大。

最后再分享一个节省时间的经验:训练过程中每5个epoch就把生成器的输出图保存到本地文件夹,肉眼观察去雾效果的演变。损失曲线只能反映数值变化,很多问题从图像上就能直接看出来,比盯着loss省钱省力得多。

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

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/8 23:55:14

AI画板实测:GPT-6 Astra在原理图与PCB设计中的能力与局限

把同一个电源域的电容分两排放在芯片两侧,结果回流路径被拉得很长,纹波指标差了30%。这种问题AI不一定能看出来,但要靠它把所有细节都安排到位,现阶段还不现实。哪些可以放心交给AI适合让GPT-6 Astra处理的,是那些“规…

作者头像 李华
网站建设 2026/9/8 23:52:26

RS-485收发器选型实测:从MAX485到THVD1550,8款芯片横向对比

去年秋天,我接手的一批无刷电机控制器在客户车间的配电柜合闸瞬间,总线上挂着的半双工RS-485收发器一片接一片被打穿。上位机一直报通信超时,现场测A、B线对地电阻,好几块板子只有几十欧姆,拆下来看,清一色…

作者头像 李华
网站建设 2026/9/8 23:51:25

PyTorch实战:基于CRNN+CTC的车牌识别全流程解析

1. 为什么把车牌识别当作 PyTorch 实战项目1.1 车牌识别看起来简单,实际上卡在哪儿先说一个反直觉的现象:车牌识别在工程里看起来特别成熟,门口停车场、高速收费口都在用,似乎是个“老掉牙”的需求。但当我自己把任务拆开&#xf…

作者头像 李华
网站建设 2026/9/8 23:51:11

AI网关全面测评:从API网关到MAI Gateway的七类方案对比

最近大半年,只要聊到 AI 应用落地,迟早会碰到一个绕不开的基础设施话题——AI 网关。模型越来越多,调用协议五花八门,团队既要接 OpenAI 又要兼容国产模型,还要控成本、管权限、做审计,这时候靠业务代码一层…

作者头像 李华
网站建设 2026/9/8 23:46:05

快手接口签名sig、sig3与NStoken原理及测试用例详解

简介:快手sig3、sig、NStoken算法资源面向移动应用开发者、爬虫及逆向分析人员,聚焦快手接口交互中的请求签名与身份令牌生成问题。压缩包共4个文件,包含2个Python脚本和2个数据文件(data0、data1),脚本分别…

作者头像 李华