DualGAN图像去雾实战:Pytorch对偶生成对抗网络原理与实现
2026/9/14 2:34:43 网站建设 项目流程

简介:基于Pytorch实现的对偶生成对抗网络图像去雾项目,面向计算机相关专业学生及需要项目实战练习的开发者,尤其适合作为课程设计、期末大作业的完整参考。项目经导师指导并获高分评价,涵盖数据加载、模型构建、训练评估、预测去雾等关键环节,代码结构清晰,便于二次开发与实际应用。资源共54个文件,主要包含Python源码(训练、预测、参数解析等脚本)、判别器与生成器模型、训练好的.pkl模型权重、多张测试图像及去雾效果对比图,并附README说明文档,压缩包整体约42.5MB,目录划分明确,方便按需查阅与复现。目前已有60人学习浏览。通过该资源可深入理解对偶生成对抗网络在图像去雾中的完整实现流程,掌握模型训练、参数调优与推理部署的实际操作,同时可基于现有代码快速复现实验,对提升深度学习与图像处理实战能力有直接帮助。

1. 图像去雾为什么需要两个生成器:DualGAN 的开箱体验

拿到这个 Pytorch 源码包,我先用预训练权重对 test_data 里的 1404_7.png 做了一次推理。即便没有配对的清晰参考图,生成器也直接输出了边缘锐利、对比度正常的清晰图,说明网络结构与权重文件是自洽的,没有归一化错位或维度不匹配这类低级问题。图像去雾不是新问题,暗通道先验、AOD-Net 各有解法,但在无配对数据上,对偶生成对抗网络(DualGAN)把去雾当成两个图像域之间的翻译:有雾图是 A 域,清晰图是 B 域,一个生成器负责翻译,另一个负责反翻译,用循环一致性约束内容不丢失。这套源码是导师指导下的高分项目(评审 98 分),从网络定义、训练到预测脚本齐全,不是封装好的黑盒软件,而是能逐行改的 Pytorch 工程,适合课程设计、期末大作业,也适合想从 GAN 实战切入图像修复的开发者。

2. 对偶生成对抗网络模型结构:生成器、判别器与循环一致性的组合

DualGAN 的“对偶”体现在两个生成器构成了一个闭环。单独用一个生成器做有雾到清晰的映射,在没有配对数据时缺少监督信号,生成器很容易产生幻觉,输出一张“看起来清晰但与输入行为无关”的图。两个生成器配合之后,G_AB 生成的清晰图会被 G_BA 重新映射回有雾域,如果重建结果和原始输入一致,说明翻译过程保留了语义内容。这个机制让无配对图像翻译成为可能。

2.1 两个生成器的分工:域翻译与反向还原

net/Generator.py 定义生成器网络,在 dual.py 中实例化 G_AB(有雾转清晰)和 G_BA(清晰转有雾)两个对象。常见做法是二者共享同一结构,输入输出均为 3 通道 RGB 图,内部采用编码器-解码器结构。编码器通过三次步长为 2 的卷积把 256×256 图像压缩到 32×32 特征图,解码器通过三次转置卷积还原到原尺寸,瓶颈处的特征图承载域迁移所需的高层语义。

import torch import torch.nn as nn class Generator(nn.Module): def __init__(self, in_channels=3, out_channels=3, ngf=64, norm=nn.BatchNorm2d): super(Generator, self).__init__() # 编码器:逐步下采样,通道数翻倍 self.enc1 = nn.Conv2d(in_channels, ngf, 4, 2, 1) self.enc2 = self._block(ngf, ngf * 2, norm) self.enc3 = self._block(ngf * 2, ngf * 4, norm) # 解码器:对称上采样,通道数减半 self.dec3 = self._up_block(ngf * 4, ngf * 2, norm) self.dec2 = self._up_block(ngf * 2, ngf, norm) self.dec1 = nn.ConvTranspose2d(ngf, out_channels, 4, 2, 1) self.tanh = nn.Tanh() def _block(self, in_c, out_c, norm): return nn.Sequential( nn.Conv2d(in_c, out_c, 4, 2, 1), norm(out_c), nn.LeakyReLU(0.2, inplace=True) ) def _up_block(self, in_c, out_c, norm): return nn.Sequential( nn.ConvTranspose2d(in_c, out_c, 4, 2, 1), norm(out_c), nn.ReLU(inplace=True) ) def forward(self, x): e1 = torch.relu(self.enc1(x)) e2 = self.enc2(e1) e3 = self.enc3(e2) d3 = self.dec3(e3) d2 = self.dec2(d3) d1 = self.tanh(self.dec1(d2)) return d1

