简介:图像去雾是计算机视觉中的经典难题,其核心在于从退化图像中恢复清晰场景。传统暗通道先验依赖物理模型,在复杂场景下易产生人工痕迹;而基于深度学习的生成对抗网络(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_size=3, padding=1) self.bn1 = nn.BatchNorm2d(in_channels) self.relu = nn.ReLU(inplace=True) self.conv2 = nn.Conv2d(in_channels, in_channels, kernel_size=3, padding=1) 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_channels=3, base_channels=64): super(PatchDiscriminator, self).__init__() self.model = nn.Sequential( nn.Conv2d(in_channels, base_channels, kernel_size=4, stride=2, padding=1), nn.LeakyReLU(0.2, inplace=True), nn.Conv2d(base_channels, base_channels * 2, kernel_size=4, stride=2, padding=1), nn.BatchNorm2d(base_channels * 2), nn.LeakyReLU(0.2, inplace=True), nn.Conv2d(base_channels * 2, base_channels * 4, kernel_size=4, stride=2, padding=1), nn.BatchNorm2d(base_channels * 4), nn.LeakyReLU(0.2, inplace=True), nn.Conv2d(base_channels * 4, 1, kernel_size=4, stride=1, padding=1) ) 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(pretrained=True).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, device='cuda'): model = build_generator() model.load_state_dict(torch.load(model_path, map_location=device)) 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千万不要迷信论文里的默认超参数。不同数据集、不同分辨率下,最优的损失权重组合差异很大。我建议训练前期先固定一组参数,只观察生成图像的视觉效果,确定某个方向趋势后,再逐步调整对应损失项的权重。这样比起一上来就微调所有参数,能更快找到一个可靠的组合。希望这个项目的思路和踩坑整理能帮你少走弯路。
本文还有配套的精品资源,点击获取