news 2026/10/1 19:30:57

基于PyTorch的对偶生成对抗网络图像去雾实战:从原理到源码解析

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
基于PyTorch的对偶生成对抗网络图像去雾实战:从原理到源码解析

简介:这份资源是面向计算机相关专业毕业设计、课程设计及期末作业场景的PyTorch实战项目,核心任务是用对偶生成对抗网络完成图像去雾。项目由生成器与判别器双网络协同训练,配套训练、预测、参数解析、数据加载与可视化等模块,适合已具备一定深度学习基础、希望借完整项目提升工程能力的学习者。压缩包共31个文件,约21.31MB,以10个Python源码为主,另有png与jpg格式的测试图像、pkl模型权重、zbak备份文件及README说明文档,覆盖从数据读取、模型定义到推理输出的完整链路。目前已有43人学习下载。代码经过系统测试与反复调试,运行稳定性与功能完整性均有保障,读者可据此理解对偶GAN的去雾原理、网络结构设计与训练流程,并在此基础上完成二次开发或论文撰写,也可作为课程作业的参考实现。

1. 从一张雾天街景说起:对偶生成对抗网络去雾到底在做什么

去年冬天帮朋友处理一批高速公路监控截图,画面里车牌和车道线全被雾吞掉,传统暗通道先验一上,天空区域直接变成色块,暗部细节也跟着糊成一团。那批图最后是用一套基于 PyTorch 的对偶生成对抗网络图像去雾系统救回来的,这也是我后来反复给别人讲这套方案的原因。它要解决的核心问题很具体:单张雾图输入,输出去雾后的清晰图,而且不需要成对的「有雾-无雾」训练数据。对偶生成对抗网络(Dual GAN)的思路是同时训练两个方向的映射,一个有雾到无雾,一个无雾到有雾,再用循环一致性约束把两边锁住,这样即使手里只有一堆无标注的雾图和一堆清晰图,也能把模型训起来。适合谁?手里有监控、遥感、车载摄像头这类真实雾天数据、又拿不到配对标签的工程师,以及想用 PyTorch 把 GAN 去雾从论文跑到自己数据上的开发者。源码解析的价值不在于逐行读,而在于搞清楚每个模块为什么这么搭、参数为什么这么设。

2. 对偶生成对抗网络去雾的原理与 PyTorch 选型理由

2.1 为什么是「对偶」而不是普通 GAN

普通 GAN 去雾要么依赖成对数据做监督,要么用单边映射加判别器硬扛,前者数据难搞,后者容易在天空、白墙这类高频区域生成伪纹理。对偶结构的关键在于两个生成器和两个判别器:生成器 G 负责雾到清晰,生成器 F 负责清晰到雾,判别器 D_Y 判断「这张清晰图是不是真的」,D_X 判断「这张雾图是不是真的」。损失由三块组成——对抗损失让生成结果逼近目标域分布,循环一致性损失保证 G(F(y))≈y、F(G(x))≈x,身份损失在部分实现里进一步稳住颜色。这样训出来的 G,即使没见过某张雾图对应的清晰版本,也能靠循环约束把结构保住。

我第一次跑通这套结构时最直观的感受是:循环一致性权重给低了,去雾图会偏色;给高了,去雾力度又不够,雾是淡了但对比度上不来。这个权衡后面会细说。

2.2 PyTorch 在这套系统里的实际优势

选 PyTorch 不是跟风。对偶 GAN 训练时要频繁在生成器、判别器之间切换梯度计算,还要对同一批数据做两次前向(一次算循环损失),PyTorch 的动态图让这种「边跑边改计算路径」的写法非常自然。另外自定义循环一致性损失、感知损失、颜色损失时,直接写 Python 函数加 autograd 就行,不用像静态图那样先搭占位符。环境搭建上,pytorch安装和pytorch环境搭建是绕不开的第一步,我一般推荐 conda 建独立环境,再按python和pytorch版本对应关系装对应 CUDA 版本,避免cuda pytorch下载装错导致 GPU 用不上。

2.3 最小可跑的训练骨架

