VisionTransformer图像去雾:Python源码与实战解析

发布时间:2026/9/28 6:36:03
VisionTransformer图像去雾:Python源码与实战解析 简介图像去雾是典型的病态反问题传统方法依赖暗通道先验估计透射率在天空、白墙等区域容易失效。VisionTransformer凭借全局自注意力机制能够有效建模长程依赖对不均匀雾和复杂场景表现出更强的稳定性。本文从大气散射模型出发讲解ViT去雾的核心原理并结合一份可运行的Python源码系统覆盖环境配置、数据准备、训练调参、常见排障与轻量化部署帮助开发者快速复现深度学习去雾基线适用于论文复现、雾天监控、车载视觉等工程场景。1. 基于VisionTransformer的图像去雾这份Python源码把路铺到了哪一步拿到标注“基于VisionTransformer的图像去雾算法研究与实现python源码项目说明数据集”的压缩包先别急着解压训练——先搞清楚这份源码解决的问题边界。图像去雾不是简单的“提高对比度”它的数学本质是一个病态反问题同一张清晰图可以对应无数张雾图你需要在缺少深度信息的前提下把透射率和大气光从像素里拆出来。VisionTransformer在这类任务里的价值在于它把输入切成patch后做全局自注意力能捕捉CNN局部感受野够不着的长程依赖对不均匀雾、物体边缘和天空区域的判断比纯卷积网络更稳。这份Python源码加项目说明加数据集的组合适合论文复现的在校生也适合被雾天监控、车载相机画面困扰、想快速验证深度学习去雾基线的开发者。下面按“原理—环境—训练—排障—落地”的顺序讲命令都能直接复制运行。2. 大气散射模型与ViT去雾原理为什么自注意力能hold住不均匀雾去雾任务不是凭空造轮子。VisionTransformer的去雾方案几乎都以大气散射模型为理论基础再决定网络结构和损失函数怎么设计。这一章先把“雾是怎么形成的”讲清楚再对比ViT和CNN在去雾上的差距最后给一组可以直接用的损失函数代码。2.1 图像去雾的数学起点一张雾图是怎么形成的图像去雾绕不开大气散射模型Atmospheric Scattering Model。它把雾图的形成拆成三部分场景反射光J(x)经过透射率t(x)衰减再加上大气光A被散射后的贡献写成公式就是I(x) J(x) * t(x) A * (1 - t(x))其中x是像素坐标t(x)exp(-beta*d(x))d(x)是场景深度。beta是散射系数雾越浓beta越大透射率越低图像信息衰减越严重。这个模型最大的问题是病态一个方程要解出J、t、A三个未知数数学上无法直接求逆。传统方法用暗通道先验估算t再反解J但遇到天空、白色墙体这类“暗通道失效”的区域就会翻车。深度学习去雾绕开了显式估计t的路径直接学习“雾图I到清晰图J”的映射。大部分研究实现的第一步是用大气散射模型在干净图像上随机生成不同浓度的雾构造训练对。你手里这份数据集里的合成雾图绝大多数就是从这个过程来的。这里给一个最小合成函数方便你理解数据集里那些雾图是怎么长出来的也方便后续扩充自己的训练数据import numpy as np import cv2 def synthesize_haze(clean_img, A, beta, depth_map): # 大气散射模型I J * t A * (1 - t) # t(x) exp(-beta * depth(x))depth 是各像素的深度值 t np.exp(-beta * depth_map) t t[..., None] # 把透射率变成与图像同形状的通道 hazy clean_img * t A * (1 - t) return np.clip(hazy, 0, 255).astype(np.uint8) A np.array([200, 200, 200], dtypenp.float64) # 大气光强度RGB 三通道一致表示灰雾 beta 0.02 # 散射系数越大雾越浓 depth_map np.random.rand(480, 640) * 10 # 随机深度图模拟不同远近的物体 clean cv2.imread(clean.png)[..., ::-1] # BGR 转 RGB hazy synthesize_haze(clean, A, beta, depth_map)代码里值得注意的参数是A和beta。A取[200,200,200]表示均匀的灰色大气光对应真实雾天最常见的环境光如果A的三个通道数值不一致比如[190,210,180]生成的就是带色偏的雾这在后续评估模型泛化性时非常有用。beta的取值直接影响雾浓度实际项目中很少用固定值而是按区间随机抽样避免模型只学会一种雾浓度导致过拟合。提示实际训练时请不要把A和beta写死。常见做法是A在[180,220]区间随机beta在[0.01, 0.05]区间随机这样合成出来的雾图分布更接近真实场景。2.2 VisionTransformer凭什么比CNN更擅长去雾VisionTransformer的核心操作是把图像切成固定大小的patch例如16x16展平成序列后加上位置编码再通过多头自注意力MSA在patch之间交换信息。相比CNN有三个特性让它特别适合去雾。第一是全局感受野来得快。CNN需要不断堆叠卷积层才能扩大感受野MSA在第一层就能看到输入图像的所有patch。去雾场景中某个像素的透射率受远处边界、景深突变和天空区域共同影响全局上下文能让模型区分“这是远处薄雾还是物体本身的白”。第二是内容自适应的权重分配。卷积核的权重在空间上是共享的所有位置用同一套卷积参数自注意力权重则根据输入内容动态计算浓雾区域和薄雾区域在特征聚合时可以走不同的路径这对不均匀雾尤其关键。第三是透射率本身的低频特性。透射率t(x)在空间中通常是平滑变化的而ViT对低分辨率patch序列的处理方式本身带有一定程度的分块语义抽象比CNN逐像素卷积更不容易把单点噪声放大。实现时的常见做法是保留ViT作为encoder后接UNet风格的decoder用跳跃连接把浅层边缘信息拉回来。也可以按Swin Transformer的思路加窗口注意力把全局注意力的平方复杂度降成线性这一点在第6章展开。需要强调的是ViT去雾不等于完全抛弃CNN工程上最稳的组合是CNN做浅层特征提取ViT做深层全局建模两头兼顾。2.3 损失函数怎么组合L1、感知损失与SSIM的权重取舍模型结构只是骨架损失函数决定优化方向。图像去雾的损失一般分三层像素级、结构级和语义级。像素级首选L1损失它对离群点不敏感边缘比L2更锐利结构级用SSIM损失保持局部对比度能抑制“去雾后整张图变灰”的问题语义级用感知损失让输出在预训练分类网络的特征空间里与目标更接近视觉质量更高但训练稳定性受backbone预训练权重影响。我一般把L1设为基准SSIM给小权重感知损失作为可选项。一个最小实现长这样import torch import torch.nn as nn import torch.nn.functional as F from torchmetrics.functional import structural_similarity_index_measure as ssim class DehazeLoss(nn.Module): def __init__(self, l1_w1.0, ssim_w0.2, perc_w0.0): super().__init__() self.l1_w l1_w self.ssim_w ssim_w self.perc_w perc_w def forward(self, pred, target): # pred 和 target 的取值必须在 [0,1] 范围 loss_l1 F.l1_loss(pred, target) loss_ssim 1.0 - ssim(pred, target, data_range1.0) total self.l1_w * loss_l1 self.ssim_w * loss_ssim if self.perc_w 0: # 接入预训练 VGG 的特征层做感知损失代码按需补 pass return total参数上l1_w1.0是底线ssim_w给0.2已经能明显改善色彩饱和度再大会导致训练早期振荡。perc_w从0.0起步试到0.1左右看验证集PSNR有没有提升没提升就退回0.0。要注意的是SSIM函数要求输入在[0,1]范围如果你的模型输出层没有做sigmoid训练时就必须在loss外面手动clamp否则SSIM会算出一堆奇怪的值这个坑我在排障章节会再提。损失项权重参考作用L11.0像素级准确度边缘保持SSIM0.1~0.3结构相似度抑制灰度化Perceptual0.05~0.1特征级视觉质量可选3. 从压缩包跑通第一张去雾图环境配置、数据集准备与代码入口这一章解决“怎么把源码跑起来”的问题。解压zip之后先按顺序做三件事建Python环境、理数据集目录、读项目说明里的训练入口。这三步做完基本就能让训练脚本转起来。3.1 Python环境从安装解释器到requirements依赖齐活如果你的机器还没装Python先去python官网下载Python 3.10或3.11Windows安装时勾选“Add Python to PATH”。然后打开终端创建环境推荐用conda隔离避免不同项目的依赖打架conda create -n dehaze python3.10 -y conda activate dehaze pip install torch torchvision --index-url https://download.pytorch.org/whl/cu118 pip install numpy opencv-python pillow tqdm torchmetrics第一行创建conda环境第二行激活第三行安装CUDA 11.8版本的PyTorch。如果你用的显卡比较新比如RTX 40系可以把cu118换成cu121对应CUDA 12.1运行时。第四行安装的是图像处理、命令行进度条和指标计算库torchmetrics用于计算SSIM指标也是上一章损失函数里的依赖。如果源码包里带了requirements.txt优先执行pip install -r requirements.txt但没有带的话就用上面这套组合基本覆盖训练和推理所需。安装完成后执行python -c import torch; print(torch.cuda.is_available())输出True说明GPU可用否则后面训练只能用CPU硬扛速度会慢几十倍。提示老显卡显存只有4GB时优先用cu118版本而不是最新版新版PyTorch对老架构支持并没有带来训练加速。3.2 把数据集理成训练脚本认识的目录结构数据组织是源码落地最容易被忽视的一环。数据集解压后常见有两种结构一种是hazy和clear两个并列目录一个存雾图、一个存同名清晰图另一种是已经按train/test切好的嵌套目录。不管哪种训练脚本最需要的都是“雾图和清晰图文件名一一对应”。我拿到新数据集的第一件事是跑一段脚本检查配对数量避免训练一半报下标越界import os hazy_dir data/hazy clear_dir data/clear pairs [] for name in os.listdir(hazy_dir): clear_path os.path.join(clear_dir, name) hazy_path os.path.join(hazy_dir, name) if os.path.exists(clear_path): pairs.append((hazy_path, clear_path)) print(valid pairs:, len(pairs)) # 假如文件名不一致比如 haze_001.png 对应 clear_001.png需要先对齐 # 常见做法是先排序再按索引配对 hazy_names sorted(os.listdir(hazy_dir)) clear_names sorted(os.listdir(clear_dir)) assert len(hazy_names) len(clear_names), 目录内图像数量不一致这里有个容易踩的细节资源管理器里排序和Python排序规则不一样haze_10.png会排在haze_2.png前面如果按索引配对就会错位。所以代码里必须用sorted()并且以数量断言兜底。如果数据集里有部分图像没有配对不要直接删掉先看看是不是文件名后缀大小写不同统一成小写再试。如果你准备用自己的照片扩充训练集目录结构就按上面的hazy/clear放文件名保持一一对应即可。3.3 项目说明与源码入口先读README再跑train.py规范的去雾源码包不会只有一个py文件通常有models目录、datasets目录、train.py、test.py、requirements.txt和README.md。项目说明README里最有用的是三块数据集来源和目录要求、超参数默认值、训练与测试命令。我见过很多人跳过README直接跑train.py结果数据集路径不对报FileNotFoundError回头再看README才发现根目录要求是data/RESIDE/train。一份可运行的训练入口命令大致长这样python train.py \ --data_root ./data \ --img_size 256 \ --patch_size 16 \ --epochs 100 \ --batch_size 8 \ --lr 1e-4 \ --save_dir ./checkpointsdata_root指向数据集根目录img_size控制输入分辨率patch_size是ViT切块大小这两个参数直接决定显存消耗和细节恢复能力。batch_size受显存限制8是一个比较稳的起点lr用1e-4配合AdamW是ViT类模型最常见的配置。save_dir指定checkpoint保存路径。如果压缩包里只有一个pytest或inference脚本而没有train.py说明这份源码偏向推理演示。你就先跑test.py把checkpoint加载起来出一张去雾图看效果再考虑要不要自己写训练循环。跑通推理比跑通训练更重要因为模型效果好坏能直接看出数据预处理是否匹配。4. 训练和评估一个去雾ViT超参数、指标与推理命令跑通第一遍之后就到了“调参”阶段。这一章先把训练超参讲清楚再给评估指标和推理脚本。调参不是玄学每一个参数都对应可观察的结果变化。4.1 训练超参数怎么设patch size、depth、lr与warmup去雾ViT里最值得调的四个维度是输入分辨率、patch大小、Transformer层数和学习率。参数推荐区间说明img_size256 / 512越大细节越丰富显存成倍上涨patch_size8 / 16越小越能保留细边缘计算量越大depth6~12深浅影响感受野层次过头会过拟合lr1e-4~3e-4AdamW下1e-4是安全起点batch_size8~32由显存决定小batch加梯度累积warmup5~10 epochs稳定ViT早期训练防止Loss爆炸patch_size的选择是个直接trade-off。patch16时序列长度短训练速度快但去雾结果在边缘处容易糊patch8时序列长度变4倍显存压力明显增大边缘细节更清晰。如果你手里的数据集大多是小目标或薄雾场景优先选patch8如果是大面积均匀雾patch16完全够用。训练策略上常见做法是先固定ViT backbone若干epoch只训练decoder的卷积层等Loss降到平台期再放开所有参数微调。这个思路和迁移学习的逻辑一致避免随机初始化的decoder把良好初始化的encoder带偏。学习率计划建议用warmup加cosine衰减warmup前5个epoch从1e-5逐渐升到1e-4后面余弦降到底。4.2 用PSNR和SSIM量化去雾效果PSNR峰值信噪比和SSIM结构相似性是去雾论文里最常出现的两个指标。PSNR只看像素误差数值越高越好SSIM看亮度、对比度和结构的综合相似度取值在0到1之间越高越好。合成雾数据集有清晰真图可以直接算真实雾图没有真图只能靠主观视觉判断。计算脚本可以自己写不用依赖重型框架import numpy as np from skimage.metrics import structural_similarity as ssim def calc_psnr(pred, target, max_val255.0): # pred 和 target 都是 uint8 图像取值范围 [0,255] mse np.mean((pred.astype(np.float64) - target.astype(np.float64)) ** 2) if mse 0: return float(inf) return 10 * np.log10(max_val ** 2 / mse) def calc_ssim(pred, target): # channel_axis 表示颜色通道所在维度对 HWC 图像取 2 return ssim(pred, target, channel_axis2, data_range255)注意PSNR算的是全图MSE如果图像里有大块天空区域PSNR会被天空的微小色差主导边缘细节的好坏反而被平均掉了。所以单看PSNR是不够的要配合SSIM和人工目检。一个好的去雾模型PSNR在RESIDE合成测试集上至少要到22以上SSIM到0.85以上才说得过去低于这个水平就优先怀疑训练数据和损失函数配置。4.3 推理单张图与批量图怎么跑训练结束后推理脚本并不复杂。核心是加载checkpoint、预处理、模型前向、反归一化四步import torch from PIL import Image import torchvision.transforms as T def dehaze_one_image(model, image_path, devicecuda, use_normTrue): transform_list [ T.Resize((256, 256)), T.ToTensor(), # 像素从 [0,255] 缩放到 [0,1] ] if use_norm: # 很多ViT训练时会做均值方差归一化推理必须保持一致 transform_list.append(T.Normalize([0.5, 0.5, 0.5], [0.5, 0.5, 0.5])) transform T.Compose(transform_list) x transform(Image.open(image_path).convert(RGB)).unsqueeze(0).to(device) model.eval() with torch.no_grad(): y model(x) y y.squeeze(0).cpu() if use_norm: y y * 0.5 0.5 # 反归一化到 [0,1] y torch.clamp(y, 0.0, 1.0) # 防止越界 return y.permute(1, 2, 0).numpy() # CHW 转 HWC这里最容易翻车的点就是归一化。如果项目说明里写了训练时用的mean和std推理时必须用同一套如果没写最安全的是把use_norm设为False只用ToTensor。用错归一化输出的图像会整体偏色或对比度异常看起来像模型坏了实际只是数据处理不一致。批量推理时把多张图凑成一个batch传入模型比一张张循环快很多典型提速两到三倍代码后面第6章再展开。5. ViT去雾项目的常见坑现象、原因与处理方案跑过几轮ViT去雾之后你会发现大部分问题不在模型结构而在训练和预处理细节。这一章整理5个高频踩坑场景每一条都是“现象—原因—解决”的记录能帮新手少走弯路。5.1 训练Loss下降但去雾效果反而变差现象训练集PSNR一路上涨验证集却越来越差真实雾图甚至出现伪影。原因合成雾数据集分布太单一模型过拟合到“训练集特有的雾浓度和颜色偏移”。很多开源数据集只用固定A和beta合成模型学会了降低对比度而不是真正去雾。解决第一步增加数据增强在合成阶段对A和beta做随机采样第二步对输入图像加随机亮度扰动和色彩抖动打破训练集的静态分布第三步用早停监控验证集PSNR连续10个epoch不涨就保存最优checkpoint并停止训练。如果你不想改训练代码最简单的办法是先做transforms.ColorJitter比如把亮度抖动范围设成0.8~1.2。5.2 去雾结果整体偏灰或偏黄现象模型输出确实变亮了但色彩饱和度低整个画面像罩了一层灰还有些情况会整体偏黄或偏蓝。原因损失函数中L1占比过高模型倾向于输出均值附近的中性色因为这样L1误差最小。偏黄偏蓝则多半是大气光A估计错误或者训练时归一化参数与推理不一致。解决先检查推理脚本里的反归一化把y * 0.5 0.5漏掉或写错是最常见原因。排除后把SSIM权重从0.1提到0.3SSIM损失能有效拉回局部对比和色彩饱和度。如果还偏黄去数据集里抽一张清晰图统计三通道均值把大气光A的初值设成接近该均值再微调网络。5.3 天空等平坦区域出现棋盘伪影或色块现象去雾结果在天空、白墙这类平滑区域出现规则的网格条纹或块状色差边缘处尤其明显。原因patch_size设置过大平坦区域被分割成了语义不连续的块同时decoder中如果有反卷积操作小数倍上采样会产生棋盘格伪影。解决优先把patch_size从16降到8序列变长但平坦区域的块状感会明显减弱。然后把反卷积替换成pixel shuffle或双线性上采样加卷积。如果还压不住就在损失里加一个轻量的边缘平滑正则让相邻像素在梯度上保持连续。5.4 GPU显存不足OOM怎么办现象训练刚跑第一个batch就报torch.cuda.OutOfMemoryError显存被瞬间打满。原因VisionTransformer的全局自注意力需要保存所有patch两两之间的注意力矩阵显存占用随序列长度平方增长。img_size512时的patch数量是256时的4倍显存占用可能直接超16GB。解决先降batch_size到2甚至1确认能跑通然后考虑缩小img_size到224或256最后用混合精度训练PyTorch自带torch.cuda.amp训练速度能提升30%以上显存占用降一小半。如果上面三步还不够就需要换轻量注意力结构第6章会讲具体方案。5.5 加载别人checkpoint报key不匹配现象执行model.load_state_dict(torch.load(...))时报Missing key(s)或Unexpected key(s)模型完全加载不进去。原因对方训练时的模型配置和你本地不一致比如depth不同、patch_size不同、输出头尺寸不同或者checkpoint里包含EMA影子权重和optimizer状态字典。解决不要硬加载先只挑选encoder部分权重用下方代码按key名过滤import torch state torch.load(checkpoints/pretrained.pth, map_locationcpu) # 兼容直接存state_dict和存整个训练对象两种情况 if state_dict in state: state state[state_dict] model_dict model.state_dict() matched {} for k, v in state.items(): if k in model_dict and v.shape model_dict[k].shape: matched[k] v model.load_state_dict(matched, strictFalse) missing [k for k in model_dict if k not in matched] print(missing keys:, missing[:20])代码里先判断checkpoint是裸state_dict还是包含其他字段的字典再用形状一致性过滤最后打印出还缺哪些key。大多时候缺的是decoder输出头这是正常的encoder权重能加载进来已经节省了大部分训练时间。如果missing keys里包含encoder前几层的embedding说明patch_size不一致这时只能统一配置重练。6. 从“能跑”到“能用”轻量化压缩与真实场景验证模型在合成雾数据集上表现好不等于在真实雾天场景就能直接落地。最后一章讲三个实用方向轻量化结构、批量推理、真实场景验证流程。如果你是想做工程交付而不是单纯交作业这一章最值得看。6.1 把标准ViT换成线性注意力变体标准全局自注意力的复杂度是O(N^2)N是patch数量img_size一上来就喘不过气。换线性注意力是开源方案里最省事的压缩路径核心是把Softmax自注意力改成核函数近似让复杂度回落到O(N)。常见的做法是把MSA换成Swin Transformer的窗口注意力或者换用Linear Attention实现其余结构完全不变。# 示意线性注意力替换多头注意力 def linear_attention(q, k, v): # q/k/v shape: (B, H, N, D) q torch.softmax(q, dim-1) k torch.softmax(k, dim-1) kv torch.einsum(bhnd,bhne-bhde, k, v) # 先合并K和V降低计算量 out torch.einsum(bhnd,bhde-bhne, q, kv) return out这个替换对其他代码几乎无侵入训练一次的成本比标准MSA低很多。缺点是细节恢复能力略有下降在追求PSNR的合成数据集上可能掉0.3~0.5个点但换来的是可接受的显存占用和推理速度工程上完全可以接受。6.2 批量推理与内存复用真实世界的去雾往往不是处理一张图而是处理视频帧或相机连拍。批量推理的关键是一次把多张图送入模型而不是在循环里逐张调用模型。用PyTorch的DataLoader包装文件夹配合batch_size8或16能稳定提速两到三倍。6.3 真实雾图验证的固定流程我的习惯是训练完先在合成测试集上算一次PSNR和SSIM然后立刻用三张不同浓度的手拍雾图做目检。目检时看两个点天空区域有没有色块近景边缘有没有伪影。如果真实雾图效果和合成指标差距很大先怀疑训练数据的问题而不是模型结构。我自己第一次用ViT做去雾时就翻过车训练时用了归一化推理时忘了反归一化直接把tensor转成uint8保存输出整片灰白我还以为是模型结构写错了排查了半个小时才发现是预处理不一致。从那以后我养成了一个习惯所有测试图先存一版原始模型输出再存一版可视化结果两张对比数据问题立刻暴露。希望帮到你。本文还有配套的精品资源点击获取