生成器内部不对输入重复归一化,因为数据加载阶段已经把像素映射到 [-1,1],输出层用 Tanh 与之匹配。编码器使用 LeakyReLU(0.2) 避免负区间梯度消失,解码器用 ReLU 保证重建数值稳定。ngf=64是基础通道数,显存紧张时可降到 32,代价是细节还原能力下降。

模块层名输出尺寸说明
编码器enc1128×128×64Conv2d + ReLU
编码器enc264×64×128Conv2d + BN + LeakyReLU
编码器enc332×32×256Conv2d + BN + LeakyReLU
解码器dec364×64×128ConvTranspose2d + BN + ReLU
解码器dec2128×128×64ConvTranspose2d + BN + ReLU
解码器dec1256×256×3ConvTranspose2d + Tanh

2.2 PatchGAN 判别器:判断局部真伪而非整图真伪

net/Discriminator.py 定义的是 PatchGAN 判别器,输出不是单一标量,而是一个 30×30 的响应矩阵。每个响应点对应输入图像的一小块感受野,生成器必须让每个局部区域都“足够真实”才能骗过判别器。雾在图像上分布不均匀,远处浓、近处淡,局部判别比整图判别更符合去雾任务的特点。

class Discriminator(nn.Module): def __init__(self, in_channels=3, ndf=64): super(Discriminator, self).__init__() # 前两层步长为 2,压缩空间分辨率 self.conv1 = nn.Conv2d(in_channels, ndf, 4, 2, 1) self.conv2 = nn.Conv2d(ndf, ndf * 2, 4, 2, 1) self.bn2 = nn.BatchNorm2d(ndf * 2) # 最后一层步长为 1,输出 Patch 矩阵 self.conv3 = nn.Conv2d(ndf * 2, 1, 4, 1, 1) def forward(self, x): out = torch.relu(self.conv1(x)) out = torch.relu(self.bn2(self.conv2(out))) out = torch.sigmoid(self.conv3(out)) return out

判别器没有池化层,分辨率全靠步长卷积压缩,这是 PatchGAN 保留空间位置信息的典型做法。ndf=64控制判别器容量,容量过大容易让判别器收敛过快,容量过小则无法约束生成器。训练时真实图和生成图拼接成同一个 batch 喂入,矩阵的每个元素经过 Sigmoid 后落在 (0,1),表示该局部区域为真实样本的概率。

2.3 循环一致性损失:无配对数据下的监督来源

对抗损失只能让生成图“看起来像清晰图”,无法保证它与输入有雾图内容一致。DualGAN 用循环一致性损失补上这个缺口:对真实有雾图 A,先由 G_AB 得到伪清晰图 B',再由 G_BA 把 B' 映射回有雾域,得到重建图 A';同理,真实清晰图 B 经过 G_BA 再经过 G_AB,应还原为 B'。这个双向闭环正是“对偶”二字的来源。

# dual.py 中损失计算的核心片段 criterion_cycle = torch.nn.L1Loss() lambda_A = 10.0 # 有雾 -> 清晰 -> 有雾 循环权重 lambda_B = 10.0 # 清晰 -> 有雾 -> 清晰 循环权重 # 正向循环:A -> G_AB(A) -> G_BA(G_AB(A)) fake_B = netG_AB(real_A) rec_A = netG_BA(fake_B) cycle_loss_A = criterion_cycle(rec_A, real_A) * lambda_A # 反向循环:B -> G_BA(B) -> G_AB(G_BA(B)) fake_A = netG_BA(real_B) rec_B = netG_AB(fake_A) cycle_loss_B = criterion_cycle(rec_B, real_B) * lambda_B cycle_loss = cycle_loss_A + cycle_loss_B

