ARTICLE · INTELLIGENCE

战地情报 · 详情页

来自尧图项目组的一线实战观察与深度解析

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

基于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()), lr2e-4, betas(0.5, 0.999)) opt_D torch.optim.Adam(list(D_X.parameters()) list(D_Y.parameters()), lr2e-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控制循环约束强度太小去雾不彻底太大颜色失真lr2e-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, transformtf) clear_ds ImageFolder(data/clear, transformtf) hazy_loader DataLoader(hazy_ds, batch_size1, shuffleTrue, num_workers4) clear_loader DataLoader(clear_ds, batch_size1, shuffleTrue, num_workers4)逻辑说明两个 loader 独立 shuffle保证每个 step 拿到的雾图和清晰图不是同一场景。参数说明batch_size1是对偶 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(inplaceTrue), 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_weight0.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_ch3, ndf64, n_layers3): 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_layers3对应约 70×70 感受野加到 4 层感受野更大但参数增多。参数说明ndf64是通道基数显存紧张可以降到 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 分辨率batch1 仍然报 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(pretrainedTrue).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够很多场景用了。希望帮到你。本文还有配套的精品资源点击获取
RELATED READING

延伸阅读

更多一线实战笔记与深度复盘,助您持续精进