下面这段是训练循环的核心骨架,去掉了日志和保存逻辑,保留对偶 GAN 最关键的四次前向和损失回传。

import torch import torch.nn as nn # G: 雾->清晰, F: 清晰->雾, D_X/D_Y: 对应域判别器 G, F = Generator(), Generator() D_X, D_Y = Discriminator(), Discriminator() opt_G = torch.optim.Adam(list(G.parameters()) + list(F.parameters()), lr=2e-4, betas=(0.5, 0.999)) opt_D = torch.optim.Adam(list(D_X.parameters()) + list(D_Y.parameters()), lr=2e-4, betas=(0.5, 0.999)) criterion_gan = nn.MSELoss() # LSGAN 比原始 GAN 稳 criterion_cyc = nn.L1Loss() # 循环一致性用 L1,边缘更锐 lambda_cyc = 10.0 # 循环损失权重,经验值 10 起步 for haze, clear in dataloader: haze, clear = haze.cuda(), clear.cuda() # ---- 生成器一步 ---- fake_clear = G(haze) # 雾 -> 清晰 rec_haze = F(fake_clear) # 再变回雾,用于循环约束 fake_haze = F(clear) # 清晰 -> 雾 rec_clear = G(fake_haze) # 再变回清晰 loss_cyc = criterion_cyc(rec_haze, haze) + criterion_cyc(rec_clear, clear) loss_gan_G = criterion_gan(D_Y(fake_clear), torch.ones_like(D_Y(fake_clear))) \ + criterion_gan(D_X(fake_haze), torch.ones_like(D_X(fake_haze))) loss_G = loss_gan_G + lambda_cyc * loss_cyc opt_G.zero_grad() loss_G.backward() opt_G.step() # ---- 判别器一步 ---- loss_D = criterion_gan(D_Y(clear), torch.ones_like(D_Y(clear))) \ + criterion_gan(D_Y(fake_clear.detach()), torch.zeros_like(D_Y(fake_clear))) \ + criterion_gan(D_X(haze), torch.ones_like(D_X(haze))) \ + criterion_gan(D_X(fake_haze.detach()), torch.zeros_like(D_X(fake_haze))) opt_D.zero_grad() loss_D.backward() opt_D.step()

逻辑说明:生成器这一步同时算了两条循环路径,rec_haze和rec_clear分别对应两个方向的循环一致性;判别器这一步对真假样本各算一次,detach()是关键,防止判别器梯度回传到生成器。参数说明:lambda_cyc控制循环约束强度,太小去雾不彻底,太大颜色失真;lr=2e-4配合betas=(0.5,0.999)是 CycleGAN 系列常用的稳定组合,比默认的 0.9 更适合 GAN 训练。

3. 从零搭一套可复现的去雾训练流程

3.1 数据准备与不成对采样

对偶 GAN 不需要配对数据,但需要两个域各自的图片。我一般把雾图放data/hazy/,清晰图放data/clear/,用两个独立的 DataLoader 分别采样,每个 batch 里雾图和清晰图互不对应,这正是对偶结构的用武之地。图片统一 resize 到 256×256 或 286×286 再随机裁剪到 256,太大显存吃不消,太小细节丢失严重。

from torch.utils.data import DataLoader from torchvision import transforms from torchvision.datasets import ImageFolder tf = transforms.Compose([ transforms.Resize(286), transforms.RandomCrop(256), transforms.RandomHorizontalFlip(), transforms.ToTensor(), transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5)), # 归一化到 [-1,1] ]) hazy_ds = ImageFolder('data/hazy', transform=tf) clear_ds = ImageFolder('data/clear', transform=tf) hazy_loader = DataLoader(hazy_ds, batch_size=1, shuffle=True, num_workers=4) clear_loader = DataLoader(clear_ds, batch_size=1, shuffle=True, num_workers=4)

逻辑说明:两个 loader 独立 shuffle,保证每个 step 拿到的雾图和清晰图不是同一场景。参数说明:batch_size=1是对偶 GAN 的常见选择,因为一个 step 要跑四次前向,显存占用是普通 GAN 的两倍左右;Normalize到 [-1,1] 是为了配合生成器输出层的 tanh 激活。

