Vision Transformer图像去雾:原理、复现与避坑指南

发布时间:2026/10/2 14:03:57
Vision Transformer图像去雾:原理、复现与避坑指南 简介面向图像去雾研究者和深度学习开发者这份压缩包提供了基于Vision Transformer的图像去雾算法完整Python实现包含项目介绍与使用说明可帮助读者快速复现算法并开展训练与测试。资源共340个文件以Python源码为主辅以YAML配置文件、PNG/GIF示例图、CSV结果记录及Markdown文档等整体约156.38MB目录结构清晰便于按模块查阅。目前已有1444人学习下载适合具备一定PyTorch基础、希望深入理解ViT在底层视觉任务中应用的中高级学习者。包内不仅包含模型实现与训练脚本还提供预训练权重和详细参数说明如训练样本patch大小设置并配有loss landscape等分析数据便于研究不同配置下的训练效果与模型特性。1. 从“像素级增强”到“全局场景建模”ViT去雾到底在解决什么很多第一次接触“基于Vision Transformer的图像去雾”的人第一反应是“这不就是给雾图做个增强吗”。我第一次跑完这类源码项目才意识到它真正的难点不是把雾“擦掉”而是让模型明白一张雾图里远山、中景建筑、近处车辆接受的是同一束大气光透射率却在空间上连续变化。Vision TransformerViT把图像切成patch后走全局attention天然适合这个建模方式。这类项目的价值在浓雾薄雾混合场景下能把色彩偏差拉回来短板是边缘细节容易糊。适合做毕设课设复现、算法对比也适合安防和车载图像的预处理。下面按我自己的复现路径讲。2. 去雾的数学模型与Vision Transformer选型理由先搞懂在优化什么2.1 大气散射模型与病态估计雾天成像的“老规矩”传统去雾的起点是大气散射模型也叫雾天成像模型I(x) J(x)·t(x) A·(1 - t(x))其中I是相机拍到的雾图J是清晰图A是全局大气光t(x)是透射率。t接近1表示该像素几乎没受雾影响t接近0表示浓雾区远处景物信息几乎被大气光盖住了。图像去雾要做的事就是从I反推J也就是把所有像素的t和A“拆”出来再按公式重算。这里有个数学意义上的麻烦一个模糊像素对应三个未知量J、t、A方程只有一个属于严重欠定。传统方法靠人工先验来约束搜索空间最出名的是暗通道先验(dark channel prior)作者观察到无雾清晰图像的局部区域里至少有一个颜色通道的亮度很低。基于这个假设能先估t再估A最后恢复J。它管用但在天空、白色墙体、雾内灯源这些“暗通道不暗”的地方会直接翻车输出灰块或光晕。深度学习时代的去雾本质上是把“人工先验”换成“数据分布先验”给网络大量成对的雾图/清晰图让它学会从像素直接映射到透射率、或直接映射到清晰图。这样就不需要手写暗通道假设模型可以在统计意义上覆盖更多场景。但注意不是所有网络结构都适合做这件事下面这条非常关键。2.2 CNN感受野的边界为什么卷积网络会输给Transformer大部分CNN去雾网络比如AOD-Net、FFA-Net结构上是编码器-解码器加各种注意力模块。CNN靠卷积核堆叠扩大感受野一个3×3卷积看3×3五层下采样之后理论上能看到全图但这种“看得见”和ViT的“全局建模”不是一回事。去雾任务里A是全局标量任意位置的像素都受同一个A影响同时t(x)又是局部连续的。这意味着网络得同时处理好“全局大气光估计”和“局部透射率连续变化”两件事。CNN在浅层只能做局部滤波要把远山和近处灯柱的信息融合到一起需要经过很多层下采样和上采样信息在那个过程里被平均掉了。表现就是CNN对浓雾区域的全局色偏恢复不到位远处容易偏蓝近处又容易发灰。Vision Transformer的思路完全不同。它把图像切成固定大小patch每个patch拉平成token再通过self-attention让任意两个token直接计算相关性。远山的token和近处灯柱的token在第一个attention block里就能交互不需要路径损耗。这个特性非常契合大气散射模型的结构全局大气光就是一张“所有像素共享的表”Transformer的global receptive field几乎是量身定做。这也是为什么“基于Vision Transformer的图像去雾”这几年复现量明显高于传统CNN方案。需要注意ViT不是没有代价。patch token化会丢失像素级细节patch尺寸越大边缘越肉self-attention的复杂度又随图像尺寸平方增长。因此在去雾任务里常见做法是patch设小8×8或4×4并只在特定分辨率下训练推理时对高分辨率图做切片。这个取舍放在第4、5章展开。2.3 与FFA-Net、DehazeFormer的选型对照ViT不是万能钥匙做技术选型时最常被拿来对比的是三个方向传统CNN注意力FFA-Net、窗口Transformer变体DehazeFormer、以及本标题用的标准ViT。方案全局建模方式训练成本浓雾场景表现复现难度FFA-Net像素注意力通道注意力局部融合较低一般色偏残留低DehazeFormer窗口注意力局部增强中较好细节保持好中标准ViT去雾全局self-attention高色偏恢复最好中高我一般这样取舍数据量在万张以上、且要重点解决全局色偏时优先ViT数据量只有几百张或跑在边缘设备上DehazeFormer这类局部注意力方案更现实。另一个参考维度是训出来的模型的“气质”不同ViT对场景理解更整体往往能把雾里远处的目标“重新画”出来而不是只做对比度拉伸这既是优点也是隐患——它会在训练数据不足时脑补细节导致纹理幻觉。还要提一点标准ViT最初为分类设计输入带class token和位置编码。用在去雾这种密集预测任务时常见改造是去掉class token把patch embedding直接接特征金字塔或逐层上采样回归头。我复现时选择的是“ViT编码器 透射率回归头 全局大气光估计”的方案而不是让ViT直接输出整张清晰图像这个设计决定会在损失函数上看到差别。3. 复现ViT去雾模型的全链路从数据处理到训练命令3.1 环境与数据准备目录约定和数据集的“脾气”先讲环境。拿到这个zip后第一步不是打开代码就训练而是先把python运行环境搭好。我习惯用conda建独立环境python版本3.8或3.9都行如果你手动从python官网下载安装记得勾选Add to PATH。用vscode python环境配置的同学重点是让解释器指向刚建好的conda环境否则pip装到了别的地方后面import报错会绕很久。依赖按这个清单装conda create -n dehaze python3.9 -y conda activate dehaze pip install torch torchvision --index-url https://download.pytorch.org/whl/cu118 pip install opencv-python scikit-image einops timm matplotlib pandas参数说明torch版本不要刻意追新跟自己的显卡驱动匹配即可timm用来加载预训练ViT权重einops方便做patch重组scikit-image提供PSNR/SSIM评估。这套组合跑通“python源码项目介绍使用说明.zip”里的训练和测试脚本基本上不会因为缺库卡住。数据方面主流选择是RESIDE的ITS子集和NH-HAZE。ITS是全合成的雾是物理模型合成的NH-HAZE是真实雾天拍摄更接近落地场景但数量少。目录我一般这样组织data/ RESIDE_ITS/ hazy/ # 模糊图 clear/ # 对应清晰图 NH_HAZE/ hazy/ clear/真实项目中不止一次遇到命名对不上ITS的hazy文件名和clear文件名不是直接同名而是hazy/xxx.png与clear/xxx.png能通过前缀关联。建议写个脚本先校验命中和数量再进训练别等DataLoader报FileNotFoundError才回头查。3.2 自定义数据加载器把成对图像裁成patch并做增强ViT吃的是patch但训练时为了控制显存通常先对原图做RandomCrop再切patch。数据加载器的核心逻辑如下import torch from torch.utils.data import Dataset import cv2, os, random import numpy as np class DehazeDataset(Dataset): def __init__(self, hazy_dir, clear_dir, crop_size256, is_trainTrue): self.hazy_paths sorted(os.listdir(hazy_dir)) self.clear_dir clear_dir self.hazy_dir hazy_dir self.crop_size crop_size self.is_train is_train def __len__(self): return len(self.hazy_paths) def __getitem__(self, idx): hz cv2.imread(os.path.join(self.hazy_dir, self.hazy_paths[idx])) cl cv2.imread(os.path.join(self.clear_dir, self.hazy_paths[idx])) hz cv2.cvtColor(hz, cv2.COLOR_BGR2RGB) / 255.0 cl cv2.cvtColor(cl, cv2.COLOR_BGR2RGB) / 255.0 if self.is_train: # 随机裁剪到固定尺寸保证ViT输入统一 x random.randint(0, hz.shape[1] - self.crop_size) y random.randint(0, hz.shape[0] - self.crop_size) hz hz[y:yself.crop_size, x:xself.crop_size] cl cl[y:yself.crop_size, x:xself.crop_size] # 水平翻转用同一随机数否则雾图和高清图对不上 if random.random() 0.5: hz hz[:, ::-1].copy() cl cl[:, ::-1].copy() hz_t torch.from_numpy(hz).permute(2, 0, 1).float() cl_t torch.from_numpy(cl).permute(2, 0, 1).float() return hz_t, cl_t逻辑说明这里把hazy和clear都按RGB读入并归一化到0~1因为ViT在0~1输入上的训练比在0~255上更稳。随机裁剪之后做翻转时必须用同一个随机数同时对两张图操作否则模型会学到错误映射。这个细节看似基础但在真实复现里经常导致PSNR上不去。参数说明crop_size256是个常见起点显存8G的卡跑batch4基本能承受。如果显存只有6G把crop降到192别降batch降得太狠因为过小的batch会放大BN和LayerNorm的统计噪声。另外排序用sorted保证每次遍历顺序一致便于复现。3.3 Vision Transformer主干构建patch embedding、编码器与透射率头常见的用法是直接用timm加载一个预训练ViT做backbone然后在上面接回归头。下面是这个方案里最核心的模型结构import torch import torch.nn as nn import timm class DehazeViT(nn.Module): def __init__(self, img_size256, patch_size8, embed_dim512, depth6, num_heads8): super().__init__() self.backbone timm.create_model( vit_base_patch8_224, # 预训练ViT-Bpatch8 pretrainedTrue, img_sizeimg_size, num_classes0, # 去掉分类头 ) in_chans self.backbone.embed_dim # 透射率回归头token序列还原成空间图输出单通道 self.deconv nn.Sequential( nn.Conv2d(in_chans, 256, kernel_size3, padding1), nn.GELU(), nn.Conv2d(256, 128, kernel_size3, padding1), nn.GELU(), nn.Conv2d(128, 1, kernel_size1), nn.Sigmoid() # 约束透射率在0~1之间 ) # 大气光每个通道一个可学习标量三个通道共享同一个空间t self.atm_light nn.Parameter(torch.tensor([0.8, 0.8, 0.8])) def forward(self, x): b, c, h, w x.shape tokens self.backbone.forward_features(x) # [B, num_tokens, dim] # 去掉cls token重排成特征图num_tokens h//patch * w//patch tokens tokens[:, 1:, :] side int((tokens.shape[1]) ** 0.5) feat tokens.permute(0, 2, 1).reshape(b, -1, side, side) t_map self.deconv(feat) t_map nn.functional.interpolate(t_map, size(h, w), modebilinear, align_cornersFalse) A torch.sigmoid(self.atm_light) * 0.95 0.05 # 保证A不落到0 # 根据大气散射模型恢复清晰图 J (x - A * (1 - t_map)) / torch.clamp(t_map, min0.05) return J, t_map逻辑说明backbone输出的token序列先丢掉class token再把序列reshape成特征图然后通过几层卷积把特征降成单通道透射率图。Sigmoid把t约束在0~1之间大气光A做成可学习参数而不是逐像素回归能有效避免网络用“伪A”骗过loss。最后的恢复公式必须带clamp防止t接近0时除发出噪声。参数说明patch_size8意味着每8×8像素合成一个token这是去雾任务里比较常用的折中。patch16能把计算量降到1/4但边缘细节会明显糊patch4效果最好但显存翻几倍。depth6在数据量只有几千张时更稳妥depth12容易过拟合。embed_dim512是ViT-B的默认宽度这里不另改。3.4 损失函数与优化器单靠L1会得到一张“会呼吸的雾霾”早期我只用L1损失出来的图亮度和结构都对但颜色边缘糊像覆盖了一层磨砂。常见做法是把像素损失、结构损失和感知损失组合起来import torch.nn.functional as F from torchvision.models import vgg16_bn from pytorch_msssim import SSIM ssim_loss SSIM(data_range1.0, size_averageTrue, channel3) def dehaze_loss(pred, target, pred_feats, target_feats, t_map): l1 F.l1_loss(pred, target) ssim 1 - ssim_loss(pred, target) # 感知损失VGG特征图上的L1逼模型在意结构 percep sum(F.l1_loss(f1, f2) for f1, f2 in zip(pred_feats, target_feats)) / len(pred_feats) # 透射率平滑损失相邻像素t变化不要太剧烈 tv torch.mean(torch.abs(t_map[:, :, 1:, :] - t_map[:, :, :-1, :])) \ torch.mean(torch.abs(t_map[:, :, :, 1:] - t_map[:, :, :, :-1])) return l1 0.2 * ssim 0.1 * percep 0.01 * tv逻辑说明L1负责绝对像素误差SSIM保证局部结构不被过度平滑感知损失用VGG中间层特征比较本质上逼模型生成“人眼看着像清晰图”的结果。最后的TV项是为了防止透射率图出现棋盘格状突变这在深雾区特别常见。参数说明权重0.2/0.1/0.01是我调过一轮后比较稳的起点不是固定标准。如果发现输出色彩偏灰把percep调到0.2如果边缘出现光晕把tv调到0.05。优化器一般用AdamW学习率从1e-4起步配cosine退火optimizer torch.optim.AdamW(model.parameters(), lr1e-4, weight_decay1e-4) scheduler torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_maxepochs)参数说明weight_decay1e-4对ViT这种大参数模型是必要的否则微调时权重很容易漂移T_max设成总训练轮数保证learning rate按余弦曲线衰减到接近0。不要用SGDViT用SGD收敛极慢这是我的血泪经验。4. 训练与指标评估把模型跑起来后怎么判断“真的能去雾”4.1 训练循环与混合精度显存不够时的必调参数训练ViT去雾模型最现实的问题是显存。8G卡跑256×256、batch4勉强够再往上就会OOM。最常见做法是打开PyTorch混合精度AMP在几乎不损失精度的情况下减少一半显存占用scaler torch.cuda.amp.GradScaler() for epoch in range(epochs): for step, (hz, cl) in enumerate(train_loader): hz, cl hz.cuda(), cl.cuda() optimizer.zero_grad() with torch.cuda.amp.autocast(): pred, t_map model(hz) loss dehaze_loss(pred, cl, vgg_feats(pred), vgg_feats(cl), t_map) scaler.scale(loss).backward() scaler.unscale_(optimizer) torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0) scaler.step(optimizer) scaler.update()逻辑说明GradScaler把loss放大再反传避免梯度在fp16下下溢unscale_之后做梯度裁剪防止ViT在刚开始微调时loss冲高。梯度裁剪这个动作很多人会跳过但对Transformer训练几乎是必需品尤其学习率设得偏大时NaN都从这一步防起。参数说明max_norm1.0是常见默认值如果训练过程中发现梯度消失可以放到5.0如果出现loss震荡把它降到0.5。batch4、crop256、fp16、ViT-B/depth6的组合在8G显存下可以在3~4小时跑完100个epoch的小型数据集。如果还想提速把dataloader的num_workers设成4并从机械硬盘换到SSD。4.2 评估指标解读PSNR、SSIM、CIEDE2000与“指标高但看着假”去雾领域最通用的评估指标是PSNR和SSIM。PSNR基于MSE数值越高表示像素误差越小SSIM比较局部亮度、对比度和结构范围0~1。另一类指标是CIEDE2000它对色差更敏感专门用来评估色彩还原。复现时跑一张测试图就能看到三个数字指标含义阈值参考陷阱PSNR像素级峰值信噪比25算可接受对颜色偏差不敏感SSIM结构相似性0.9算良好对模糊边缘不敏感CIEDE2000色差8算合格对压缩格式敏感比较诡异的情况是PSNR和SSIM都很漂亮但图看起来“塑料感”很重。原因是模型学会了近似平均色SSIM只关心局部结构人眼却会盯大面积天空和墙面。我的习惯是不只跑合成测试集还要在真实雾图上看两个地方天空色彩是否自然、远处建筑边缘是否有光晕。一张把雾去干净的图应该在90%缩放的情况下看不出网格感和色块。如果你手头没有真实雾图可以用NH-HAZE的验证集替代。注意NH-HAZE图像偏暗直接套用RESIDE上训好的模型会输出偏灰的图这是数据集域差异不是模型坏了。评估脚本用scikit-image一行就能算from skimage.metrics import peak_signal_noise_ratio, structural_similarity psnr peak_signal_noise_ratio(clear, pred, data_range1.0) ssim structural_similarity(clear, pred, channel_axis-1, data_range1.0)参数说明data_range必须和输入数据范围一致这里输入是0~1所以写1.0如果写255数值会非常难看在star。channel_axis-1表示通道在最后一维因为cv2或imread读出来的是HWC。4.3 测试脚本与模型保存用最佳checkpoint做推理训练过程中我一般每个epoch结束都在验证集上算一次PSNR把最优checkpoint另存不拿最后一个epoch的结果直接用因为最后一个epoch往往已经收敛到头了选验证集最优的才可靠import torch def inference(model, hazy_tensor, devicecuda): model.eval() with torch.inference_mode(): haz hazy_tensor.unsqueeze(0).to(device) pred, t_map model(haz) return pred.squeeze(0).cpu(), t_map.squeeze(0).cpu() # 保存最佳模型 best_state model.state_dict() torch.save({model: best_state, config: vit_b_patch8}, best_dehaze.pt)逻辑说明torch.inference_mode()比no_grad()更激进会关闭一些跟踪记录推理显存占用更低。模型输出的是0~1的tensor最终保存成图片时要先permute回HWC再乘255用matplotlib的imshow可以直接看tensor。透射率图也值得存下来它能直观展示模型对浓雾区域的判断。参数说明如果推理时输入分辨率不是训练时的256ViT的位置编码会因为token数量不同而报错。常见做法是重新插值位置编码或像第3.3节那样让backbone支持img_size动态变化。timm的create_model传img_size参数后位置编码会自动插值但插值会损失一点精度追求最可靠性能时最好训练和推理分辨率保持一致。5. ViT去雾避坑指南训练翻车与“去雾幻视”的4个典型现场5.1 加载预训练模型报错位置编码长度与patch尺寸不匹配现象代码里用timm.load_pretrained加载ViT-B/16权重但模型定义的是patch8第一次forward就报size mismatch提示pos_embed维度对不上。原因ViT的位置编码数量由patch尺寸决定patch16训练出的权重是196个位置编码patch8模型需要784个位置编码。直接加载肯定报错就算用了插值视觉transformer的位置编码插值后会引入高频伪影早期训练阶段会出现方块状条纹。解决两种做法选一个。一是让timm的create_model自己处理img_size和patch_size再加载权重二是干脆不加载预训练从零训练ViT去雾。我个人经验是数据量超过2万张时从零训练不会比预训练微调差而且省去一堆兼容性麻烦。如果一定要用预训练记得设置vit_base_patch8_224这个架构名而不是patch16的。5.2 训练loss下降但输出图发灰大气光回归“骗loss”了现象PSNR在涨SSIM也还行但生成的图灰蒙蒙的像整体加了一层半透明灰度遮罩天空和白色物体完全看不出层次。原因网络发现一个“作弊”路径——把透射率t整体推向1大气光A推向训练集平均亮度。这样输出的J≈x减去一个常量导致loss下降但全局色偏没被纠正。本质是我前面提到过的大气光不需要逐像素回归逐像素A会被loss压向中值。解决把A改成可学习标量每个通道一个用sigmoid限制在0.05~0.95之间透射率头加Sigmoid后再把t做一次约束例如t_min0.05防止纯黑像素被除爆。训练完成后如果仍发灰可以做一个后处理用torch.clamp对t做0~1裁剪后再做gamma增强gamma值取0.7左右通常能把亮度拉回来。5.3 浓雾区域死黑或过曝合成数据的透射率分布太“温柔”现象合成测试集上表现不错一到真实浓雾图上最浓的雾区要么被压成全黑要么过曝成白色一片。原因RESIDE的合成雾在t0.1~1之间均匀采样真实雾图的t在浓雾区经常低于0.05。模型没见过后半段分布遇到分布外输入就开始乱猜。这不是网络结构问题是训练数据分布的盲区。解决训练时对hazy图做在线扰动——随机gamma变换、亮度偏移、饱和度缩放让模型见过更极端的输入另一个有效手段是在透射率回归头的输出上做随机遮挡增强强制模型在局部信息缺失时依赖全局上下文。更可靠的方案是加一个domain adaptation分支把真实雾图作为无标签输入对齐特征分布但这属于进阶内容初期不建议啃。5.4 推理显存爆炸与速度慢全局attention在大图上“起飞”现象训练时256×256很顺推理时拿1080P监控画面直接跑OOM或等得人想放弃。原因self-attention的复杂度按像素数平方增长。256×256是16384个token1080P按patch8算超过1万个token但实际计算时是先切patch再计算视频帧1080P对应百万级像素token数直接爆掉。解决常见做法是切片推理加边缘羽化。把大图切成256×256的块块与块之间保留overlap16像素推理完成后对overlap区域做线性加权融合这样既能避免切块接缝也不用吃满显存。视频流场景再叠加帧间EMA简单说就是new_frame0.8last_result0.2current_result能压掉闪烁。ONNX导出后配合OpenVINO做半精度推理1080P单帧能跑到百毫秒级别这个优化放到最后一章细说。6. 验证模型的最后一公里可视化解读与复用ViT去雾模型6.1 注意力图可视化解释“凭什么把这块雾去掉”训练完光看PSNR不够还需要让人信服。最简单有效的验证手段是把ViT最后一层的attention map画出来。我习惯直接搬python数据分析与可视化那套matplotlib三件套提取backbone最后一个block的注意力权重在所有头上取平均再upsample回输入尺寸叠加到原图上。高亮区域代表模型在重建该像素时重点依赖的远处信息。在这个可视化里能看到一个有趣的可靠现象远处山体和天空区域的高亮通常指向图像顶部或左右边缘说明模型确实在用全局大气光线索推断透射率而不是只在局部磨皮。这个图放进论文或答辩材料里说服力比一行PSNR强得多。6.2 把去雾模型接进自己的算法链路从单张图到简单视频流实际项目里去雾很少是终点更多是检测或识别的预处理。接进链路时要注意三件事一是色彩校准去雾模型输出的颜色会偏暖或偏冷下游如果是车牌识别加一个颜色校正矩阵会更可靠二是分辨率适配检测模型一般吃640×640或更小完全没必要在推理时对大图做全分辨率去雾直接切片后resize即可三是时序平滑对监控视频用前文提到的EMA就能消除闪烁但注意运动目标会拖尾所以只对大气光A和透射率做平滑不对最终像素做强平滑能在平滑和拖尾之间找到平衡点。6.3 坑踩完之后我留下的验证习惯回头总结这个项目我最大的教训是“指标只是门票视觉才是答案”。每轮实验跑完我都不只看PSNR而是固定挑三张图盯一张大雾天空、一张中景街道、一张近景车辆。哪怕指标涨了0.3dB如果三张图的边缘更脏或颜色更假这轮改动就不进主线。另一个习惯是把每次训练的超参数、loss曲线和配置记录在一个文本里防止一周后自己都记不清哪个配置产生了哪张图。去雾模型这种“黑匣子”属性很强的方向可控实验记录比模型本身更值钱。希望帮到你。本文还有配套的精品资源点击获取