简介图像去雾是计算机视觉中的经典难题其核心在于从退化图像中恢复清晰场景。传统暗通道先验依赖物理模型在复杂场景下易产生人工痕迹而基于深度学习的生成对抗网络GAN通过生成器与判别器的对抗博弈能够学习更真实的纹理分布成为图像恢复领域的重要技术路径。在Pytorch框架下对偶生成对抗网络利用循环一致性约束进一步提升了去雾模型的泛化能力尤其适用于合成数据与真实场景的适配。本文围绕Pytorch环境搭建、数据集处理、U-Net与PatchGAN结构设计、损失函数组合及训练技巧展开系统梳理了图像去雾项目从训练到推理部署的全流程并总结了显存优化、模式崩溃、颜色偏移等实战问题。无论是入门深度学习图像处理还是工程落地这套实践方案都提供了可复用的经验。 做图像去雾这个方向也有一段时间了从传统暗通道先验一路摸到深度学习最后真正落地用的还是这套基于Pytorch的对偶生成对抗网络方案。单纯做学术demo的话训练出能看的模型并不难难在数据预处理、网络结构设计、损失函数搭配和训练稳定性的平衡上。这篇文章把整个项目的核心拆开揉碎从环境搭建到推理部署讲清楚顺便把所有踩过的坑都整理出来希望对打算入坑GAN去雾的朋友有实际帮助。对于要复现这个项目的读者默认你至少会Python基础语法、懂一点卷积神经网络的概念并且能在本地或者服务器上把Pytorch跑起来。只要具备这些下面内容可以按顺序一步步操作不需要额外补充太多前置知识。1. 项目背景与核心思路1.1 为什么图像去雾适合用生成对抗网络图像去雾本质上是图像到图像的翻译问题输入是带雾图像输出是清晰无雾图像。早期方法主要依赖物理散射模型比如暗通道先验通过估计透射率和大气光来反演清晰图像。这种方法在没有雾的平坦区域经常失效而且处理高分辨率图像时的计算量非常大恢复出来的人造痕迹也很明显。后来基于深度学习的方法兴起用卷积神经网络直接回归映射关系比如MSCNN、AOD-Net这类模型。它们的思路是把去雾当作一个回归任务训练时用L2损失或L1损失来约束输出和真实清晰图像接近。这样做的优点是训练稳定、推理快但问题也很突出——回归损失倾向于生成平均化的结果细节纹理容易被平滑掉看起来像蒙了一层灰。GAN生成对抗网络加入之后相当于给去雾加了一个“质量评委”。生成器负责从有雾图像恢复出清晰图判别器负责判断输出是真实照片还是生成结果。两者相互博弈生成器被迫去学习更真实、更锐利的纹理分布而不是简单求一个平均值。这个思路和超分、图像修复、风格迁移领域的GAN应用是共通的都依赖对抗损失把输出逼到真实数据的流形上。这里说的对偶生成对抗网络指的是建立两个互为逆向的生成器雾图到清晰图是一个方向清晰图到雾图是另一个方向形成闭环约束。这样即使整理训练数据时没办法做到百分百像素级配对也能通过循环一致性损失让两个映射都保持语义一致。实际项目中我用的数据多是人工合成雾图虽然有配对但加入对偶结构后生成器的泛化能力明显更稳定野外真实雾图上的虚化感和色偏也少很多。1.2 基于Pytorch实现的技术选型思路选Pytorch而不是TensorFlow主要三个原因。第一Pytorch的动态图机制在自定义GAN结构时特别舒服生成器、判别器、损失函数都是普通Python对象调试时可以随时打断点检查中间张量不会有静态图那种“先构图后执行”的割裂感。第二社区生态对GAN研究向的项目支持很成熟很多预训练模型、数据加载器和训练trick都有现成实现可以用。第三在写这个项目时Pytorch已经支持分布式训练和AMP混合精度显存不够的时候可以灵活调整。网络结构上生成器选了带有跳跃连接的U-Net变体深层特征加残差块来增强感受野。判别器用了PatchGAN它的输出不是一个标量而是一个N×N矩阵每个元素对应图像局部区域的真实性判断。相比普通判别器只输出0到1的全局概率PatchGAN能更细腻地约束局部纹理对去雾这种“全局提亮局部还原”的任务更友好。损失函数方面没有只用对抗损失而是凑了三部分对抗损失让输出掉进清晰图像的分布域循环一致性损失维持雾图和无雾图之间的内容结构感知损失用预训练的VGG网络提取高层特征来计算差异保证视觉感知上的一致性。三者加权组合训练出的模型在指标和观感上才比较均衡。2. 环境准备与数据集处理2.1 Pytorch与CUDA环境搭建最基本的依赖就是Pytorch配合对应的CUDA版本。如果电脑有NVIDIA显卡建议直接安装GPU版本否则纯CPU训练一个周期就要等很久基本没法实际用。安装时先确认自己的显卡驱动版本和CUDA版本然后再选择对应的Pytorch安装命令。比如在Linux环境里可以用下面这个方式安装# 先查看显卡驱动支持的CUDA版本 nvidia-smi # 安装Pytorch注意选择和CUDA匹配的版本 pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu118这里要提醒的是Pytorch安装包自带CUDA运行时所以并不需要你单独安装完整的CUDA Toolkit。只要显卡驱动版本足够新nvidia-smi看到的CUDA版本大于等于你安装的cuda运行时版本就行。Windows环境的话要注意Python版本和Pytorch版本的兼容建议直接用Python 3.8到3.10之间的版本太新或太旧都可能出现import错误。除了Pytorch本体还需要装几个配套库OpenCV图像读取和预处理、NumPy数组操作、Matplotlib可视化画图、tqdm训练进度条。这些可以一次性装完。pip install opencv-python numpy matplotlib tqdm2.2 有雾/无雾数据集准备训练数据是去雾项目最重要的部分没有之一。最理想的情况是同一场景同时拍摄有雾和无雾照片但真实场景很难等来完全一致的光照和雾气条件所以绝大多数研究项目都使用合成雾图。经典的方案是用NYU-Depth-V2数据集它提供室内场景的深度图配合大气散射模型可以生成合成的有雾图像。大气散射模型的公式是[ I(x) J(x)t(x) A(1 - t(x)) ]其中 ( J(x) ) 是清晰图像( t(x) e^{-\beta d(x)} ) 是透射率由大气散射系数 ( \beta ) 和深度 ( d(x) ) 决定( A ) 是全局大气光。实际操作中随机在多个 ( \beta ) 值下生成不同浓度的雾比如从0.5到1.5之间均匀采样再把A设成接近1的随机RGB值这样能模拟从薄雾到浓雾的多种情况提高模型的泛化能力。如果你手头没有深度图也可以用RESIDE这类公开数据集它提供了大量室内外场景的雾图和对应的清晰图直接下载解压就能用。数据目录结构尽量按项目标准方式组织data/ ├── train/ │ ├── hazy/ │ │ ├── 000001.png │ │ ├── 000002.png │ │ └── ... │ └── clear/ │ ├── 000001.png │ ├── 000002.png │ └── ... └── val/ ├── hazy/ └── clear/数据加载时不要直接原图丢进网络先做预处理统一缩放到256×256或512×512随机水平翻转、随机裁剪、颜色抖动。这些数据增强手段虽然简单但能明显提升模型的泛化能力尤其是随机翻转和裁剪几乎零成本。要注意的是颜色抖动需要谨慎因为去雾本身和颜色相关增强幅度过大会导致色偏。3. 对偶GAN网络结构设计3.1 生成器结构U-Net加残差块生成器是对偶GAN里最关键的部分它决定了去雾结果的天花板。我用的是U-Net形状的编解码结构编码器部分逐层提取特征并压缩空间尺寸解码器部分逐步恢复分辨率。为了让高低层特征能互相流动在每层之间添加了跳跃连接。这个设计对图像的局部结构恢复特别有效雾气的边缘和物体的轮廓都能保留得更完整。在编码器和解码器之间的瓶颈区域我串联了多个残差块。残差块的好处是可以让梯度从深层网络直接回流避免网络太深时梯度消失。每个残差块由两层卷积加BN加ReLU组成通过跨层连接把输入和输出加到一起。这里贴一个简化版的残差块实现import torch.nn as nn class ResidualBlock(nn.Module): def __init__(self, in_channels): super(ResidualBlock, self).__init__() self.conv1 nn.Conv2d(in_channels, in_channels, kernel_size3, padding1) self.bn1 nn.BatchNorm2d(in_channels) self.relu nn.ReLU(inplaceTrue) self.conv2 nn.Conv2d(in_channels, in_channels, kernel_size3, padding1) self.bn2 nn.BatchNorm2d(in_channels) def forward(self, x): identity x out self.conv1(x) out self.bn1(out) out self.relu(out) out self.conv2(out) out self.bn2(out) return self.relu(out identity)生成器的输入是3通道有雾图输出也是3通道无雾图所以最后一层卷积后不需要加BN和ReLU直接用Tanh激活函数把像素值归一化到-1到1之间。这个细节很关键因为Pytorch里图像如果归一化到0到1网络收敛速度会慢和感知损失、判别器输入的分布也不匹配。3.2 判别器结构PatchGAN局部判断判别器直接决定生成器是不是在“骗人”。如果用普通CNN判别器输出一个全局标量它能捕捉到整张图像的整体风格但对局部区域容易出现判别盲区。PatchGAN的结构是把输入图像切分成多个Patch每个Patch独立判断真假。这里的切分不是真的把图像裁开而是通过堆叠卷积层让最后一个特征图的每个神经元对应输入图像的一个感受野区域。PatchGAN的实现非常简洁核心就是几层卷积加LeakyReLU步长设为2来做下采样最后输出一个张量比如16×16×1。这个16×16矩阵中的每个值都代表输入图像中某个区域的真实性。用这个结构做对抗训练生成器需要同时保证每个局部区域都足够真实细节纹理自然就被逼出来了。class PatchDiscriminator(nn.Module): def __init__(self, in_channels3, base_channels64): super(PatchDiscriminator, self).__init__() self.model nn.Sequential( nn.Conv2d(in_channels, base_channels, kernel_size4, stride2, padding1), nn.LeakyReLU(0.2, inplaceTrue), nn.Conv2d(base_channels, base_channels * 2, kernel_size4, stride2, padding1), nn.BatchNorm2d(base_channels * 2), nn.LeakyReLU(0.2, inplaceTrue), nn.Conv2d(base_channels * 2, base_channels * 4, kernel_size4, stride2, padding1), nn.BatchNorm2d(base_channels * 4), nn.LeakyReLU(0.2, inplaceTrue), nn.Conv2d(base_channels * 4, 1, kernel_size4, stride1, padding1) ) def forward(self, x): return self.model(x)3.3 对偶网络的两个生成器设计对偶结构里有两个生成器( G_{H \to C} ) 负责把有雾图变成清晰图( G_{C \to H} ) 负责把清晰图变成有雾图。两个生成器的网络结构完全相同但参数相互独立也就是说一共要训练四套网络参数。训练时有雾图 ( h ) 先经 ( G_{H \to C} ) 得到伪清晰图 ( c_{fake} )再经 ( G_{C \to H} ) 重建回有雾图 ( h_{rec} )用循环一致性损失约束 ( h_{rec} ) 和 ( h ) 尽量一致。反过来清晰图也走一遍同样的流程。这种对偶结构的价值在于它相当于同时在两个方向上约束映射关系即便缺少严格配对的训练数据也能保持内容信息的完整性。对去雾任务来说这意味着模型不容易在去除雾气的同时丢失边缘和颜色信息输出的清晰图会更自然。4. 损失函数设计与训练策略4.1 对抗损失、循环一致性损失和感知损失的权重组合项目最终使用的损失函数是三部分加权求和。对抗损失采用最小二乘GAN形式判别器输出后直接计算均方误差。相比传统二元交叉熵的对抗损失最小二乘损失可以缓解训练时梯度消失的问题尤其在判别器已经很强的情况下生成器的梯度依然足够。循环一致性损失用的是L1范数对比L2它不惩罚太大误差的平方所以恢复出来的图像边缘更锐利。感知损失则是把生成图和真实图分别送入预训练的VGG19网络取出中间层的特征图计算它们之间的L1距离。这里贴出感知损失的简单实现import torch import torch.nn.functional as F from torchvision import models class PerceptualLoss(nn.Module): def __init__(self): super(PerceptualLoss, self).__init__() vgg models.vgg19(pretrainedTrue).features.eval() for param in vgg.parameters(): param.requires_grad False self.layers nn.Sequential(*list(vgg)[:20]).cuda() def forward(self, pred, target): pred_feat self.layers(pred) target_feat self.layers(target) return F.l1_loss(pred_feat, target_feat)三个损失项的具体权重我最后调成了对抗损失权重1.0循环一致性损失10.0感知损失5.0。循环一致性损失权重最大因为它是保证图像内容保真的主力。感知损失次之提供高层语义约束。对抗损失虽然权重最小但它决定图像纹理的“真实感”缺了它会明显感觉输出图像偏平滑。4.2 训练参数与学习率策略优化器我选了Adam生成器和判别器的初始学习率都设为0.0002指数衰减率beta1取0.5。这里和常规分类任务的Adam参数不同GAN训练中beta1建议设小一些让优化器不要累积过多历史梯度能更快响应生成器和判别器之间的动态变化。批量大小视显卡显存而定我用RTX 3090时设为8输入图片分辨率256×256。如果你显存只有8G建议批量大小降到4甚至2分辨率也可以下调到224×224否则很容易OOM。训练周期设置为80个epoch学习率在前40个epoch保持不变后面40个epoch线性衰减到0。这种固定后再衰减的学习率策略是GAN训练里的常见做法目的是前半段让模型充分探索后半段稳定收敛。训练时需要每间隔一定迭代次数保存一次模型的checkpoint建议保存内容包括生成器、判别器、优化器状态和当前epoch/iteration。这样即使训练中断也能恢复现场不至于几天的训练白跑。5. 训练循环实现与监控5.1 训练主循环代码解析对偶GAN的训练流程比普通GAN复杂一点每次迭代要交替更新两次生成器方向和两次判别器方向。核心流程如下从数据加载器取一批有雾图和清晰图把它们归一化到[-1,1]送入网络。判别器的训练里直接用真实清晰图训练判别器D_C把由有雾图生成的伪清晰图当作假样本训练。这里有个细节生成器的梯度只在更新生成器时计算更新判别器时要把生成器的梯度冻结通常用detach()方法切掉梯度回传路径。更新生成器时除了对抗损失还要把循环重建的结果拿去算循环一致性损失加上感知损失。由于两个生成器的参数都要更新所以总损失要同时回传到两个生成器的计算图里。关键代码逻辑如下# 判别器D_C训练 fake_clear gen_H2C(hazy) loss_dc criterion_GAN(disc_C(clear), valid) \ criterion_GAN(disc_C(fake_clear.detach()), fake) disc_C_optimizer.zero_grad() loss_dc.backward() disc_C_optimizer.step() # 生成器训练 fake_clear gen_H2C(hazy) rec_hazy gen_C2H(fake_clear) loss_adv criterion_GAN(disc_C(fake_clear), valid) loss_cycle criterion_L1(rec_hazy, hazy) criterion_L1(rec_clear, clear) loss_percep perceptual_loss(fake_clear, clear) loss_gen loss_adv 10.0 * loss_cycle 5.0 * loss_percep gen_optimizer.zero_grad() loss_gen.backward() gen_optimizer.step()这里要特别注意训练顺序。我习惯先更新判别器再更新生成器模拟真实对抗过程。如果你发现loss值波动很剧烈说明判别器和生成器节奏失配可以试试每更新一次判别器就更新两次生成器或者反过来这种调节方式比改学习率更直接。5.2 训练过程监控与Loss曲线判读训练GAN不像普通分类任务那样容易判断收敛loss曲线也不是越小越好。关键看生成器和判别器的对抗是否处于动态平衡。理论上二者会收敛到一个纳什均衡点但实际训练中经常出现震荡。我在训练时主要查两样东西。第一周期性把生成器的输出图像保存到本地比如每500个iteration保存一张对比图清晰度和颜色都能直观看到变化。第二记录三个损失分量各自的值画出曲线。如果对抗损失一直居高不下可能判别器太强或者生成器容量不够如果感知损失降到很低但直观效果很差可能是感知损失权重太大把生成器推向“特征匹配”但不真实的解。实用技巧方面推荐使用TensorBoard或者wandb。Pytorch自带SummaryWriter少量代码就能记录图像和标量from torch.utils.tensorboard import SummaryWriter writer SummaryWriter(logs/train) writer.add_image(result/fake_clear, (fake_clear[0].detach().cpu() 1) / 2, global_step) writer.add_scalar(loss/gen_total, loss_gen.item(), global_step)6. 推理部署与效果调优6.1 模型导出与去雾推理流程训练完进入推理阶段需要把模型从训练状态切到eval模式这一步很容易忽略。因为BN和Dropout在训练和推理时的行为不同不切eval模式的话BN层会用batch内的统计量导致输出出现异常色块。更合理的做法是加载训练过程中保存的最好的checkpoint再做一次推理。推理代码的核心就是读取模型权重、加载图像、预处理、前向计算、后处理。这里贴一个完整的推理脚本import torch import cv2 import numpy as np from torchvision import transforms def inference(model_path, image_path, output_path, devicecuda): model build_generator() model.load_state_dict(torch.load(model_path, map_locationdevice)) model.to(device).eval() img cv2.imread(image_path) img cv2.cvtColor(img, cv2.COLOR_BGR2RGB) img cv2.resize(img, (256, 256)) img_tensor transforms.ToTensor()(img).unsqueeze(0) img_tensor img_tensor * 2.0 - 1.0 with torch.no_grad(): output model(img_tensor.to(device)) output (output.squeeze().cpu().numpy() 1) / 2.0 output np.transpose(output, (1, 2, 0)) output (output * 255).astype(np.uint8) output cv2.cvtColor(output, cv2.COLOR_RGB2BGR) cv2.imwrite(output_path, output)推理阶段的分辨率有讲究。如果训练时用256×256推理时直接对高分辨率大图输出模型可能产生严重的块状伪影。一般来说可以先将原图缩放到训练分辨率推理再把结果resize回原尺寸或者用滑窗方式在局部块上分块推理。去雾任务里我建议用前一种方式因为滑窗方式在拼接处容易产生不一致的亮度变化。6.2 主观与客观指标评估项目效果不能只看眼睛还需要客观指标。图像去雾领域常用的指标有PSNR和SSIM。PSNR衡量像素级误差值越高越好SSIM衡量结构相似性越接近1越好。如果数据集有对应的清晰图直接计算这两个指标。代码短小简单import cv2 from skimage.metrics import peak_signal_noise_ratio, structural_similarity gt cv2.cvtColor(cv2.imread(clear.png), cv2.COLOR_BGR2GRAY) out cv2.cvtColor(cv2.imread(output.png), cv2.COLOR_BGR2GRAY) print(PSNR:, peak_signal_noise_ratio(gt, out)) print(SSIM:, structural_similarity(gt, out))如果数据集没有清晰图比如户外真实雾图就只能做主观评估了。几个常见的观测点天空区域是否出现过度增强导致色带白色物体是否忠实还原边缘区域有没有光晕伪影整体对比度是否自然。7. 常见问题与排查技巧实录7.1 训练不收敛或模式崩溃GAN训练最常见的问题是模式崩溃典型表现是生成器输出图像长时间变化很小或者出现大片重复纹理。排查步骤一般是先调判别器太弱的判别器会给不了生成器足够的学习压力太强的判别器又会把生成器梯度打没。我的经验是优先调整训练顺序比如判别器每更新一次生成器更新两次这个做法在很多GAN实战项目里都有效。如果发现对抗损失在很小的值附近震荡生成图像虽然真实但是和目标关系不大这也属于模式崩溃的一种可以尝试增大循环一致性损失的权重强制生成结果和输入保持内容一致性。另外适当加大判别器网络dropout的比例也能缓解此类问题。7.2 显存不足与训练速度慢去雾训练输入图像的分辨率和batch大小直接影响显存。遇到OOM时除了降低batch size还可以尝试开启混合精度训练。AMP自动选择在GPU上使用半精度计算能在几乎不损失效果的前提下减少一半左右的显存占用。使用方法很简单from torch.cuda.amp import autocast, GradScaler scaler GradScaler() with autocast(): loss criterion(model(inputs), targets) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()Pytorch官方的AMP方案在Pytorch 1.6以后已经非常好用不需要手动改网络结构。如果用了AMP发现某些层输出Nan可以先把动态损失缩放打开通常能解决。7.3 图像颜色偏灰或偏暗这种情况多半是数据归一化或生成器输出激活函数不匹配导致的。生成器最后一层如果用Sigmoid输出范围是[0,1]那输入图像也要归一化到[0,1]如果用Tanh输出范围是[-1,1]输入图像也要对应归一化到[-1,1]。一旦两边不匹配就会出现颜色整体偏移。另外训练数据里的有雾图和清晰图如果色域不一致比如一个有偏蓝调一个有暖调模型会学出颜色映射偏差。可以通过白色均衡预处理来统一色域或者人为在数据增强里加入颜色扰动让网络不依赖特定色偏。8. 项目可扩展方向与部署建议8.1 从Pytorch模型到实际服务部署训练好的Pytorch模型如果只跑本地推理能发挥的价值有限。部署到实际应用场景时一个常规思路是导出为ONNX格式再转换为TensorRT加速推理或者直接在Pytorch里用torch.jit做TorchScript导出。ONNX的好处是模型格式标准化后续可以部署到不同深度学习框架和边缘设备上。导出ONNX的代码大致如下model.eval() dummy_input torch.randn(1, 3, 256, 256).cuda() torch.onnx.export( model, dummy_input, dehaze.onnx, input_names[input], output_names[output], dynamic_axes{input: {0: batch}, output: {0: batch}} )8.2 在真实雾图场景的适应策略实际使用中模型在合成雾图上训练后在真实雾图上表现经常打折扣。一个切实可行的方法是在推理阶段对输入图像做一定的增强预处理比如先适当提高对比度再输入模型推理完再做后处理降噪。另一个思路是收集少量真实雾图用大小合适的patch和人工标注的参考图像做微调不需要重新训练整个数据集通常只需几十张图就能明显提升泛化效果。8.3 项目继续优化的方向这套项目还能向多个方向延伸。比如把生成器换成当前更流行的Vision Transformer结构或者引入多头注意力机制提升全局信息捕捉能力加入感知驱动的边缘损失或者暗通道损失作为额外约束用知识蒸馏的方式把大模型压缩成轻量化模型方便在移动端实时运行。这些方向都是在现有框架基础上做局部替换不会推翻整套项目设计。最后分享一个小技巧训练GAN千万不要迷信论文里的默认超参数。不同数据集、不同分辨率下最优的损失权重组合差异很大。我建议训练前期先固定一组参数只观察生成图像的视觉效果确定某个方向趋势后再逐步调整对应损失项的权重。这样比起一上来就微调所有参数能更快找到一个可靠的组合。希望这个项目的思路和踩坑整理能帮你少走弯路。本文还有配套的精品资源点击获取