资讯详情

PyTorch对偶生成对抗网络图像去雾实战:从原理到部署

📅 2026/10/9 18:58:22 | 华诺云谱 👁 阅读
PyTorch对偶生成对抗网络图像去雾实战:从原理到部署
简介本资源为基于PyTorch实现的对偶生成对抗网络图像去雾项目面向计算机相关专业正在做毕业设计的学生以及需要项目实战练习的学习者也可作为课程设计或期末大作业参考。项目经导师指导并认可通过评审分99分代码完整可运行适合具备一定Python与深度学习基础、希望深入理解GAN去雾原理的读者。压缩包共25个文件约21.23MB包含10个py源码文件、6个png与5个jpg示例图片、2个pkl训练好的模型文件以及md说明文档和gitignore配置覆盖网络结构、训练、预测与数据加载等模块。已有144人学习关注。读者可获得完整的对偶生成对抗网络去雾实现方案包括生成器与判别器代码、预训练权重、测试图片与预测脚本便于快速复现去雾效果、理解模型训练流程并在此基础上进行二次开发或撰写论文实验对比。1. 对偶生成对抗网络做图像去雾为什么它比单向 GAN 更值得上手有雾图像的本质是场景辐射在传播路径上被大气散射函数“污染”了去雾要做的就是从观测图像里反解出干净场景。早期做法靠暗通道先验参数一多就玄学换一批数据就翻车。这几年基于 PyTorch 的对偶生成对抗网络Dual GAN方案逐渐成为主流落地路径它用两个生成器分别学“有雾→无雾”和“无雾→有雾”两个方向的映射再靠循环一致性约束把两个方向绑在一起避免单向 GAN 那种“生成得挺好看但和原图对不上”的老毛病。这套方案适合谁如果你手头有一批成对的雾图/清晰图想训一个能直接跑推理的去雾模型或者想拿训练好的权重做二次微调那对偶 GAN 是性价比很高的选择。它不需要成对数据也能训非成对模式下靠循环一致性有成对数据时收敛更快、细节保留更好。下面从网络结构、数据准备、训练脚本、推理部署一路讲到踩坑代码全部基于 PyTorch能直接抄。2. 对偶生成对抗网络去雾的原理与选型为什么是 CycleGAN 这一路2.1 单向 GAN 去雾的硬伤在哪单向 GAN 的思路很直接生成器 G 把雾图映射成清晰图判别器 D 判断“这张图是不是真实清晰图”。问题出在损失函数上——对抗损失只约束生成图的“分布”接近清晰图分布并不约束“这张生成图对应的是哪张雾图”。结果就是生成器可能把 A 雾图去雾成一张完全无关的清晰图 B判别器照样给高分因为 B 确实像真实清晰图。这个现象在去雾任务里特别明显雾的浓度、颜色偏移在不同区域差异很大单向 GAN 容易学到“整体提亮加对比度”这种偷懒映射遇到浓雾区域直接糊成一片。我见过不少单向 GAN 的去雾结果远看通透放大一看纹理全丢边缘还带伪影。2.2 对偶结构怎么把两个方向绑死对偶 GAN 的核心是两组映射G_AB 负责雾→清晰G_BA 负责清晰→雾。循环一致性损失要求 G_BA(G_AB(x)) ≈ x也就是雾图经过“去雾再重新加雾”后要能回到原样。这个约束逼着 G_AB 保留原图的内容结构不能乱生成。数学上完整损失由三部分组成对抗损失两个判别器 D_A、D_B 分别判断清晰域和雾域的真假循环一致性损失λ_cyc 加权通常取 10身份损失可选G_AB(y) ≈ yy 是清晰图用来稳定颜色选型上生成器我一般用 ResNet 风格的 9 残差块结构下采样两次、上采样两次中间堆残差块。判别器用 PatchGAN输出 70×70 的感受野比整图判别更关注局部纹理去雾这种细节敏感任务用 PatchGAN 明显更稳。提示如果你只有非成对数据循环一致性是唯一的内容约束λ_cyc 不能调太小否则两个生成器会各玩各的。有成对数据时建议额外加 L1 监督损失收敛快很多。2.3 生成器与判别器的 PyTorch 实现先看生成器的残差块和整体结构。下面这段代码可以直接用输入输出都是 3 通道 RGB尺寸不限制全卷积。import torch import torch.nn as nn class ResidualBlock(nn.Module): 标准残差块两个 3x3 卷积 InstanceNorm ReLU带跳跃连接 def __init__(self, channels): super().__init__() self.block nn.Sequential( nn.ReflectionPad2d(1), # 反射填充避免边缘伪影 nn.Conv2d(channels, channels, 3), nn.InstanceNorm2d(channels), # 去雾任务用 InstanceNorm 比 BatchNorm 稳 nn.ReLU(inplaceTrue), nn.ReflectionPad2d(1), nn.Conv2d(channels, channels, 3), nn.InstanceNorm2d(channels) ) def forward(self, x): return x self.block(x) # 跳跃连接保留原始信息 class Generator(nn.Module): ResNet 生成器下采样 - 9 残差块 - 上采样 def __init__(self, in_ch3, out_ch3, ngf64, n_blocks9): super().__init__() layers [ nn.ReflectionPad2d(3), nn.Conv2d(in_ch, ngf, 7), nn.InstanceNorm2d(ngf), nn.ReLU(inplaceTrue) ] # 两次下采样每次通道翻倍 for i in range(2): mult 2 ** i layers [ nn.Conv2d(ngf * mult, ngf * mult * 2, 3, stride2, padding1), nn.InstanceNorm2d(ngf * mult * 2), nn.ReLU(inplaceTrue) ] # 堆 9 个残差块 for _ in range(n_blocks): layers.append(ResidualBlock(ngf * 4)) # 两次上采样用转置卷积恢复分辨率 for i in range(2): mult 2 ** (2 - i) layers [ nn.ConvTranspose2d(ngf * mult, ngf * mult // 2, 3, stride2, padding1, output_padding1), nn.InstanceNorm2d(ngf * mult // 2), nn.ReLU(inplaceTrue) ] layers [nn.ReflectionPad2d(3), nn.Conv2d(ngf, out_ch, 7), nn.Tanh()] self.model nn.Sequential(*layers) def forward(self, x): return self.model(x)逻辑说明ReflectionPad2d 在卷积前做反射填充比零填充更能减少边界伪影去雾图边缘经常出问题这一步别省。InstanceNorm2d 对每个样本每个通道单独归一化不依赖 batch 统计量小 batch 训练时比 BatchNorm 稳定得多。9 个残差块是 CycleGAN 原论文的默认配置显存不够可以降到 6 个但去雾细节会略降。最后用 Tanh 把输出压到 [-1,1]和训练时的归一化范围对齐。判别器用 PatchGAN输出一个 N×N 的 patch 概率图class Discriminator(nn.Module): PatchGAN 判别器输出 70x70 感受野的 patch 真假图 def __init__(self, in_ch3, ndf64): super().__init__() self.model nn.Sequential( nn.Conv2d(in_ch, ndf, 4, stride2, padding1), nn.LeakyReLU(0.2, inplaceTrue), nn.Conv2d(ndf, ndf * 2, 4, stride2, padding1), nn.InstanceNorm2d(ndf * 2), nn.LeakyReLU(0.2, inplaceTrue), nn.Conv2d(ndf * 2, ndf * 4, 4, stride2, padding1), nn.InstanceNorm2d(ndf * 4), nn.LeakyReLU(0.2, inplaceTrue), nn.Conv2d(ndf * 4, 1, 4, padding1) # 输出单通道 patch 图 ) def forward(self, x): return self.model(x)参数说明ndf64 是基础通道数显存紧张可以降到 32。LeakyReLU 斜率 0.2 是 GAN 判别器的常规选择防止梯度死亡。最后一层不加激活输出 logits配合 BCEWithLogitsLoss 使用。3. 数据准备与训练脚本从成对/非成对数据到能跑的 train.py3.1 数据集组织与预处理去雾数据集常见两种组织方式。成对数据如合成雾图放两个文件夹文件名一一对应非成对数据放两个独立文件夹不需要对应关系。目录结构建议这样dataset/ ├── trainA/ # 雾图 │ ├── 001.png │ └── ... ├── trainB/ # 清晰图 │ ├── 001.png │ └── ... ├── valA/ # 验证集雾图 └── valB/ # 验证集清晰图预处理我一般做三件事统一缩放到 256×256或 286×286 再随机裁剪到 256、归一化到 [-1,1]、随机水平翻转做增强。注意雾图不能做颜色抖动否则会破坏雾的物理特性模型学到的映射就偏了。from torch.utils.data import Dataset from PIL import Image import torchvision.transforms as T import os class DehazeDataset(Dataset): def __init__(self, root_a, root_b, size256, pairedFalse): self.files_a sorted(os.listdir(root_a)) self.files_b sorted(os.listdir(root_b)) self.root_a, self.root_b root_a, root_b self.paired paired self.transform T.Compose([ T.Resize((size, size), Image.BICUBIC), T.ToTensor(), T.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5)) # 归一化到 [-1,1] ]) def __len__(self): # 非成对时取较大长度索引取模 return max(len(self.files_a), len(self.files_b)) def __getitem__(self, idx): a Image.open(os.path.join(self.root_a, self.files_a[idx % len(self.files_a)])).convert(RGB) if self.paired: # 成对模式B 用同名文件 b Image.open(os.path.join(self.root_b, self.files_a[idx % len(self.files_a)])).convert(RGB) else: b Image.open(os.path.join(self.root_b, self.files_b[idx % len(self.files_b)])).convert(RGB) return {A: self.transform(a), B: self.transform(b)}逻辑说明非成对模式下 A、B 各自独立采样索引取模保证不越界。成对模式下 B 用 A 的同名文件保证内容对应。Normalize 用 0.5 均值和标准差是 GAN 训练的惯例和生成器最后的 Tanh 对应。3.2 训练循环与损失函数配置训练脚本的核心是四个网络交替更新。下面给出关键部分完整脚本按这个骨架补全即可。import torch import torch.nn as nn from torch.utils.data import DataLoader # 初始化四个网络 G_AB Generator().cuda() # 雾 - 清晰 G_BA Generator().cuda() # 清晰 - 雾 D_A Discriminator().cuda() # 判别清晰域 D_B Discriminator().cuda() # 判别雾域 # 损失函数 criterion_gan nn.BCEWithLogitsLoss() criterion_cyc nn.L1Loss() criterion_idt nn.L1Loss() # 优化器生成器和判别器分开 opt_G torch.optim.Adam( list(G_AB.parameters()) list(G_BA.parameters()), lr2e-4, betas(0.5, 0.999)) # betas 0.5 是 GAN 训练惯例 opt_D torch.optim.Adam( list(D_A.parameters()) list(D_B.parameters()), lr2e-4, betas(0.5, 0.999)) lambda_cyc 10.0 # 循环一致性权重 lambda_idt 5.0 # 身份损失权重 for epoch in range(num_epochs): for batch in dataloader: real_A batch[A].cuda() # 雾图 real_B batch[B].cuda() # 清晰图 # ---- 训练生成器 ---- opt_G.zero_grad() fake_B G_AB(real_A) # 去雾结果 fake_A G_BA(real_B) # 加雾结果 rec_A G_BA(fake_B) # 循环回来 rec_B G_AB(fake_A) # 对抗损失骗过判别器 loss_gan_AB criterion_gan(D_B(fake_B), torch.ones_like(D_B(fake_B))) loss_gan_BA criterion_gan(D_A(fake_A), torch.ones_like(D_A(fake_A))) # 循环一致性 loss_cyc criterion_cyc(rec_A, real_A) criterion_cyc(rec_B, real_B) # 身份损失稳定颜色 loss_idt criterion_idt(G_AB(real_B), real_B) \ criterion_idt(G_BA(real_A), real_A) loss_G loss_gan_AB loss_gan_BA \ lambda_cyc * loss_cyc lambda_idt * loss_idt loss_G.backward() opt_G.step() # ---- 训练判别器 ---- opt_D.zero_grad() # 真样本判真 loss_D_A_real criterion_gan(D_A(real_B), torch.ones_like(D_A(real_B))) loss_D_B_real criterion_gan(D_B(real_A), torch.ones_like(D_B(real_A))) # 假样本判假detach 切断生成器梯度 loss_D_A_fake criterion_gan(D_A(fake_A.detach()), torch.zeros_like(D_A(fake_A))) loss_D_B_fake criterion_gan(D_B(fake_B.detach()), torch.zeros_like(D_B(fake_B))) loss_D (loss_D_A_real loss_D_B_real loss_D_A_fake loss_D_B_fake) * 0.5 loss_D.backward() opt_D.step()逻辑说明生成器损失里对抗损失用 ones_like 作为目标意思是“让判别器以为这是真的”。判别器训练时对假样本要 detach否则梯度会回传到生成器把两个网络的更新搅在一起。lambda_cyc10 是 CycleGAN 论文的默认值去雾任务里我试过 5 到 2010 比较均衡lambda_idt5 用来防止颜色漂移如果发现去雾图偏色严重可以加到 10。参数说明学习率 2e-4、betas(0.5, 0.999) 是 GAN 训练的标准配置betas 第一个值调小是为了让动量不要太大避免判别器更新过猛。batch size 建议 1 到 4PatchGAN 对小 batch 友好显存 8G 也能跑 256×256。3.3 训练监控与 checkpoint 保存训练过程中要盯三个指标G 的循环一致性损失、D 的对抗损失、以及验证集上的 PSNR/SSIM。循环损失持续下降说明内容保留在变好判别器损失如果长期接近 0说明判别器太强生成器学不动这时候要降低 D 的学习率或者给 D 加噪声。# 每 5 个 epoch 存一次 checkpoint同时保存两个生成器 if epoch % 5 0: torch.save({ G_AB: G_AB.state_dict(), G_BA: G_BA.state_dict(), D_A: D_A.state_dict(), D_B: D_B.state_dict(), epoch: epoch, opt_G: opt_G.state_dict(), opt_D: opt_D.state_dict() }, fcheckpoints/dehaze_epoch_{epoch}.pth)保存优化器状态是为了断点续训去雾模型通常要训 100 到 200 个 epoch中途断了没有优化器状态就得重来这个后悔药不好吃。4. 推理部署与效果验证把训练好的模型跑起来4.1 加载权重做单图推理训练好的模型推理很简单只需要 G_AB 一个网络。下面这段代码可以直接拿去用import torch from PIL import Image import torchvision.transforms as T def dehaze_image(model_path, input_path, output_path, img_size256): device torch.device(cuda if torch.cuda.is_available() else cpu) # 重建生成器结构并加载权重 G_AB Generator().to(device) ckpt torch.load(model_path, map_locationdevice) G_AB.load_state_dict(ckpt[G_AB]) G_AB.eval() # 切到推理模式InstanceNorm 行为会变 # 预处理 img Image.open(input_path).convert(RGB) transform T.Compose([ T.Resize((img_size, img_size), Image.BICUBIC), T.ToTensor(), T.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5)) ]) x transform(img).unsqueeze(0).to(device) with torch.no_grad(): # 关闭梯度省显存 fake_B G_AB(x) # 反归一化回 [0,1] 并保存 fake_B (fake_B.squeeze(0).cpu() * 0.5 0.5).clamp(0, 1) out T.ToPILImage()(fake_B) out.save(output_path) print(f去雾结果已保存到 {output_path}) dehaze_image(checkpoints/dehaze_epoch_100.pth, hazy.png, clear.png)逻辑说明eval() 必须调用InstanceNorm 在训练和推理模式下行为不同不切会出问题。torch.no_grad() 关闭梯度计算推理速度能快 30% 左右。反归一化用乘 0.5 加 0.5和训练时的 Normalize 对应clamp 防止溢出。参数说明img_size 要和训练时一致训练用 256 推理也用 256。如果原图分辨率很高建议先缩放到 256 去雾再放大回去或者用全卷积特性直接跑大图显存够的话但效果可能和训练分布不一致。4.2 用 PSNR 和 SSIM 量化去雾效果光看肉眼看不出模型好坏得有量化指标。成对数据上直接算 PSNR 和 SSIMfrom skimage.metrics import peak_signal_noise_ratio as psnr from skimage.metrics import structural_similarity as ssim import numpy as np def evaluate(clear_path, dehazed_path): clear np.array(Image.open(clear_path).convert(RGB)) dehazed np.array(Image.open(dehazed_path).convert(RGB)) # 确保尺寸一致 if clear.shape ! dehazed.shape: dehazed np.array(Image.fromarray(dehazed).resize( (clear.shape[1], clear.shape[0]))) p psnr(clear, dehazed, data_range255) s ssim(clear, dehazed, channel_axis2, data_range255) print(fPSNR: {p:.2f} dB, SSIM: {s:.4f}) return p, s逻辑说明PSNR 衡量像素级误差SSIM 衡量结构相似度。去雾任务里 PSNR 到 20dB 以上、SSIM 到 0.8 以上算可用到 25dB/0.9 以上算不错。注意这两个指标对颜色偏移不敏感如果去雾图整体偏蓝PSNR 可能还行但肉眼很难看所以指标要结合肉眼一起看。4.3 非成对数据上的验证方法没有清晰图做参考时PSNR/SSIM 用不了。我一般用两个替代指标一是雾密度估计用暗通道先验算去雾前后暗通道的均值均值越低说明雾越少二是无参考图像质量评价比如 BRISQUE 分数。这两个指标不完美但能横向对比不同 checkpoint 的好坏。5. 避坑与排查对偶 GAN 去雾最常见的 5 个翻车现场5.1 生成图整体偏色像蒙了一层蓝膜现象去雾结果整体偏蓝或偏黄PSNR 还行但肉眼没法看。原因身份损失权重太低或者训练数据里清晰图的颜色分布和雾图差异太大生成器学偏了。另一个常见原因是判别器太强生成器为了骗过判别器走了“整体调色”的捷径。解决把 lambda_idt 从 5 提到 10或者在生成器损失里加一个颜色一致性损失约束去雾图和原图的均值差异。如果判别器损失长期低于 0.1把 D 的学习率降到 1e-4。5.2 循环一致性损失降不下去卡在某个值不动现象loss_cyc 训了几十个 epoch 还在 0.3 以上去雾图内容对不上原图。原因生成器容量不够或者 lambda_cyc 太小。还有一种可能是数据里雾图和清晰图的内容差异太大非成对模式下常见循环映射本身就不成立。解决先把 lambda_cyc 加到 15 试试不行就把残差块从 9 加到 12或者把 ngf 从 64 提到 96。如果是非成对数据检查两个域的内容分布是否接近差太远的话循环一致性约束会互相打架。5.3 训练到一半判别器损失变成 0生成器完全不更新现象D 的损失突然掉到接近 0G 的损失开始震荡或爆炸。原因判别器太强把真假样本完全分开了生成器梯度消失。这是 GAN 训练的经典问题对偶结构里两个判别器同时变强会加速这个过程。解决给判别器加标签平滑把真样本目标从 1.0 改成 0.9或者给判别器输入加高斯噪声标准差 0.1。另一个办法是判别器每训 1 次、生成器训 2 次让生成器多学一点。5.4 推理时显存爆了或者速度慢得没法用现象单张 1080p 图推理要好几秒或者直接 OOM。原因全卷积网络对输入尺寸没限制1080p 图直接跑中间特征图会非常大。另外没加 torch.no_grad() 也会多占显存。解决推理前把图缩到 256 或 512去雾完再放大回去。如果必须处理大图用滑动窗口分块推理每块 256×256块之间重叠 32 像素避免接缝。torch.no_grad() 和 model.eval() 一个都不能少。5.5 换一批数据效果就崩泛化性差现象在自己数据集上训得挺好换一批雾图去雾效果明显下降。原因训练数据太单一雾的浓度、颜色、场景类型覆盖不够。对偶 GAN 虽然比单向 GAN 泛化好但也扛不住训练分布和测试分布差太远。解决训练时做更强的数据增强除了翻转还可以加随机裁剪、轻微缩放。如果目标域数据能拿到一点做微调最有效冻结判别器只训生成器 10 个 epoch 就能明显改善。另外合成雾图训练时雾的浓度参数要随机化别只用固定浓度。6. 进阶技巧把对偶 GAN 去雾推到能用的水平前面讲的都是能跑通的基础版但真要用起来还有几个技巧值得试。第一个是感知损失在生成器损失里加一个 VGG 特征匹配损失权重取 0.1 左右能明显改善去雾图的纹理细节。VGG 用预训练权重取 relu3_3 层的特征计算生成图和清晰图的特征 L1 距离。这个损失对浓雾区域的细节恢复特别有效代价是训练慢 20% 左右。第二个技巧是判别器用多尺度结构。单个 PatchGAN 只关注 70×70 的感受野对全局雾分布不敏感。加一个下采样 2 倍的判别器分支两个尺度一起判能同时约束局部纹理和全局通透度。实现上就是把输入图缩一半再过一个判别器损失加权求和。第三个是学习率调度。前 50 个 epoch 用 2e-4 恒定之后每 50 个 epoch 线性衰减到 0。我试过余弦退火效果不如线性衰减稳GAN 训练里学习率突变容易让判别器崩掉。最后一个技巧关于模型导出。如果要去雾后接其他视觉任务建议把生成器导出成 ONNX用 onnxruntime 推理比 PyTorch 快 1.5 到 2 倍。导出时注意固定输入尺寸动态轴虽然支持但某些算子会出问题。# 导出 ONNX dummy torch.randn(1, 3, 256, 256).cuda() torch.onnx.export( G_AB, dummy, dehaze.onnx, input_names[hazy], output_names[clear], opset_version11, # 11 对 InstanceNorm 支持好 dynamic_axes{hazy: {2: h, 3: w}, clear: {2: h, 3: w}} )导出后拿 onnxruntime 跑一遍对比 PyTorch 输出误差在 1e-4 以内算正常。如果误差大检查 opset 版本和算子支持情况。我自己训去雾模型踩过最大的坑是过早看指标——前 20 个 epoch PSNR 涨得很快以为要成了结果 50 epoch 后开始过拟合验证集指标掉头向下。后来养成习惯每 5 个 epoch 存一次 checkpoint最后从验证集指标最好的那个往回挑而不是用最后一个。这个习惯帮我省了不少重训的时间。希望帮到你。本文还有配套的精品资源点击获取
📝

华诺云谱内容团队

资深建站顾问 · 行业研究员

10年+企业数字化服务经验,专注智能建站、SEO优化与品牌营销,持续输出建站技巧、行业洞察与营销干货,已帮助5000+企业实现数字化增长。

你可能需要的服务

订阅华诺云谱资讯周报

每周一封,精选建站技巧、SEO与营销干货,直达邮箱。已有 8,000+ 企业主订阅,助你少走弯路。

↑