3.2 生成器与判别器的结构选择

生成器我一般用 ResNet 风格的编码器-解码器:编码器下采样两次,中间堆 6 到 9 个残差块,解码器用转置卷积或上采样加卷积恢复分辨率。判别器用 PatchGAN,输出 70×70 的感受野判别图,比全局判别器更能抓住局部纹理。这套结构在去雾任务上比 U-Net 直连更稳,因为残差块保留了低频结构信息,雾的去除主要发生在高频细节上。

class ResidualBlock(nn.Module): def __init__(self, dim): super().__init__() self.block = nn.Sequential( nn.ReflectionPad2d(1), nn.Conv2d(dim, dim, 3), nn.InstanceNorm2d(dim), # InstanceNorm 比 BatchNorm 更适合风格迁移类任务 nn.ReLU(inplace=True), nn.ReflectionPad2d(1), nn.Conv2d(dim, dim, 3), nn.InstanceNorm2d(dim), ) def forward(self, x): return x + self.block(x) # 残差连接,稳住梯度

逻辑说明:ReflectionPad2d避免边缘出现黑框,InstanceNorm2d在 batch 很小时比 BatchNorm 稳定得多。参数说明:残差块数量 6 适合 256 分辨率,9 适合 512,再多收益递减还容易过拟合。

3.3 训练参数与显存控制

训练这套系统,显存是第一个拦路虎。一个 step 四次前向加两次反向,256 分辨率下 8GB 显存基本是底线。如果显存不够,常见做法是把残差块减到 4 个、batch 保持 1、开启混合精度。学习率用 2e-4,前 10 个 epoch 保持恒定,之后线性衰减到 0。判别器更新频率可以设成生成器的 1 倍,也可以每两步更新一次生成器,后者在早期更稳。

# 混合精度训练片段,省显存约 30% python train.py --amp --batch 1 --res_blocks 6 --lr 2e-4 --lambda_cyc 10 --epochs 200

逻辑说明:--amp开启自动混合精度,--lambda_cyc对应前面损失里的循环权重。参数说明:--epochs 200是经验值,100 epoch 左右去雾效果开始稳定,200 之后提升有限但颜色会更自然。

4. 源码解析:损失函数、判别器与推理脚本的关键细节

4.1 循环一致性损失的实现差异

很多开源实现里循环损失直接写L1(F(G(x)), x),但实际训练时如果两个方向权重一样,清晰到雾那个方向往往学得慢,因为雾的生成比去雾简单。我一般给两个方向分别设权重,去雾方向权重 10,加雾方向权重 5,这样生成器更关注去雾质量。另外循环损失可以加一个 SSIM 项,边缘保持会更好,但计算开销增加约 15%。

def cycle_loss(rec, real, ssim_weight=0.0): l1 = nn.L1Loss()(rec, real) if ssim_weight > 0: # 简化 SSIM,实际用 pytorch-ssim 或自己实现 ssim = 1 - ssim_loss(rec, real) return l1 + ssim_weight * ssim return l1

逻辑说明:SSIM 项在去雾任务里主要保边缘,ssim_weight一般设 0.1 到 0.3,太大反而让颜色偏灰。参数说明:如果数据里天空占比高,SSIM 权重可以调低,避免天空区域被过度平滑。

4.2 判别器的感受野与 PatchGAN 输出

PatchGAN 判别器输出的是 N×N 的判别图,每个点对应原图一个感受野。70×70 是常用配置,对应 5 层卷积。如果去雾后出现网格状伪影,多半是判别器感受野太小,可以加到 5 层以上或改用多尺度判别器。源码里判别器最后一层不加 sigmoid,配合 LSGAN 的 MSE 损失,训练更稳。

class PatchDiscriminator(nn.Module): def __init__(self, in_ch=3, ndf=64, n_layers=3): super().__init__() layers = [nn.Conv2d(in_ch, ndf, 4, 2, 1), nn.LeakyReLU(0.2, True)] mult = 1 for i in range(1, n_layers): prev, mult = mult, min(2 ** i, 8) layers += [nn.Conv2d(ndf * prev, ndf * mult, 4, 2, 1), nn.InstanceNorm2d(ndf * mult), nn.LeakyReLU(0.2, True)] layers += [nn.Conv2d(ndf * mult, 1, 4, 1, 1)] # 输出 1 通道判别图 self.model = nn.Sequential(*layers) def forward(self, x): return self.model(x)

逻辑说明:n_layers=3对应约 70×70 感受野,加到 4 层感受野更大但参数增多。参数说明:ndf=64是通道基数,显存紧张可以降到 32,判别能力会弱一些但训练更快。

4.3 推理脚本与 ONNX 导出

训练完的生成器 G 单独拿出来做推理,输入一张雾图,输出清晰图。推理时不需要判别器和 F,所以可以把 G 单独保存成G_final.pth。如果要在 C++ 或移动端部署,常见做法是pytorch转onnx,导出时注意固定输入尺寸,动态轴只保留 batch 维。

G.load_state_dict(torch.load('G_final.pth')) G.eval() with torch.no_grad(): out = G(haze_tensor) # haze_tensor: 1x3x256x256, 归一化到 [-1,1] out = (out * 0.5 + 0.5).clamp(0, 1) # 反归一化到 [0,1]

逻辑说明:eval()关掉 InstanceNorm 的训练态统计,no_grad省显存。参数说明:反归一化系数 0.5 对应前面 Normalize 的均值和方差,如果训练时用了别的归一化参数,这里要同步改。

5. 避坑与排查:对偶 GAN 去雾训练中最容易翻车的五件事

5.1 去雾图整体偏蓝或偏黄

现象:训练几十个 epoch 后,输出图颜色明显偏离真实场景,天空发紫或地面发黄。原因:循环一致性权重过高,生成器为了满足循环约束牺牲了颜色保真;或者判别器太强,生成器被迫生成「讨判别器喜欢」的色调。解决:把lambda_cyc从 10 降到 5 到 7,同时给判别器加标签平滑(真样本标签用 0.9 而不是 1.0),削弱判别器优势。

5.2 训练中期损失突然爆炸

现象:前 20 个 epoch 正常,之后生成器损失或判别器损失突然飙到几百甚至 NaN。原因:学习率没衰减,或者判别器更新太快导致梯度爆炸。解决:加梯度裁剪torch.nn.utils.clip_grad_norm_(G.parameters(), 1.0),学习率在第 30 个 epoch 后线性衰减,判别器每两步更新一次。

5.3 去雾后细节全丢,像被磨皮

现象:雾是去掉了,但车牌、树枝、文字全糊成一片。原因:循环损失只用 L1,生成器倾向于输出平滑结果;或者判别器感受野太小,只关注局部颜色不关注纹理。解决:循环损失加 SSIM 项,判别器加到 4 层或改用多尺度判别器,训练数据里增加纹理丰富的样本。

5.4 显存不够,batch 只能设 1 还 OOM

现象:8GB 显存跑 256 分辨率,batch=1 仍然报 CUDA out of memory。原因:一个 step 四次前向加两次反向,中间激活值占用远超普通 GAN。解决:开启混合精度--amp,残差块从 9 降到 6,判别器通道从 64 降到 32,或者把图片裁到 192 分辨率先跑通再逐步加。

5.5 推理结果和训练时看到的不一致

现象:训练日志里生成的图看着不错,单独用 G 推理却发灰、发暗。原因:推理时忘了eval(),InstanceNorm 用了 batch 统计;或者反归一化参数写错。解决:推理前一定G.eval()加torch.no_grad(),反归一化系数和训练时的 Normalize 严格对应,最好把预处理和后处理写成同一个函数复用。

6. 进阶技巧:用感知损失和分阶段训练把去雾质量再拉一档

如果前面五章跑通后觉得效果还差口气,可以上两个进阶手段。第一个是加感知损失,用预训练 VGG 的前几层特征算 L1,让生成图在语义层面更接近清晰图。这个损失对去雾任务特别有用,因为它约束的是「内容」而不是「像素」,能明显减少伪纹理。实现上把 VGG 前 16 层冻结,取 relu2_2 和 relu3_3 两层特征,权重分别设 0.1 和 0.05,太大反而会让颜色偏。

vgg = torchvision.models.vgg16(pretrained=True).features[:16].cuda().eval() for p in vgg.parameters(): p.requires_grad = False def perceptual_loss(fake, real): f_fake, f_real = vgg(fake), vgg(real) return nn.L1Loss()(f_fake, f_real) * 0.1

第二个是分阶段训练:前 50 个 epoch 只用循环损失加对抗损失,让结构先稳住;50 到 150 epoch 加入感知损失,提升细节;150 之后加入身份损失,稳住颜色。这样比一上来全损失一起上更容易收敛,也少了很多玄学调参。我自己的习惯是每个阶段结束存一个 checkpoint,最后横向对比挑最好的,而不是死磕最后一个 epoch。验证时不要只看几张图,用 FID 或 LPIPS 在留出集上算一遍,数字比肉眼靠谱。这套方案值不值得投入?如果你手里有真实雾天数据又缺配对标签,对偶 GAN 加 PyTorch 是目前落地成本最低的路线之一,训一次大概两三天,推理单张 256 图在 1080Ti 上不到 50ms,够很多场景用了。希望帮到你。

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

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

Mac配置Java环境变量全指南:从JAVA_HOME到PATH的深度解析

1. 为什么在Mac上配环境变量经常翻车——先理解Mac的路径机制很多从Windows转过来的朋友,第一次在Mac上配置Java环境变量,都会对着终端一脸茫然:明明照着网上的教程敲了export JAVA_HOME...,重启终端又失效了;明明已经…

作者头像 李华
网站建设 2026/10/1 19:30:32

马德拉岛旅游攻略:7天6夜经典路线与避坑指南

航程单上写着“Madeira”的时候,我旁边那位葡萄牙大叔笑着说了句:“You will come back again.”我当时觉得是客套,落地第三天就明白了,他没在客套。马德拉,葡萄牙在大西洋深处的群岛,离欧洲大陆一千多公里…

作者头像 李华
网站建设 2026/10/1 19:29:12

手工标注高质量人车识别VOC数据集1000张:从VOC格式到YOLO训练全流程

简介:手工标注的1000张人车识别VOC数据集,面向计算机视觉开发者与深度学习算法工程师,用于解决行人及车辆检测任务中标注数据不足、标注质量不稳定的问题。整个压缩包共1994个文件,包括997个xml标注文件、729张jpg与268张png原始图…

作者头像 李华
网站建设 2026/10/1 19:28:55

AI工程从零构建:完整路线图、最小闭环与踩坑实战

把 ai-engineering-from-scratch 当项目名的人,大概率不是想再装个环境跑通 demo 了事,而是想把这门技术栈从地基开始重新立一遍。这几年我前后面试过不少候选人,简历上写着“熟悉 AI 开发”,但一聊到数据怎么准备、模型怎么评估、…

作者头像 李华
网站建设 2026/10/1 19:28:37

Unity切割模型实战:从Mesh切割到凸包封口与性能优化

简介:这份Unity切割模型案例面向游戏引擎初学者与希望掌握物理交互的开发者,围绕“模型切割”这一常见需求,提供可运行的实践项目。案例重点讲解碰撞检测、鼠标左键蓄力与右键触发切割的交互逻辑,以及通过修改Mesh顶点与索引数据实…

作者头像 李华
网站建设 2026/10/1 19:28:21

MySQL事务与索引实战:从原理到排障的完整指南

1. 把事务和索引拆开看:它们到底在解决什么问题先讲个我在实际项目中遇到的场景。去年帮朋友排查一个电商后台的订单接口,用户下单后页面一直转圈,数据库CPU直接飙到100%。查了半天,发现是两个程序员写代码时对同一张订单表做了不…

作者头像 李华