☰
PyTorch对偶生成对抗网络图像去雾实战:从原理到部署
2026/10/9 18:58:18 网站建设 项目流程

简介:本资源为基于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) ≈ y,y 是清晰图,用来稳定颜色

选型上,生成器我一般用 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(inplace=True), 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_ch=3, out_ch=3, ngf=64, n_blocks=9): super().__init__() layers = [ nn.ReflectionPad2d(3), nn.Conv2d(in_ch, ngf, 7), nn.InstanceNorm2d(ngf), nn.ReLU(inplace=True) ] # 两次下采样,每次通道翻倍 for i in range(2): mult = 2 ** i layers += [ nn.Conv2d(ngf * mult, ngf * mult * 2, 3, stride=2, padding=1), nn.InstanceNorm2d(ngf * mult * 2), nn.ReLU(inplace=True) ] # 堆 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, stride=2, padding=1, output_padding=1), nn.InstanceNorm2d(ngf * mult // 2), nn.ReLU(inplace=True) ] 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_ch=3, ndf=64): super().__init__() self.model = nn.Sequential( nn.Conv2d(in_ch, ndf, 4, stride=2, padding=1), nn.LeakyReLU(0.2, inplace=True), nn.Conv2d(ndf, ndf * 2, 4, stride=2, padding=1), nn.InstanceNorm2d(ndf * 2), nn.LeakyReLU(0.2, inplace=True), nn.Conv2d(ndf * 2, ndf * 4, 4, stride=2, padding=1), nn.InstanceNorm2d(ndf * 4), nn.LeakyReLU(0.2, inplace=True), nn.Conv2d(ndf * 4, 1, 4, padding=1) # 输出单通道 patch 图 ) def forward(self, x): return self.model(x)

参数说明:ndf=64 是基础通道数,显存紧张可以降到 32。LeakyReLU 斜率 0.2 是 GAN 判别器的常规选择,防止梯度死亡。最后一层不加激活,输出 logits,配合 BCEWithLogitsLoss 使用。

3. 数据准备与训练脚本:从成对/非成对数据到能跑的 train.py

3.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, size=256, paired=False): 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()), lr=2e-4, betas=(0.5, 0.999)) # betas 0.5 是 GAN 训练惯例 opt_D = torch.optim.Adam( list(D_A.parameters()) + list(D_B.parameters()), lr=2e-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_cyc=10 是 CycleGAN 论文的默认值,去雾任务里我试过 5 到 20,10 比较均衡;lambda_idt=5 用来防止颜色漂移,如果发现去雾图偏色严重可以加到 10。

参数说明:学习率 2e-4、betas=(0.5, 0.999) 是 GAN 训练的标准配置,betas 第一个值调小是为了让动量不要太大,避免判别器更新过猛。batch size 建议 1 到 4,PatchGAN 对小 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() }, f'checkpoints/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_size=256): device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') # 重建生成器结构并加载权重 G_AB = Generator().to(device) ckpt = torch.load(model_path, map_location=device) 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 和 SSIM:

from 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_range=255) s = ssim(clear, dehazed, channel_axis=2, data_range=255) print(f'PSNR: {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 的损失突然掉到接近 0,G 的损失开始震荡或爆炸。

原因:判别器太强,把真假样本完全分开了,生成器梯度消失。这是 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_version=11, # 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,最后从验证集指标最好的那个往回挑,而不是用最后一个。这个习惯帮我省了不少重训的时间。希望帮到你。

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

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询