循环一致性损失用 L1 而不是 L2,因为 L1 对边缘梯度更友好,重建结果更锐利,L2 会把模棱两可的像素平均化,导致重建图偏模糊。lambda_Alambda_B控制内容保持与对抗博弈的平衡,项目默认取 10,这个数值在多数室内外去雾场景下都能稳定收敛。若生成图出现偏色,可适当增大循环权重;若生成图过于平滑、缺少纹理细节,说明循环损失过强,需要往 5 的方向调低。

3. Pytorch 训练流程:从数据加载、对抗更新到权重保存

训练对偶生成对抗网络,关注点不在“能跑通”,而在判别器和生成器的更新节奏。常见错误是判别器收敛过快,生成器完全学不到东西;或者循环一致性权重过高,输出退化成输入的轻微提亮。下面以项目中的 train.py、util/loader.py、util/parseArgs.py 为线索拆解。

3.1 数据加载与预处理:loader.py 的关键操作

util/loader.py 负责读取两个图像域目录,返回可迭代的 Dataset。常见做法是把有雾图放在一个目录、清晰图放在另一个目录,自定义 Dataset 的__getitem__独立读取,不要求一一配对。预处理环节包括 resize 到 256×256、随机水平翻转、归一化到 [-1,1]。翻转能有效扩充有雾样本,因为雾的分布对手性不敏感。

# util/loader.py 核心逻辑 import torch from torch.utils.data import Dataset from PIL import Image import numpy as np class UnpairedDataset(Dataset): def __init__(self, dir_A, dir_B, size=256, flip=True): self.files_A = sorted(make_dataset(dir_A)) # A 域:有雾图 self.files_B = sorted(make_dataset(dir_B)) # B 域:清晰图 self.size = size self.flip = flip def __getitem__(self, idx): img_A = self._load(self.files_A[idx % len(self.files_A)]) img_B = self._load(self.files_B[idx % len(self.files_B)]) if self.flip and np.random.rand() > 0.5: img_A = img_A.transpose(Image.FLIP_LEFT_RIGHT) img_B = img_B.transpose(Image.FLIP_LEFT_RIGHT) return img_A, img_B def _load(self, path): img = Image.open(path).convert('RGB').resize((self.size, self.size)) arr = np.array(img, dtype=np.float32) / 127.5 - 1.0 return torch.from_numpy(arr).permute(2, 0, 1)

这里默认把像素从 [0,255] 映射到 [-1,1],与生成器输出的 Tanh 激活函数区间匹配。如果训练用这套归一化,推理时也必须用同一套,否则输出要么整体偏亮要么偏暗。idx % len(...)的写法让两个域样本数量不一致时也能继续训练,但要注意每个 epoch 里数量多的域会被重复采样,属正常现象。

3.2 训练主循环:生成器与判别器的交替更新

train.py 在每一轮迭代中先更新判别器,再更新生成器。判别器更新的目标是拉大真实样本与生成样本的输出差异;生成器更新要同时骗过两个判别器,并满足循环一致性约束。这里的对抗损失采用 LSGAN 的 MSE 形式,训练比原始 GAN 的 log 损失更稳定。

# train.py 单个 batch 的训练流程 for epoch in range(opt.epochs): for i, (real_A, real_B) in enumerate(train_loader): real_A, real_B = real_A.to(device), real_B.to(device) # 判别器输出 30x30 Patch,标签需对齐 real_label = torch.ones((real_A.size(0), 1, 30, 30), device=device) fake_label = torch.zeros_like(real_label) # 第一步:更新判别器 D_A 和 D_B fake_B = netG_AB(real_A) fake_A = netG_BA(real_B) dA_loss = criterion_D(netD_A(real_A), real_label) + \ criterion_D(netD_A(fake_A.detach()), fake_label) dB_loss = criterion_D(netD_B(real_B), real_label) + \ criterion_D(netD_B(fake_B.detach()), fake_label) d_loss = 0.5 * (dA_loss + dB_loss) optimizer_D.zero_grad() d_loss.backward() optimizer_D.step() # 第二步:更新生成器 G_AB 和 G_BA fake_B = netG_AB(real_A) fake_A = netG_BA(real_B) rec_A = netG_BA(fake_B) rec_B = netG_AB(fake_A) g_loss_adv = criterion_D(netD_B(fake_B), real_label) + \ criterion_D(netD_A(fake_A), real_label) g_loss_cyc = criterion_CYC(rec_A, real_A) * opt.lambda_A + \ criterion_CYC(rec_B, real_B) * opt.lambda_B g_loss = g_loss_adv + g_loss_cyc optimizer_G.zero_grad() g_loss.backward() optimizer_G.step()

