
简介基于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_channels3, out_channels3, ngf64, normnn.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, inplaceTrue) ) 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(inplaceTrue) ) 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 保证重建数值稳定。ngf64是基础通道数显存紧张时可降到 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 Tanh2.2 PatchGAN 判别器判断局部真伪而非整图真伪net/Discriminator.py 定义的是 PatchGAN 判别器输出不是单一标量而是一个 30×30 的响应矩阵。每个响应点对应输入图像的一小块感受野生成器必须让每个局部区域都“足够真实”才能骗过判别器。雾在图像上分布不均匀远处浓、近处淡局部判别比整图判别更符合去雾任务的特点。class Discriminator(nn.Module): def __init__(self, in_channels3, ndf64): 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 保留空间位置信息的典型做法。ndf64控制判别器容量容量过大容易让判别器收敛过快容量过小则无法约束生成器。训练时真实图和生成图拼接成同一个 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_A和lambda_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, size256, flipTrue): 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, dtypenp.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), devicedevice) 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_locationdevice) 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, size256): img Image.open(img_path).convert(RGB).resize((size, size), Image.BICUBIC) arr np.array(img, dtypenp.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指定为cpu或cuda解决训练与推理设备不一致的问题。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, multichannelTrue) print(fPSNR: {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()), lr2e-4, betas(0.5, 0.999)) optimizer_D torch.optim.Adam( list(netD_A.parameters()) list(netD_B.parameters()), lr2e-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_locationcpu) 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)用strictFalse可以容忍缺失键但要注意缺失过多说明前缀没剥干净先打印ckpt.keys()确认结构再动手。把这些点逐项核对再回头对比 g_loss、d_loss 曲线和输出图像基本能定位问题出在判别器还是生成器。本文还有配套的精品资源点击获取