公司动态
DehazeNet图像去雾实战:PyTorch实现大气散射模型与传输图估计
简介本资源是面向深度学习初学者与图像复原研究者的PyTorch版DehazeNet图像去雾实现方案聚焦于单幅图像雾霾去除这一经典低层视觉任务适用于遥感、自动驾驶、监控视频增强等实际场景。压缩包共21个文件包含9个核心Python模块如网络定义net.py、数据加载train_data.py、预处理pre.py、推理demo.m、2个已训练.pth模型best_indoor.pth与best_outdoor.pth、3个MATLAB辅助函数guidedfilter.m等及配套说明文档与README整体仅114KB轻量易部署。已有40人学习下载资源提供开箱即用的完整训练-验证-推理流程支持直接加载预训练权重进行效果演示亦可基于现有架构快速开展消融实验或迁移适配。代码结构清晰、注释规范特别适合具备PyTorch基础的研究者理解端到端去雾网络的设计逻辑与工程落地细节。 雾天拍出来的照片除了发白发灰、对比度低更麻烦的是后续的人脸识别、目标检测、语义分割统统会掉点。以前做去雾大家第一反应是暗通道先验那套物理方法效果确实不错但碰上天空区域就很容易翻车因为暗通道先验在天空这种大面积明亮区域本身就不成立。后来CNN方法慢慢成了主流DehazeNet就是其中比较早也相当经典的一个它用卷积网络直接回归大气散射模型里的传输图没有暗通道那种强先验假设对天空场景友好得多而且网络很小、推理很快一块普通GPU就能流畅跑视频帧级别的去雾。这次我把自己从零搭建的完整流程整理出来从大气散射模型原理、网络结构拆解到PyTorch环境配置、合成训练数据生成再到预训练模型加载与推理部署最后把几个容易踩的坑也一并列出来。整个过程基于PyTorch实现提供可直接运行的代码块你在自己的机器上逐步执行就能复现也可以直接把预训练权重拿来处理自己的雾天图片。1. 雾天图像为什么难处理大气散射模型与DehazeNet的解决思路1.1 从物理退化到数学表达在有雾天气下相机传感器接收到的光分成两部分一部分是场景物体反射光经过大气衰减后到达相机的直接辐射另一部分是大气光被空气中悬浮颗粒散射进入相机的环境光。这两部分叠加才形成了我们看到的发白、模糊的雾天图像。这个物理过程被McCartney等人整理成了经典的大气散射模型I(x) J(x) · t(x) A · (1 - t(x))其中x是像素坐标I(x)是观测到的有雾图像J(x)是我们要恢复的无雾清晰图像A是全局大气光t(x)是传输图表示物体反射光没有经过衰减直接到达相机的比例。传输图的物理范围在0到1之间t1意味着完全没有雾t0意味着信息完全被雾气遮挡。雾天拍摄时远处的物体传输图趋近0所以看起来白茫茫一片近处物体传输图接近1纹理相对清晰。去雾任务的本质就是在只有I(x)已知的情况下同时求解J(x)、A和t(x)。这是典型的不适定问题——三个未知量一个方程信息严重不足必须借助额外的约束或先验。1.2 DehazeNet为什么选择直接估计传输图传统方法里最有代表性的就是暗通道先验DCP统计发现无雾图像的局部区域在至少一个颜色通道上像素值接近0而对有雾图像做同样统计时这个最小值不再是0而是被大气光抬高了由此可以反推出传输图。DCP在多数户外场景效果很好但一旦遇到天空、白色墙壁、雾本身厚重的区域暗通道假设失效估计出的传输图就会出现色偏和光晕。DehazeNet的思路绕开了暗通道先验的强假设用CNN从大量合成有雾图像中直接学习从有雾图像到传输图的非线性映射。网络输入是一张有雾图像输出是对应位置的传输图t(x)然后在已知传输图的条件下用大气散射模型反解出清晰图像。这样做有两个关键收益。第一CNN通过数据驱动的方式自动学到了比固定先验更鲁棒的映射关系对天空、高亮区域不再那么敏感。第二DehazeNet是一个全卷积网络不依赖输入图像尺寸推理时对任意分辨率图像都能直接处理不用像传统方法那样做复杂的块匹配和细化后处理。1.3 BReLU的物理含义与设计动机DehazeNet论文里最有辨识度的设计是BReLU激活函数全称是Bilateral Rectified Linear Unit双边修正线性单元。它做的事情就是一句话把输出值限制在[0,1]区间。为什么做这个限制因为传输图t(x)的物理含义是光线透射率取值范围天然就在0到1之间。如果网络输出层随便输出一个负值或者大于1的值反解出来的J(x)就完全不符合物理规律——负数亮度在图像里没法解释。用BReLU替代普通ReLU相当于在模型的最后一层加入了对传输图物理范围的先验约束不需要额外正则项网络也很难跑到不合理解空间里去。这个思路看着简单但在当时挺有启发性。后来很多低层视觉任务里的输出层比如深度估计、光流估计都会根据输出量的物理意义设计裁剪或激活函数本质上都是同一个思路。2. 环境搭建的版本陷阱与PyTorch安装实测2.1 版本选择Python、CUDA与PyTorch的搭配DehazeNet本身的依赖非常少只要PyTorch和torchvision就能跑不涉及什么冷门库。但PyTorch的安装版本搭配如果不留意很容易在import阶段就报CUDA相关的错。我这边的建议组合是这样组件推荐版本说明Python3.10PyTorch官方wheel覆盖最全的版本区间CUDA Toolkit11.8或12.1取决于显卡驱动新卡优先12.xPyTorch2.4 ~ 2.6建议装GPU版CPU版也能跑但慢很多torchvision与PyTorch匹配对齐版本号安装不要混装显卡驱动550.x以上较新向下兼容各CUDA版本需要注意一点PyTorch的CUDA运行时是随wheel打包的不需要单独安装完整CUDA Toolkit。只要显卡驱动版本足够新在conda环境里用pip指定index-url安装GPU版PyTorch就能正常调用GPU。很多人在这步被各种教程误导去NVIDIA官网装了一大堆驱动和SDK最后发现全是多余的。2.2 conda环境创建与GPU验证用conda创建独立环境是必须的别偷懒直接装在base环境里。不同项目的PyTorch版本要求不一样跑DehazeNet的时候你可能还在跑其他项目环境隔离能省掉大量为什么我这个环境突然import报错的烦恼。conda create -n dehaze python3.10 -y conda activate dehaze pip install torch torchvision --index-url https://download.pytorch.org/whl/cu121安装完成后验证GPU是否可用python -c import torch; print(torch.__version__); print(torch.cuda.is_available()); print(torch.cuda.get_device_name(0) if torch.cuda.is_available() else CPU only)正常输出应该类似2.6.0cu121 True NVIDIA GeForce RTX 4060如果cuda.is_available()返回False优先检查显卡驱动版本。在Linux下用nvidia-smi看Driver VersionWindows下在系统信息或NVIDIA控制面板里查。驱动版本过旧时即使PyTorch装的是CUDA 12.1的版本也会加载不了CUDA runtime。2.3 依赖安装清单DehazeNet项目本身只用到了下面这些库装起来很快torch2.0 torchvision0.15 numpy Pillow matplotlib opencv-python建议把opencv-python也装上虽然用Pillow也能做图像读取但后面做推理脚本或者数据预处理时OpenCV的颜色空间转换和resize接口更顺手。matplotlib则用于训练曲线可视化。如果你用的是Jetson设备比如项目中常见到的JetPack 6.2.2那镜像源里的PyTorch版本和PC端不太一样最好直接刷NVIDIA官方为JetPack预编译的PyTorch wheel包别用pip默认源装的CPU版本否则会浪费设备上的GPU算力。3. 手写DehazeNetMaxout、BReLU与全卷积结构3.1 网络主干结构逐层拆解DehazeNet的完整结构不算复杂原文的框架可以分成四层特征提取层一个3×3卷积将RGB三通道映射到16个特征图然后接Maxout激活输出4个特征图。多尺度映射层分别用3×3、5×5、7×7三种尺寸的卷积核对特征图并行卷积每种尺寸输出4个特征图拼接后得到12个特征图。这一步的设计目的是捕捉不同感受野下的雾气特征。局部极值层对12个特征图再做一次Maxout输出3个特征图。非线性回归层先通过5×5卷积和3×3卷积将通道数压缩到1最后接BReLU输出单通道的传输图。在写代码时我习惯在注释里标注每层的输入输出尺寸方便后续调整和排查问题。import torch import torch.nn as nn import torch.nn.functional as F class BReLU(nn.Module): 双边ReLU将输出限制在[0, 1]区间 def __init__(self): super(BReLU, self).__init__() def forward(self, x): return torch.clamp(x, min0.0, max1.0) class Maxout(nn.Module): Maxout激活按通道维度分组取最大值 def __init__(self, groups): super(Maxout, self).__init__() self.groups groups def forward(self, x): # x: [B, C, H, W] B, C, H, W x.shape x x.view(B, self.groups, C // self.groups, H, W) x, _ x.max(dim2) return x class DehazeNet(nn.Module): def __init__(self): super(DehazeNet, self).__init__() # 特征提取层 self.conv1 nn.Conv2d(3, 16, kernel_size3, padding1, biasFalse) self.maxout1 Maxout(4) # 多尺度映射层 self.conv2_1 nn.Conv2d(4, 4, kernel_size3, padding1, biasFalse) self.conv2_2 nn.Conv2d(4, 4, kernel_size5, padding2, biasFalse) self.conv2_3 nn.Conv2d(4, 4, kernel_size7, padding3, biasFalse) # 局部极值层 self.maxout2 Maxout(4) # 非线性回归层 self.conv3 nn.Conv2d(3, 1, kernel_size5, padding2, biasFalse) self.conv4 nn.Conv2d(1, 1, kernel_size3, padding1, biasFalse) self.brelu BReLU() def forward(self, x): # x: [B, 3, H, W] 有雾图像 x self.conv1(x) x self.maxout1(x) # [B, 4, H, W] x1 self.conv2_1(x) x2 self.conv2_2(x) x3 self.conv2_3(x) x torch.cat([x1, x2, x3], dim1) # [B, 12, H, W] x self.maxout2(x) # [B, 3, H, W] x F.relu(self.conv3(x)) x self.conv4(x) x self.brelu(x) # [B, 1, H, W] return x3.2 Maxout激活的实现细节Maxout不是DehazeNet独有的但它把Maxout用在多通道特征聚合上这里有一个容易理解错的地方。我第一次实现的时候以为Maxout只是对每个通道图做元素级取最大后来细看论文才发现它是沿通道维度把特征图分组然后在每个组内取逐位置的最大值。拿第一层举例卷积层输出16个特征图Maxout(4)的含义是每4个通道合成1个通道——分成4组每组4个通道在4张特征图的相同位置取最大值得到4个输出通道。这个操作放在今天看就像是一种特定方式的channel pooling它缩小了通道数同时保留了每个位置上最显著的特征响应。实现上就一句话x x.view(B, self.groups, C // self.groups, H, W) x, _ x.max(dim2)维度重排后再取最大值代码很简短但逻辑要理清楚分组数groups要能整除通道数C否则view会直接报错。3.3 从模型输出到去雾图像的闭环模型输出的是传输图t(x)这还不是最终的清晰图像。要得到去雾结果还需要估计全局大气光A然后用大气散射模型反解J(x) (I(x) - A) / max(t(x), t0) A其中t0是一个极小值防止分母出现0导致除零。实际取值一般在0.1到0.2之间。t(x)越小说明雾气越浓、信息损失越多恢复时稍微压低增益避免出现噪声放大和色块。全局大气光A的估计有几种方式。最简单的是直接取输入图像中亮度最高的像素值更鲁棒的做法是先求暗通道在暗通道中取最亮的前0.1%像素然后回原图找这些位置的平均像素值作为A。后者避免了纯白物体误判为大气光的风险。def estimate_atmospheric_light(hazy_img, top_k_ratio0.001): 估计全局大气光A Args: hazy_img: [H, W, 3] float32, 取值范围[0,1] top_k_ratio: 暗通道中最亮像素比例 Returns: A: [3] 三个通道的大气光 # 计算暗通道每个像素取三个通道的最小值 dark torch.min(hazy_img, dim-1).values # [H, W] flat dark.flatten() k max(int(flat.shape[0] * top_k_ratio), 1) _, indices torch.topk(flat, k) # 在原图中取这些位置像素的平均值作为A img_flat hazy_img.reshape(-1, 3) A img_flat[indices].mean(dim0) return A取到A之后按公式恢复即可。这里我给出一套完整的推理代码后面会复用torch.no_grad() def dehaze(model, hazy_img, t00.1): 有雾图像 - 无雾图像 Args: model: DehazeNet实例 hazy_img: torch.Tensor [H, W, 3] RGB顺序, 取值范围[0,1] t0: 传输图下限 Returns: clear_img: torch.Tensor [H, W, 3] RGB顺序 t_map: torch.Tensor [H, W] 传输图 model.eval() x hazy_img.permute(2, 0, 1).unsqueeze(0) # [1, 3, H, W] t model(x).squeeze(0).squeeze(0) # [H, W] A estimate_atmospheric_light(hazy_img) t_clamped t.clamp(mint0) # J (I - A) / t A clear (hazy_img - A.view(1, 1, 3)) / t_clamped.unsqueeze(-1) A.view(1, 1, 3) clear clear.clamp(0.0, 1.0) return clear, t注意A.view(1, 1, 3)这个维度处理A是三维向量要广播到[H, W, 3]必须保证维度顺序对齐。4. 合成雾数据与训练流程设计4.1 基于NYU Depth v2合成有雾图像DehazeNet训练用的是合成雾图因为真实世界中同时拿到同一场景的有雾/无雾图像对极其困难。最经典的做法是用NYU Depth v2数据集它包含RGB图和对应的深度图利用深度信息d(x)和大气散射模型合成有雾图像。合成公式也很直接t(x) exp(-β · d(x)) I(x) J(x) · t(x) A · (1 - t(x))其中β是大气散射系数控制雾的浓度A是全局大气光。训练时对每张图像随机采样β和A等于给模型提供了不同雾浓度和光照条件下的样本增强泛化能力。def synthesize_hazy(clear_img, depth_img, betaNone, ANone): 根据深度图合成有雾图像 Args: clear_img: [H, W, 3] float32, [0,1] depth_img: [H, W] float32, 深度值 beta: 散射系数None则随机采样 A: 大气光None则随机采样 Returns: hazy_img: [H, W, 3] t_map: [H, W] if beta is None: beta np.random.uniform(0.4, 1.6) if A is None: A np.random.uniform(0.5, 1.0, size(1, 1, 3)) t np.exp(-beta * depth_img) hazy clear_img * t[..., None] A * (1 - t[..., None]) return hazy.astype(np.float32), t.astype(np.float32)每次迭代随机采样β和A相当于无限扩充训练集。我在实际训练中还加了随机裁剪、水平翻转、颜色抖动等数据增强网络泛化能力好了不少。4.2 训练损失与优化器配置DehazeNet原文用的是MSE损失也就是让预测传输图和真实传输图逐像素接近。MSE的优点是梯度稳定、实现简单缺点是容易导致恢复图像偏平滑纹理细节有丢失。我在复现时试过MSE、MAE和两者的组合最终效果是MSE和MAE按0.5:0.5加权最好边缘保留得更干净但收敛速度比纯MSE稍慢。loss_fn lambda pred, target: 0.5 * F.mse_loss(pred, target) 0.5 * F.l1_loss(pred, target)优化器选择Adam初始学习率1e-3余弦退火衰减到1e-5。批大小根据显存来我用的RTX 4060 8G显存batch size可以开到32输入patch是64×64。如果你的显存只有4G左右建议降为16或把patch缩到48×48。4.3 训练过程中的监控与调参建议训练时不要只盯loss曲线更重要的是定期在验证集上肉眼查看去雾效果。因为MSE这类像素损失优化的方向和人眼感知不总是一致loss很低但恢复图可能发灰、过曝或出现伪影。我的做法是每训练5个epoch保存一组验证样本的去雾结果图丢到一个固定目录里随时翻看。这样能及时发现以下问题输出传输图整体偏暗恢复图过曝检查损失权重是否合理或者BReLU的0-1约束是否被意外移除。天空区域出现色斑A的估计可能存在偏差训练数据里要保证A的采样范围覆盖0.7到1.0的高亮场景。边缘区域有光晕尝试增加MAE损失比重或多尺度卷积层的卷积核尺寸再调大。训练轮数方面DehazeNet参数量不大在单张RTX 4060上用10万张合成的64×64 patch训练大约40个epoch就能收敛耗时2小时左右。不需要像大模型那样跑几天几夜。5. 预训练模型加载与去雾推理全流程5.1 加载预训练模型的方式与注意事项项目里附带了一个在合成雾图上训练好的预训练权重文件dehazenet_epoch40.pth。加载方法分为两种只加载state_dict或者从checkpoint中恢复完整内容。推荐用state_dict方式兼容性最好model DehazeNet() checkpoint torch.load(dehazenet_epoch40.pth, map_locationcpu, weights_onlyTrue) model.load_state_dict(checkpoint[state_dict] if state_dict in checkpoint else checkpoint) model.eval()这里有个细节要注意如果在保存时把整个模型对象torch.save(model, ...)存进去了那么加载时torch.load默认的weights_onlyTrue会报错因为模型对象里包含的Python对象类型不在安全权重列表中。这时候需要显式设置weights_onlyFalse或者在保存时就只保存model.state_dict()从根源上避开这个问题。我建议保存时写成checkpoint { state_dict: model.state_dict(), opt_state: optimizer.state_dict(), epoch: epoch, best_psnr: best_psnr, } torch.save(checkpoint, dehazenet_epoch40.pth)这样保存的是字典形式PyTorch 2.6及以上版本加载时用默认的weights_onlyTrue也能正确载入不会被安全校验拦截。5.2 完整推理脚本与效果评估给出一份完整的推理脚本import argparse import torch import numpy as np from PIL import Image import torchvision.transforms as T def main(): parser argparse.ArgumentParser(descriptionDehazeNet 推理) parser.add_argument(--input, typestr, requiredTrue, help输入有雾图像路径) parser.add_argument(--output, typestr, defaultoutput.png, help输出去雾图像路径) parser.add_argument(--weights, typestr, defaultdehazenet_epoch40.pth) parser.add_argument(--gpu, typeint, default0) args parser.parse_args() device torch.device(fcuda:{args.gpu} if torch.cuda.is_available() else cpu) model DehazeNet().to(device) checkpoint torch.load(args.weights, map_locationcpu, weights_onlyTrue) model.load_state_dict(checkpoint[state_dict]) model.eval() # 读取图像转为float32 [0,1]RGB顺序 img Image.open(args.input).convert(RGB) img_tensor T.ToTensor()(img).permute(1, 2, 0) # [H, W, 3] # 为避免显存压力长边超过2000时做resize max_side max(img_tensor.shape[:2]) if max_side 2000: scale 2000.0 / max_side new_size (int(img_tensor.shape[1] * scale), int(img_tensor.shape[0] * scale)) img_tensor F.interpolate( img_tensor.permute(2, 0, 1).unsqueeze(0), sizenew_size, modebilinear, align_cornersFalse, ).squeeze(0).permute(1, 2, 0) # 推理 clear_img, t_map dehaze(model, img_tensor.to(device), t00.1) # 保存结果 clear_np (clear_img.cpu().numpy() * 255).clip(0, 255).astype(np.uint8) t_np (t_map.cpu().numpy() * 255).clip(0, 255).astype(np.uint8) Image.fromarray(clear_np, RGB).save(args.output) Image.fromarray(t_np, L).save(transmission_map.png) print(f去雾图像已保存到 {args.output}) if __name__ __main__: main()在测试集上预训练模型的PSNR大约在22-24dBSSIM在0.88到0.92之间具体取决于测试集的雾浓度分布和场景内容。单独看数值其实只是参考真正的判断标准还是肉眼看有没有偏色、雾是否除干净、边缘是否保留完整。5.3 推理性能优化DehazeNet参数量约5万左右计算量集中在多尺度卷积层。测试下来CPUi5-12400上跑一张1024×768图像约0.3秒GPURTX 4060上跑同样尺寸图像约10毫秒批量推理时把多张图拼成batch一次前向吞吐量提升更明显。如果要做实时视频去雾可以进一步用TensorRT或ONNX Runtime量化。实测把模型导出为ONNX后用FP16模式跑比PyTorch原版快1.5到2倍且精度损失可忽略。但导出ONNX时要注意Maxout层里的view操作确保输入形状是动态的否则换分辨率会报错。torch.onnx.export( model, torch.randn(1, 3, 256, 256).to(device), dehazenet.onnx, input_names[input], output_names[transmission], dynamic_axes{input: {0: batch, 2: height, 3: width}, transmission: {0: batch, 2: height, 3: width}}, opset_version17, )6. 实测踩坑记录与后续优化方向6.1 PyTorch 2.6加载旧模型的新坑这个话题在社区里最近讨论得特别多我也真实踩到了。PyTorch 2.6发布后torch.load()函数的weights_only参数默认值改成了True目的是提升安全性防止加载恶意pickle文件执行任意代码。但这也意味着以前那些直接保存整个模型对象的老checkpoint在新版PyTorch下会直接报错。典型报错信息是UnpicklingError: Weights only load failed: Unsupported global: builtins.getattr排查思路很简单。先分清楚你的.pth文件是两种类型中的哪一种保存了完整模型对象torch.save(model, path)必须在加载时加weights_onlyFalse或者重新训练后用state_dict方式保存。保存的是state_dict或包含键值的字典默认的weights_onlyTrue就能加载。我自己的做法是统一改成了字典保存顺带着把训练配置、优化器状态一起存进去这样复现训练和做断点续训都方便。6.2 图像预处理不一致导致的伪影推理时最容易忽略的问题是预处理不一致。训练时图像是浮点数[0,1]、RGB通道顺序、BGR和RGB的问题在训练数据里踩过一次。OpenCV读图默认是BGRPIL读图是RGB。如果训练数据用的是PIL读图但推理脚本用OpenCV读图且忘了转换色彩空间去雾结果会整体偏蓝偏红效果惨不忍睹。另一个常见问题是归一化方式不一致。如果训练时把像素除以255映射到[0,1]推理时却忘了做这一步模型输入分布完全不对输出的传输图会整体偏移恢复图发灰。这个问题我在代码里刻意规避了在推理脚本中统一使用T.ToTensor()完成归一化。还有一类问题更容易被忽略图像resize时插值方式不一致。如果训练时用双线性插值推理时用最近邻边缘处会出现锯齿状伪影肉眼看得特别明显。建议在训练和推理的预处理中统一使用modebilinear。6.3 后续可以这样扩展DehazeNet作为2016年的方法和现在基于Transformer、扩散模型的去雾算法相比在极端浓雾场景下的恢复能力确实有差距。但它的结构小、速度快、训练成本低在工程落地中仍然很有价值。我建议按下面几个方向扩展把DehazeNet的输出作为粗传输图再接一个轻量的细化模块比如梯度域引导滤波或空间注意力能明显改善边缘锯齿。在损失函数中加入感知损失用ImageNet预训练的VGG提取高层特征算MSE恢复图像的纹理细节会更自然。做数据集扩充时加入真实雾图数据哪怕只有几十张和合成数据混合训练泛化能力也能提升不少。部署到Jetson这类边缘设备时可以用ONNX Runtime或TensorRT做推理加速步骤在前文已给出。模型冻结部分指定层做迁移学习也是可行方向。比如在无人机航拍图、水下图像等特定场景下用少量真实数据微调最后一层卷积训练时间很短效果提升往往立竿见影。# 冻结前两层只训练后面层 for name, param in model.named_parameters(): if conv1 in name or conv2 in name: param.requires_grad False不过要注意冻结层之后学习率可以适当调大一些但微调轮数不宜过多不然容易出现灾难性遗忘把原来在合成数据上学到的通用去雾能力丢掉。我在实际使用中的一个心得是DehazeNet用好了更像是一个预处理模块。不要只把它当作最终去雾方案可以接在目标检测、语义分割模型前面做数据清洗雾天图像经过它再进检测网络识别精度提升往往比直接换更强的检测模型还明显而且额外耗时只有十几毫秒。这个思路在很多实际项目里比死磕去雾指标更实用。本文还有配套的精品资源点击获取