判别器更新时fake_A.detach()fake_B.detach()必不可少,目的是阻断对抗梯度回流到生成器,否则判别器和生成器会在同一步内互相拉扯,损失曲线剧烈震荡。标签张量的形状(batch, 1, 30, 30)必须与netD_B输出严格一致,如果改了输入尺寸或判别器步长,这里的 30 要同步换算,这是最容易被隐藏的维度坑。

3.3 超参数配置与训练监控

util/parseArgs.py 用 argparse 暴露训练参数,典型配置如下表。Pytorch 环境搭建好之后,安装好依赖直接执行训练命令即可。

参数默认值作用
--epochs200总训练轮数
--batch-size4单卡建议值,显存不足可降到 2
--lr0.0002Adam 初始学习率
--beta10.5Adam 一阶矩衰减系数
--lambda-A10.0正向循环损失权重
--lambda-B10.0反向循环损失权重
--size256训练图像尺寸
python train.py --epochs 200 --batch-size 4 --lr 0.0002

训练期间用 util/logger.py 记录 generator loss、discriminator loss 和 cycle loss。当 g_loss 不再下降而 d_loss 持续走低时,往往不是训练完成,而是判别器过强、生成器梯度消失的信号。

提示:正式训练前先用一个 batch 跑通前向和反向,确认损失都能正常反传,再启动完整训练,能省掉大半调试时间。

4. 去雾推理与效果评估:predict.py 从权重到清晰图

训练完成后,推理本身不复杂,但有一个容易被忽略的环节:训练时的归一化与推理时的归一化必须完全一致。项目里 predict.py 负责加载生成器,读取 test_data 中的有雾图,把输出写到 predict 目录。

4.1 predict.py 推理流程

权重文件以 .pkl 形式保存。加载时先确认 checkpoint 的键结构:是每个模块单独保存,还是把 G_AB、G_BA、D_A、D_B 打包进一个字典。下面给出兼容两种情况的写法。

# predict.py 核心逻辑 import torch from PIL import Image import numpy as np from net.Generator import Generator device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') netG_AB = Generator(3, 3).to(device) ckpt = torch.load('model/dual_gan.pkl', map_location=device) if isinstance(ckpt, dict) and 'G_AB' in ckpt.keys(): netG_AB.load_state_dict(ckpt['G_AB']) else: netG_AB.load_state_dict(ckpt) netG_AB.eval() def preprocess(img_path, size=256): img = Image.open(img_path).convert('RGB').resize((size, size), Image.BICUBIC) arr = np.array(img, dtype=np.float32) / 127.5 - 1.0 return torch.from_numpy(arr).permute(2, 0, 1).unsqueeze(0) def postprocess(tensor): arr = tensor.squeeze(0).permute(1, 2, 0).detach().cpu().numpy() arr = np.clip((arr + 1.0) * 127.5, 0, 255).astype(np.uint8) return Image.fromarray(arr) with torch.no_grad(): x = preprocess('test_data/1404_7.png').to(device) out = netG_AB(x) postprocess(out).save('predict/1404_7.jpg')

map_location指定为cpucuda,解决训练与推理设备不一致的问题。netG_AB.eval()会关闭 dropout 和 BatchNorm 的统计更新,这点在生成器含 BN 层时尤其重要。resize 插值方式使用 BICUBIC,需要与训练保持一致,否则细节区域可能出现轻微振铃。推理全程放在torch.no_grad()下,避免构建计算图浪费显存。

4.2 用 PSNR 与 SSIM 验证去雾效果

