简介:基于Pytorch实现的对偶生成对抗网络(DualGAN)图像去雾项目,面向计算机视觉方向的课程设计、毕业设计以及入门GAN的开发者。项目采用U-Net作为生成器、PatchGAN作为判别器,通过G_A(有雾→无雾)与G_B(无雾→有雾)两个生成器和对应的两个判别器完成对抗训练,代码结构清晰,便于理解对偶对抗机制。压缩包共26个文件,以10个Python源码为核心,涵盖训练、预测、数据加载与参数解析等模块,另含预训练模型(.pkl)、效果图与测试图片、说明文档等,总大小21.23MB。已有210人学习下载,资源内附文档说明和关键注释,并提供了完整可运行的训练与预测流程,下载后可直接体验去雾效果,也方便在此基础上进行改进和二次开发,非常适合作为毕设、课设或项目初期演示使用。
1. 雾图不配对也能学:DualGAN 把去雾做成了循环翻译
有雾图像到无雾图像的转换,最大的瓶颈不是模型容量,而是训练数据。真实场景里你很难拿到同一机位、同一视角的“有雾/无雾”严格配对样本,而合成数据训练的模型拿到户外又容易瘫。Pytorch 生态里基于生成对抗网络的去雾方案不少,但大部分需要配对监督。DualGAN 走的是另一条路线:用两个生成器和两个判别器构成循环翻译结构,不需要像素级配对,只要“有雾图集合”和“无雾图集合”两个目录就能训练。对于做课程设计、毕业设计,或者想快速在图像到图像翻译上落地的 Python 开发者来说,这套代码的结构非常规整:生成器是 U-Net,判别器是 PatchGAN,训练和预测脚本分离,还带了预训练权重。下面从网络结构拆到训练参数,再把推理阶段容易踩的坑一起过一遍。
2. DualGAN 网络结构与 Pytorch 实现拆解
2.1 生成器为什么选 U-Net:跳连结构保住纹理细节
DualGAN 里有对称的两个生成器:G_A 负责把有雾图映射成无雾图,G_B 负责把无雾图映射成有雾图。代码里两个生成器共用同一个 U-Net 结构,只是实例化时区分角色。U-Net 的骨架是“编码器-解码器”,编码器逐层下采样提取语义特征,解码器逐层上采样恢复分辨率。单看这条路径,细节信息会在下采样过程中丢失,去雾结果容易出现边缘模糊、纹理被抹平的问题。
U-Net 的关键是每层下采样后的特征图会通过跳连(skip connection)直接拼到对应层上采样结果上。这样解码器在恢复细节时,手里同时握着高层语义和低层纹理,去雾后的图像能保留更多原始边缘和纹理信息。在 Pytorch 里,这种结构通常写在net/Generator.py,核心就是一个双路径的模块封装。常见实现里编码器用卷积加 InstanceNorm,解码器用转置卷积或者上采样加卷积,跳连通过torch.cat完成。一般写法类似这样:
# 伪代码级实现,对应 net/Generator.py 的内部逻辑 class UNetBlock(nn.Module): def __init__(self, in_ch, out_ch, inner=False): super().__init__() if inner: self.conv = nn.Conv2d(in_ch, out_ch, 3, padding=1) self.norm = nn.InstanceNorm2d(out_ch) else: self.conv = nn.Sequential( nn.Conv2d(in_ch, out_ch, 3, padding=1), nn.InstanceNorm2d(out_ch), nn.ReLU(True), nn.Conv2d(out_ch, out_ch, 3, padding=1), nn.InstanceNorm2d(out_ch), ) self.relu = nn.ReLU(True) def forward(self, x): return self.relu(self.norm(self.conv(x)))这里inner=True表示 bottleneck 层只做一次卷积,不做下采样和上采样;非 bottleneck 层则被包在编码和解码路径里多次调用。选用 InstanceNorm 而不是 BatchNorm,是因为图像翻译任务里单张样本的统计量往往比 batch 统计量更稳定,去雾这种任务输入图像的亮度分布差异很大,BatchNorm 容易受同 batch 其他图的统计量干扰。Pytorch 里nn.InstanceNorm2d不需要维护 running mean,对单卡小 batch 训练更友好。
2.2 PatchGAN 判别器:6 通道输入怎么拼
判别器在 DualGAN 里也不是普通的“输出一个真/假标量”的分类器,而是 PatchGAN。PatchGAN 输出的是一个N x N的矩阵,每个元素对应输入图像上一块局部的真假判断。用这个设计的好处是参数少、关注局部纹理和颜色分布,能逼着生成器把每个局部区域都做得像真的,而不是只骗过全局统计。
项目中 D_A 和 D_B 的输入都是 6 通道。以 D_B 为例,它负责判别 G_A 生成的无雾图像,输入由两部分拼成:真实的 clear 图像和 G_A 生成的无雾图像,沿着通道维度拼接,得到[B, 6, H, W]的张量。D_A 同理,输入是真实的 hazy 图像和 G_B 生成的有雾图像拼接。判别器的任务就是区分拼接对里的“第二张图”到底来自真实集合还是生成器。
| 网络 | 输入拼接构成 | 判断目标 | 训练符号 |
|---|---|---|---|
| D_A | 真实有雾图hazy+ G_B 生成的fake_hazy | 分辨fake_hazy的真假 | 收 G_B 的生成结果 |
| D_B | 真实无雾图clear+ G_A 生成的fake_clear | 分辨fake_clear的真假 | 收 G_A 的生成结果 |
| G_A | 有雾图hazy | 欺骗 D_B | 与 D_B 对抗 |
| G_B | 无雾图clear | 欺骗 D_A | 与 D_A 对抗 |
从 Pytorch 代码看,拼接操作就是torch.cat([real, fake], dim=1),dim=1 对应通道维度。这里有一个很多初学容易忽略的点:判别器第一层卷积的in_channels必须是 6,而不是 3。改网络结构时如果只改了生成器没改判别器的输入通道,训练一启动就会报尺寸不匹配的错误。训练脚本里 dual 模型会把 G_A、G_B、D_A、D_B 四个网络实例统一封装,再分别传入各自的优化器。
2.3 循环一致性损失:对偶结构训练的支点
只有对抗损失的对偶训练是不稳定的,因为 G_A 可以把任意有雾图都映射成同一张“像无雾”的图,只要骗过 D_B 就行,这会丢失原图内容。DualGAN 用循环一致性来约束这种退化:一张有雾图经过 G_A 去雾后,再用 G_B 加雾回来,结果应该和原图基本一致;反过来对无雾图做一次 G_B → G_A 的循环也是如此。这样两个生成器就被绑成一个闭环,谁也偷不了懒。
Pytorch 实现里循环损失一般用 L1 距离,因为 L1 对边缘的惩罚比 L2 柔和,不容易把图磨平。代码上就是两次前向:
fake_clear = G_A(hazy) # 有雾 -> 无雾 recon_hazy = G_B(fake_clear) # 无雾 -> 有雾(还原) cycle_loss = L1_loss(recon_hazy, hazy) fake_hazy = G_B(clear) # 无雾 -> 有雾 recon_clear = G_A(fake_hazy) # 有雾 -> 无雾(还原) cycle_loss += L1_loss(recon_clear, clear)recon_hazy和recon_clear就是循环重建结果。训练时用optimizer_G.zero_grad()清空梯度后,把 GAN 损失和 cycle loss 一起反传。项目文档里写 D_A 和 D_B 的输入是 6 通道,也正是为了配合这种“原图 + 生成图”成对判断的模式。理解了这个闭环,后面调参才知道改哪个损失、动哪个网络。
3. 数据组织与训练复现:train.py 参数怎么给
3.1 数据目录与配对的两种组织方式
train.py的训练入口要求数据放在data_path下,项目约定把成对图片分别放进clear和hazy两个文件夹。这里“成对”指的是文件名一一对应,比如1404_7.png在hazy里,对应clear里也有一张同名的无雾版本。这种组织方式在代码里通常通过按文件列表顺序读取实现,因此两个目录下的文件名必须完全一致,否则会读到错位的图。
如果确实拿不到配对数据,DualGAN 理论上也能训练,因为循环一致性损失不要求样本配对。但项目的数据加载器是在预加载阶段按索引对齐的,所以非配对场景需要改data_loader的读取逻辑,把“按索引取数”改成“按目录随机采样”。常见做法是直接保留现有结构,用合成数据集(比如 RESIDE 的子集)做训练,这样最简单、坑最少。
数据准备阶段还有一个容易忽略的环节:训练集图片尺寸。U-Net 对分辨率不敏感,但 PatchGAN 的 patch 数会随输入尺寸变化,训练时最好统一长宽。项目中图片 1404_7.png 这类样本是 600×400 左右,为了稳定训练,建议在parseArgs.py里加上--size参数,将图像缩放到固定尺寸,比如 256×256,或者保持 2 的幂次倍数,避免下采样时出现奇数维度。
3.2 训练命令与关键参数说明
环境方面,项目基于 Pytorch 实现,需要先确认本机 Python 版本和 Pytorch 环境能否对得上。以 Python 3.8 搭配 Pytorch 1.10 左右的组合比较省心,Pytorch 2.x 也能跑,但要注意旧代码里torchvision的接口变化。GPU 不是必须的,但这个模型用 CPU 训练会很慢,建议优先配置好 CUDA 版的 Pytorch,显存 6G 以上跑 256×256 的 batch_size=2 没有压力。
启动训练的命令写法如下:
python train.py \ --data_path ./data \ --size 256 \ --batch_size 2 \ --epoch 100 \ --lr 0.0002 \ --lambda 10.0我一般会把参数含义和使用建议放到表格里,方便对照:
| 参数名 | 常见取值 | 作用与调整建议 |
|---|---|---|
data_path | ./data | 指向包含clear和hazy子文件夹的根目录 |
size | 256 | 统一输入尺寸,显存小就降到 128 |
batch_size | 1~4 | 显存不足时优先减这个值,不要直接降分辨率 |
epoch | 100~200 | 100 epoch 能看到基本效果,200 效果更好 |
lr | 0.0002 | 生成器和判别器共用,训练不稳定时降为 0.0001 |
lambda | 10 | 循环一致性损失权重,权重太大导致重建模糊 |
--lambda控制的是循环一致性损失的整体占比。值太大,生成器会偏向保守,输出图像更接近原图,去雾不彻底;值太小,对抗损失主导,图像可能产生伪影。论文里的常见设置是 10,实际训练中如果发现重建图细节保留得很好但去雾强度不够,可以适当下调到 5;如果图像开始出现奇怪的色斑和纹理,就调高一些。
3.3 从 loss.png 判断收敛状态
训练过程中项目会持续生成loss.png,里面记录了生成器损失、判别器损失以及循环一致性损失的曲线。很多新手盯着一个 total loss 看,这是不够的。正确姿势是拆开看三条曲线的相对关系:
- 生成器 G 的对抗损失下降,说明生成图像越来越能骗过判别器;
- 判别器 D 的损失波动但整体不坍缩为 0,说明判别器没有被生成器“打趴”,对抗是健康的;
- 循环一致性损失稳定下降,说明闭环约束在起作用,内容保真度在提升。
如果发现 D 的损失一路跌到接近 0,而 G 的损失还在涨,通常是判别器太强了,生成器梯度不稳定。常见处理是把判别器的学习率调低,比如 G 保持2e-4,D 降到5e-5,或者给判别器加谱归一化。如果 cycle loss 降不下去,先检查数据目录是否配对正确,再确认--lambda是否存在手滑设成了 0,这两个原因占了循环损失不收敛的大多数情况。、
顺带提一句,项目中train.py与dual.py分工明确,前者处理命令行参数、日志和 checkpoint 保存,后者负责四个网络的实例化与前向计算。想改损失权重,去dual.py后半部分找 loss 计算段;想改网络层数,去net/目录下的两个网络文件。这个分层对课程设计和毕设答辩来说,代码结构本身就是一个加分项。
4. predict.py 推理与模型加载实战
4.1 推理指令与输出
去雾效果验证不需要把整套训练再跑一遍,项目提供了predict.py来加载预训练模型并对单张或目录内图片做推理。推理目录下默认放了几张测试图,比如1408_10.png和1423_5.png,对应的预测结果会输出到predict文件夹,格式是 jpg。执行方式一般是:
python predict.py \ --model_dir ./model \ --input ./test_data \ --output ./output \ --device cuda这段命令的逻辑是:从model_dir读取预训练权重,把test_data下所有图片逐个过一遍 G_A,结果写到output。由于去雾走的是 G_A(有雾→无雾),所以推理阶段真正加载的生成器是 G_A,D_B 只是在对抗训练中参与判别,不参与预测。很多初学者把discriminator_*.pkl当成了可以直接调用的去雾模型,这是理解上的偏差。
4.2 pkl 模型加载:先建结构再灌权重
model目录下预训练文件是discriminator_b.pkl和discriminator_a.pkl。.pkl后缀只是保存时的文件名习惯,并不代表文件格式特殊,实际内容仍是 Pytorch 的torch.save序列化结果。加载时最容易踩的坑在于:如果保存的是state_dict(推荐方式),必须先实例化对应结构的网络对象,再调用load_state_dict,而不是直接用torch.load拿到的字典去 forward。
import torch from net.Discriminator import PatchDiscriminator from net.Generator import UNetGenerator device = torch.device("cuda" if torch.cuda.is_available() else "cpu") G_A = UNetGenerator(in_channels=3, out_channels=3).to(device) G_A.eval() # 若模型目录中保存的是生成器权重 G_A.load_state_dict(torch.load("model/generator_a.pkl", map_location=device))Pytorch 2.x 的torch.load默认weights_only=True,加载旧版 sage 模型时可能直接抛参数不兼容的错,这个错误信息往往会误导人以为是文件坏了。我的处理习惯是先在 CPU 上加载看看torch.load(..., map_location='cpu')再定位问题。另外要检查 pkl 里的 key 是否与网络层名完全一致,如果网络结构被改过而权重是旧的,load_state_dict会提示missing keys或unexpected keys。项目文档里专门提到代码测试通过、答辩评审平均分较高,说明预训练权重和原始结构是对应关系,改动网络前最好把原始权重备份一份。
4.3 GPU/CPU 切换与预处理对齐
推理时设备切换不只是model.to(device)那么简单,还涉及数据张量的迁移。Pytorch 在这方面的常见报错是Expected all tensors to be on the same device,解决办法是把输入图片张量也显式to(device)。另一个坑是图像预处理必须和训练时保持一致:训练时如果做了归一化(比如把像素缩放到[-1, 1]),推理时也要用同一套规则,否则生成器的输入分布漂移,输出会偏色。
常见做法是推理时读取图像后,先转 RGB、resize 到训练尺寸、转 float 类型再归一化:
from PIL import Image import torchvision.transforms as T transform = T.Compose([ T.Resize((256, 256)), T.ToTensor(), T.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5]), ]) img = Image.open("test_data/1408_10.png").convert("RGB") x = transform(img).unsqueeze(0).to(device) with torch.no_grad(): fake_clear = G_A(x)unsqueeze(0)是为了把单张图扩展成 batch 维度,torch.no_grad()关闭梯度计算省显存。推理完成后,如果输出是归一化张量,保存前要反归一化回[0,1],再转成uint8存图。常见做法是:
out = (fake_clear.squeeze(0).permute(1, 2, 0) * 0.5 + 0.5).clamp(0, 1).numpy()这里* 0.5 + 0.5是Normalize(0.5, 0.5)的逆变换,.clamp(0, 1)防止浮点误差导致像素越界。这一步漏了的话,输出图像会偏暗或者出现大面积灰色。
5. 对偶训练中容易被忽略的三个细节与效果验证技巧
5.1 判别器输入拼接顺序与数据增强同步
对偶训练里 G_A 和 G_B 是镜像对称的,任何一侧的网络更新都会影响另一侧的输入分布,因此训练时最忌讳的是生成器和判别器的更新频率不对称。常见做法是每个 iteration 先更新生成器两次、再更新判别器一次,让判别器慢半拍,避免它过强。第二个细节是 6 通道输入的拼接顺序必须全局一致:要么永远是real在前fake在后,要么永远相反,混用会导致判别器学会“按位置判断真假”,而不是按内容判断。
数据增强也要保持对偶关系。如果对输入图做了随机翻转或裁剪,同一个 batch 里hazy和它的配对clear必须做完全相同的增强。Pytorch 里用torch.rand生成一个随机状态,然后对两张图应用相同变换即可。不要用两次独立RandomCrop,因为两次裁剪位置不同等价于破坏了配对约束,循环一致性损失会瞬间飙升,而且这种不稳定很难从 loss 曲线上直接判断出来。
5.2 用 PSNR/SSIM 指标验证去雾效果
去雾结果不能只靠肉眼判断,尤其是毕设和课程设计答辩时,量化指标比十张效果图更有说服力。PSNR 衡量重建的像素级误差,SSIM 衡量结构相似度,两者对配对测试集都能直接计算。工具可以用 skimage,实现起来非常短:
import cv2 import numpy as np from skimage.metrics import peak_signal_noise_ratio, structural_similarity hazy = cv2.imread("test_data/1408_10.png") clear = cv2.imread("clear/1408_10.png") predict = cv2.imread("output/1408_10.jpg") psnr = peak_signal_noise_ratio(clear, predict) ssim = structural_similarity(clear, predict, channel_axis=2) print(f"PSNR: {psnr:.2f} dB, SSIM: {ssim:.3f}")注意structural_similarity在 skimage 0.19 以上的版本里channel_axis=2取代了旧的multichannel=True,版本不对会直接报参数错误。评估时最好同时计算 hazy 原图的 PSNR/SSIM 作为 baseline,去雾后的指标只有明显高于 hazy 原图,才能说明模型确实起了作用。比如 hazy 对 clear 的 PSNR 如果是 24 dB,预测结果到 28 dB 才算有效提升。
5.3 一个实用小技巧:定期导出中间样本做趋势巡检
训练日志里的 loss 曲线会被平滑掉很多细节,直观可靠的验证方式是在每个 epoch 结束时,拿同一批固定测试图做一次推理,把结果拼成网格图保存。这批测试图不参与训练,相当于盲测。Pytorch 里实现很简单,训练循环末尾加几行:
if epoch % 5 == 0: G_A.eval() with torch.no_grad(): grid = torchvision.utils.make_grid( G_A(fixed_hazy), nrow=4, normalize=True ) torchvision.utils.save_image(grid, f"eval_epoch_{epoch}.jpg") G_A.train()make_grid会把fixed_hazy里多个样本的去雾结果拼在一张图里,normalize=True自动把张量映射到可视范围。我一般每 5 个 epoch 存一次,连续看几张就能发现去雾强度是否过冲:早期图像偏灰、中期通透度上来、后期如果出现色块说明判别器被压制得太狠。相比盯着 loss 猜测,用固定样本的演化趋势判断训练状态,效率高得多,出现模式坍缩也能第一时间发现。
本文还有配套的精品资源,点击获取