如果测试集带有配对清晰图,可以用 PSNR 和 SSIM 做量化评估。PSNR 衡量像素级误差,SSIM 衡量结构相似性,两个指标结合能避免单一指标被少量极端像素带偏。没有参考图时,改用 BRISQUE 这类无参考指标或直接目视检查。

# 评估脚本片段 import math import cv2 import numpy as np from skimage.metrics import structural_similarity as ssim def psnr(img1, img2): mse = np.mean((img1.astype(np.float64) - img2.astype(np.float64)) ** 2) if mse == 0: return float('inf') return 20 * math.log10(255.0 / math.sqrt(mse)) gt = cv2.imread('data/clear/1404_7.png') pred = cv2.imread('predict/1404_7.jpg') psnr_val = psnr(gt, pred) ssim_val = ssim(gt, pred, multichannel=True) print(f'PSNR: {psnr_val:.2f} dB, SSIM: {ssim_val:.4f}')
评估方式是否需要参考图适用场景
PSNR有配对测试集,衡量像素重建误差
SSIM有配对测试集,衡量结构信息保持
BRISQUE真实场景无参考图,衡量图像自然度

PSNR 对整体亮度偏移敏感,SSIM 对局部结构变化敏感。在去雾场景中,输出比参考图稍亮可能使 PSNR 下降很多,但视觉上反而更通透,因此不能只看单一指标。如果输出图出现偏蓝或偏灰,通常是对抗损失权重过高、生成器牺牲颜色换取域分布逼近的结果。

5. 调参与踩坑:从权重文件到稳定收敛的细节

对偶生成对抗网络在 Pytorch 里调参,有几个反复出现的坑值得单独记下来。

5.1 判别器收敛过快

损失曲线上 d_loss 迅速压到接近 0、g_loss 却停在原地,就是判别器太强的信号。常见处理:把判别器学习率下调到生成器的十分之一,用两个独立优化器分组管理。

optimizer_G = torch.optim.Adam( list(netG_AB.parameters()) + list(netG_BA.parameters()), lr=2e-4, betas=(0.5, 0.999)) optimizer_D = torch.optim.Adam( list(netD_A.parameters()) + list(netD_B.parameters()), lr=2e-5, betas=(0.5, 0.999))

也可以改成每迭代两次生成器才更新一次判别器,给生成器更多追赶时间。标签平滑同样有效:real_label 用 0.9、fake_label 用 0.1,降低判别器置信度过冲,训练过程会更稳。

5.2 循环一致性权重与身份损失

lambda 默认 10 在多数场景可用,但不同数据集差异很大。如果去雾不彻底,输出只是轻微提亮,说明循环约束过强、生成器不敢做大幅迁移,把 lambda 降到 5。如果出现伪影或色斑,再往 15 方向提高。此外,去雾任务加一个身份损失能明显改善颜色保持。

identity_loss = criterion_CYC(netG_AB(real_B), real_B) * 5 + \ criterion_CYC(netG_BA(real_A), real_A) * 5

身份损失的直观含义是:输入已经清晰的图,生成的清晰图应该基本等于输入,约束生成器不要随意改动本已清晰的区域。颜色敏感场景下,这个损失值得保留,代价是生成图对比度会略保守。

5.3 权重保存与加载的兼容性

.pkl 文件保存的是 state_dict,不是完整模型对象。用torch.save(model.state_dict(), path)保存,用load_state_dict加载。如果网络结构或键名前缀不一致,Pytorch 会抛键名不匹配,这是最容易排查的问题。若训练脚本用了 DataParallel,权重键名会带module.前缀,加载时需要剥掉:

from collections import OrderedDict raw = torch.load('model/dual_gan.pkl', map_location='cpu') clean = OrderedDict((k.replace('module.', ''), v) for k, v in raw.items()) netG_AB.load_state_dict(clean['G_AB'] if 'G_AB' in clean else clean)

strict=False可以容忍缺失键,但要注意缺失过多说明前缀没剥干净,先打印ckpt.keys()确认结构再动手。把这些点逐项核对,再回头对比 g_loss、d_loss 曲线和输出图像,基本能定位问题出在判别器还是生成器。

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

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

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

立即咨询