简介:基于Vision Transformer的图像去雾算法研究与实现资源,面向图像去雾、视觉Transformer方向的算法工程师、研究生及项目开发者,提供可直接运行的Python源码、预训练权重和完整使用说明。项目将ViT引入图像去雾任务,预训练权重存放于My_best_model文件夹,支持按数据集划分选择对应权重;训练参数集中配置在option.py中,可通过--pretrain_weights指定权重路径,通过--train_ps控制输入补丁大小(默认128),方便复现实验、继续训练或调整模型输入尺度。资源包共340个文件,以204个py源码文件为核心,配套16个yaml模型与训练配置、12个csv训练过程记录、9个ipynb交互式示例、9个txt及8个md说明文档,另有39个png、10个gif等可视化结果图,压缩包整体156.34MB,目录按代码、配置、权重、文档划分,便于查阅与二次开发。目前已有468人学习/下载,适合想要快速理解ViT去雾原理、开展模型训练调参或在既有代码基础上扩展应用的读者。
1. 基于Vision Transformer的图像去雾算法:从源码跑通到效果调优的全过程
图像去雾在工程落地里是个看着简单、做起来却很闹心的方向。基于Vision Transformer的图像去雾算法,把卷积的局部感受野换成全局自注意力,确实在浓雾区域和远景细节上比传统CNN方法稳不少,但真正把这份python源码跑通、再迁移到自己的数据上,中间要跨过的坑比想象中多得多。这篇文章从一个一线工程师的角度,把这条技术路线从原理到训练、推理、评估完整拆开,适合两类人:一是想做图像复原方向毕业设计或课题预研的学生,二是想把去雾能力集成到现有视觉系统中的开发者。读完你会知道这个方案能不能用、参数怎么设、崩了看哪里。
2. 去雾算法为什么需要Vision Transformer:从大气散射模型到全局感受野
2.1 大气散射模型与物理先验:去雾到底在解什么方程
图像去雾不是简单的对比度拉伸,几乎所有经典方法都建立在大气散射模型上。这个模型用一句话概括:相机接收到的光 = 物体反射光经过雾气衰减后的部分 + 环境光被散射后进入相机的部分。写成公式就是 I(x) = J(x)·t(x) + A·(1 - t(x)),其中 I(x) 是观察到的雾图,J(x) 是我们要恢复的清晰图像,t(x) 是透射率,A 是全局环境光。
去雾任务的核心就是从这个方程里反解出 J(x)。但问题在于,一张雾图里 J、t、A 全是未知数,一个方程三个未知量,数学上这叫病态问题。传统方法比如暗通道先验,是通过统计规律先估 A 和 t,再反算 J,但对天空区域和白色物体经常失效。深度学习方法换了个思路:用大量成对的雾图和清晰图训练网络,让网络直接从 I 回归到 J,把物理模型隐式地学进网络参数里。这就是基于Vision Transformer的图像去雾算法和传统CNN去雾的根本区别。
2.2 CNN去雾的边界:局部卷积为什么搞不定浓雾区域
CNN做去雾已经有很多成熟工作,比如AOD-Net、DehazeNet,它们用卷积层堆叠出一个映射网络。卷积操作的感受野是局部的,一层3×3卷积只能看到周围几个像素。虽然通过加深网络可以扩大感受野,但实际效果有限:浓雾区域的像素值被环境光严重污染,局部邻域里的信息几乎全是雾,卷积核学到的特征也就缺乏区分度。
另一个实际问题是空间不变性。CNN的卷积核是权值共享的,同一套卷积核作用在图像的不同位置。但雾的浓度在空间上分布不均匀,近处雾薄、远处雾厚,局部卷积很难同时适应不同雾浓度的区域。你可以这样理解:CNN像是在用同一个放大镜看整幅图,而雾图需要的是不同区域给不同倍率的矫正,这恰恰是全局建模才能做到的。
2.3 ViT的全局建模:patch化与自注意力如何改变去雾的解题路径
Vision Transformer把图像切成固定大小的patch,比如16×16,然后把每个patch展平成一个token,通过自注意力机制在token之间计算相关性。这意味着任意两个patch之间可以直接建立依赖,距离远的像素也能彼此关联。对去雾来说,这个特性的价值很直接:远处物体的颜色信息可以通过自注意力传递到近处被雾污染的区域,帮助网络还原真实颜色。
实际结构上,去雾ViT通常采用encoder-decoder架构。Encoder部分用标准ViT的self-attention做特征提取,Decoder部分用transposed convolution或者pixel shuffle把特征图恢复到原始分辨率。skip connection在去雾任务里尤其重要,因为encoder下采样会丢失边缘细节,通过跳跃连接把浅层特征直接送到decoder,能保住纹理结构。这套思路和U-Net类似,但骨干网络从CNN换成了ViT,全局建模能力是核心收益。
提示:ViT做去雾的代价是计算量远高于CNN,尤其是输入分辨率大的时候,self-attention的复杂度是O(n²)。对1080p图像直接跑ViT-Base是不现实的,一般会先下采样或者用分块策略。
3. 跑通源码前的数据与工程准备:RESIDE数据集与Python环境三板斧
3.1 数据格式:合成雾图的生成逻辑与目录结构
图像去雾算法的训练数据主流是RESIDE数据集,它包含室内和室外场景的合成雾图。合成的逻辑就是前面说的大气散射模型:从清晰图像J出发,随机生成透射率t和环境光A,再把它们合成雾图I。也就是说,RESIDE虽然叫真实世界数据集,但训练集的雾是人为合成的,这让模型在真实雾图上天然存在域差异。
拿到源码包后先看目录结构,通常包含 train/、val/、test/ 三个目录,每个目录下又分 hazy/(雾图)和 clear/(清晰图)两个子目录,文件名一一对应。有些版本还会附上透射率图 depth/,做消融实验时用得上。准备自己的数据时,尽量保持同样的目录命名规则,这样源码里的数据加载器不需要改动就能直接用。
3.2 Python环境配置:CUDA、PyTorch与ViT依赖的最小清单
跑ViT去雾代码不需要花哨的深度学习框架,PyTorch就够了。Python版本建议3.8到3.10,PyTorch用1.12以上或者2.x版本。GPU方面,训练ViT-Base至少需要11GB显存,推荐16GB以上;如果只有8GB显存,需要把patch size调大、batch size调小,或者换ViT-Tiny。
环境配置最常见的问题就是CUDA版本和PyTorch对不上,训练时莫名其妙报CUDA error。我一般会用conda单独建一个环境,按下面的顺序装依赖,基本上不会翻车:
conda create -n dehaze python=3.9 -y conda activate dehaze pip install torch==2.1.1 torchvision==0.16.1 --index-url https://download.pytorch.org/whl/cu118 pip install timm==0.9.12 opencv-python==4.8.1.78 numpy==1.24.4 tensorboard==2.14.0 pip install einops==0.7.0 tqdm scikit-image==0.21.0这里把torch和torchvision通过--index-url指定了CUDA 11.8的wheel包,避免pip默认安装CPU版本或者CUDA版本不匹配。timm库是用来加载预训练ViT权重的,einops用于重排张量维度,scikit-image提供PSNR和SSIM的计算接口。
注意:CUDA 11.8对应PyTorch的cu118,如果你本机是CUDA 12.x,要用cu121或cu124的index-url,否则import torch会报libcudart.so找不到的错误。用nvidia-smi看到的是驱动支持的CUDA版本,不是PyTorch实际用的运行时版本,两者不要求一致。
3.3 预处理与数据加载器:把成对雾图/清晰图喂给模型的代码
数据加载器是整个训练流程的入口,也是新手最容易写错的地方。核心要做三件事:读取配对的雾图和清晰图、做随机裁剪和数据增强、转成Tensor并归一化。下面这个数据加载器是去雾项目里最常见的一种实现,我基于源码包里的loader做了简化说明:
import os import random import cv2 import numpy as np import torch from torch.utils.data import Dataset class DehazeDataset(Dataset): def __init__(self, hazy_dir, clear_dir, crop_size=256, augment=True): self.hazy_paths = sorted([os.path.join(hazy_dir, f) for f in os.listdir(hazy_dir)]) self.clear_paths = sorted([os.path.join(clear_dir, f) for f in os.listdir(clear_dir)]) assert len(self.hazy_paths) == len(self.clear_paths), "雾图和清晰图数量必须一致" self.crop_size = crop_size self.augment = augment def __len__(self): return len(self.hazy_paths) def __getitem__(self, idx): hazy = cv2.imread(self.hazy_paths[idx]) clear = cv2.imread(self.clear_paths[idx]) hazy = cv2.cvtColor(hazy, cv2.COLOR_BGR2RGB) clear = cv2.cvtColor(clear, cv2.COLOR_BGR2RGB) # 随机裁剪到固定尺寸,训练时不用整图,省显存 h, w = hazy.shape[:2] x = random.randint(0, max(0, w - self.crop_size)) y = random.randint(0, max(0, h - self.crop_size)) hazy = hazy[y:y+self.crop_size, x:x+self.crop_size, :] clear = clear[y:y+self.crop_size, x:x+self.crop_size, :] if self.augment: # 随机水平翻转和旋转90度,提升数据多样性 if random.random() < 0.5: hazy = hazy[:, ::-1, :] clear = clear[:, ::-1, :] if random.random() < 0.5: hazy = np.rot90(hazy, k=1) clear = np.rot90(clear, k=1) # HWC转CHW,并归一化到[0,1],ViT对输入范围敏感 hazy = torch.from_numpy(hazy.transpose(2, 0, 1).copy()).float() / 255.0 clear = torch.from_numpy(clear.transpose(2, 0, 1).copy()).float() / 255.0 return hazy, clear这个loader的关键点在于随机裁剪尺寸crop_size,默认256×256。ViT的patch size如果不是整除关系会有问题,比如patch size是16,那输入尺寸必须是16的整数倍,256正好整除。augment开关在验证时要关掉,保证评测结果可复现。排序时用sorted()可以保证雾图和清晰图按文件名一一对应,否则训练时配对错乱,损失曲线直接崩掉。
4. 模型训练与推理落地:从ViT-Base改造到去雾头的最小可跑配置
4.1 模型结构拆解:encoder-decoder主干与去雾头的连接方式
基于Vision Transformer的去雾模型基本可以理解为:ViT做encoder,一个轻量decoder还原分辨率,最后接一个去雾输出头。用timm库加载预训练ViT-Base的权重是最省事的做法,关键是把末尾的分类头去掉,只保留encoder部分。下面是一个典型的模型构建代码:
import torch import torch.nn as nn import timm class ViTDehaze(nn.Module): def __init__(self, img_size=256, patch_size=16, in_chans=3, embed_dim=768, decoder_dim=256): super().__init__() # 加载预训练ViT-Base encoder,去分类头 self.encoder = timm.create_model('vit_base_patch16_224', pretrained=True, num_classes=0) # 直接让ViT接受256x256的输入,不需要改位置编码 # 因为timm默认会自动插值位置编码到不同分辨率 self.decoder = nn.Sequential( nn.ConvTranspose2d(embed_dim, decoder_dim, kernel_size=4, stride=2, padding=1), nn.ReLU(inplace=True), nn.ConvTranspose2d(decoder_dim, decoder_dim // 2, kernel_size=4, stride=2, padding=1), nn.ReLU(inplace=True), nn.ConvTranspose2d(decoder_dim // 2, decoder_dim // 4, kernel_size=4, stride=2, padding=1), nn.ReLU(inplace=True), ) self.output_head = nn.Conv2d(decoder_dim // 4, 3, kernel_size=3, padding=1) self.upsample = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=False) def forward(self, x): B, C, H, W = x.shape # ViT输出shape: (B, N+1, embed_dim),需要去掉cls token并重塑 feats = self.encoder.forward_features(x) # (B, 257, 768) feats = feats[:, 1:, :] # 去掉cls token,因为是像素级任务 N = feats.shape[1] side = int(N ** 0.5) feats = feats.permute(0, 2, 1).reshape(B, -1, side, side) out = self.decoder(feats) out = self.upsample(out) out = self.output_head(out) return out def forward_features(self, x): return self.encoder.forward_features(x)decoder用转置卷积逐步把16倍下采样的特征图恢复分辨率,输出head是一个3×3卷积层输出3通道RGB。最后的Upsample层把分辨率对齐到输入大小。这个结构能跑通的前提是timm内部自动处理了位置编码插值,如果换成自定义的ViT实现,输入尺寸变了会直接报位置编码维度不匹配的错误。
4.2 损失函数设计:感知损失与频域约束的搭配
去雾模型不能用单一的L1或MSE损失,不然训练出来的图偏平滑、细节糊。业界常用的组合是L1损失 + 感知损失,有些工作还会加频域损失。L1损失保证像素级别的颜色准确,感知损失用VGG网络的特征图计算距离,保证恢复出来的纹理和结构感知上接近清晰图。下面是训练脚本里常见的损失函数组合:
import torch.nn.functional as F from torchvision.models import vgg16 class DehazeLoss(nn.Module): def __init__(self, perceptual_weight=0.05): super().__init__() vgg = vgg16(pretrained=True).features[:16] # 取到conv3 self.perceptual = vgg.eval() for p in self.perceptual.parameters(): p.requires_grad = False # 冻结感知网络参数 self.perceptual_weight = perceptual_weight def forward(self, pred, target): l1 = F.l1_loss(pred, target) # 感知损失:在VGG特征空间计算L1距离 pred_feat = self.perceptual(pred) target_feat = self.perceptual(target) perc = F.l1_loss(pred_feat, target_feat) return l1 + self.perceptual_weight * perc感知损失权重perceptual_weight设成0.05比较安全。设太大,模型会过度关注高频纹理,雾区域会出现伪影;设太小,感知约束不起作用,退化成普通L1。还有个细节:perceptual网络输入要求归一化到ImageNet的均值和标准差,去雾模型的输出是[0,1]范围,直接喂给VGG会分布不匹配,最好在perceptual损失计算前做一次normalize。
4.3 训练脚本与超参数:learning rate、batch size与warmup的参考配置
ViT训练和CNN训练有个很大的不同:ViT对优化器非常敏感,直接上SGD基本不收敛,需要用AdamW加warmup。下面是一个可以直接套用的训练循环核心逻辑:
import torch from torch.optim import AdamW from torch.optim.lr_scheduler import CosineAnnealingLR model = ViTDehaze(img_size=256, patch_size=16).cuda() optimizer = AdamW(model.parameters(), lr=1e-4, weight_decay=1e-4) scheduler = CosineAnnealingLR(optimizer, T_max=200, eta_min=1e-6) criterion = DehazeLoss(perceptual_weight=0.05).cuda() dataloader = torch.utils.data.DataLoader( DehazeDataset('train/hazy', 'train/clear'), batch_size=8, shuffle=True, num_workers=4, pin_memory=True ) for epoch in range(200): model.train() # warmup前5个epoch,lr从1e-5线性升到1e-4 if epoch < 5: lr = 1e-5 + (1e-4 - 1e-5) * epoch / 5 for g in optimizer.param_groups: g['lr'] = lr for hazy, clear in dataloader: hazy, clear = hazy.cuda(), clear.cuda() pred = model(hazy) loss = criterion(pred, clear) optimizer.zero_grad() loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) optimizer.step() scheduler.step() if epoch % 10 == 0: print(f"Epoch {epoch}, Loss: {loss.item():.4f}") torch.save(model.state_dict(), f"checkpoints/vit_dehaze_{epoch}.pth")batch size在8GB显存下只能设到4,16GB可以设8。clip_grad_norm_是ViT训练的常规操作,不加的话训练刚开始loss会剧烈震荡。warmup阶段手动覆盖optimizer的lr,能有效避免ViT在初始阶段梯度爆炸。我用这个配置在RTX 3090上训练200个epoch大概需要18到24小时,Loss能从0.15降到0.04左右。
| 超参数 | 推荐值 | 备注 |
|---|---|---|
| batch size | 8(16GB显存)/ 4(8GB显存) | 显存不足时优先减batch |
| learning rate | 1e-4 | AdamW专用,SGD要调小 |
| warmup epochs | 5 | 线性从1e-5升到1e-4 |
| weight decay | 1e-4 | 防止过拟合 |
| crop size | 256 | patch size为16时的安全值 |
| 感知损失权重 | 0.05 | 太大会出伪影,太小没效果 |
5. 去雾模型训练与部署避坑指南:五个高频翻车现场
5.1 图像尺寸与patch大小不匹配:位置编码张量崩溃
现象:forward的时候报错"size mismatch for pos_embed",或者输出的特征图边长开方后是小数。
原因:ViT把输入切成patch,位置编码数量是预先定义好的。切出来的patch数量对不上位置编码的数量,最常见的就是输入尺寸不是patch size的整数倍。
解决:统一约定输入的H和W能被patch size整除。用patch size 16时,输入尺寸选256、224、384都可以,但别用300×300这种数。代码里最好加一行断言,提前报错而不是在forward中途莫名其妙崩。
5.2 损失函数震荡不收敛:感知损失权重与学习率不匹配
现象:训练前几个epoch损失函数上下跳动,过了50个epoch还没有下降趋势,甚至越来越大。
原因:感知损失来自预训练VGG特征空间,它的梯度量级和L1损失不在一个数量级上,学习率稍微大一点就会出现梯度震荡。另外ViT本身对学习率就敏感。
解决:把perceptual_weight先降到0.01试跑10个epoch看趋势,收敛稳定后再调回0.05。同时确认是否做了warmup,没有warmup的ViT在初始阶段大概率震荡。还有一个常见低级错误:perceptual网络的参数没有冻结,导致感知损失的梯度把VGG也更新了,这时候损失曲线会出现诡异的周期性波动。
5.3 推理结果偏灰:反透射率映射的数值陷阱
现象:模型训练时PSNR很高,但推理出来的图整体偏灰,像蒙了一层纱,颜色饱和度不足。
原因:数据集里的清晰图是sRGB颜色空间,但有些源码在预处理时做了线性化变换,或者归一化时用了错误的均值和标准差。更隐蔽的原因是模型输出层的激活函数,如果用了Sigmoid但训练时数据是线性归一化的,输出会被压缩到[0.3, 0.7]这个区间附近,看起来就是灰蒙蒙的。
解决:检查模型最后一层是直接输出还是要过Sigmoid/Tanh,和训练时保持一致。推理时建议图先除以255归一化,输出再乘255转回,不要整出两套归一化标准。实在偏灰,可以在后处理时做一个自动色阶拉伸,但要谨慎,过度拉伸会引入色带。
5.4 显存溢出:ViT在大分辨率图像上的显存优化策略
现象:训练好之后想推理一张1920×1080的雾图,直接喂给模型报CUDA out of memory。
原因:ViT对序列长度极度敏感,1920×1080除以patch 16,序列长度是120×68约8000多个token,自注意力的中间激活直接撑爆显存。
解决:常见做法是分块推理,把大图切成512×512的重叠块分别推理再拼接,重叠区域用线性加权融合避免接缝。另一个更省事的方法是用adaptive average pool把特征图压到固定尺寸再送ViT,但会丢失细节。我一般用第一种,用下面的代码做切块推理:
def inference_large_image(model, image, crop_size=512, stride=256): model.eval() h, w = image.shape[:2] output = np.zeros((h, w, 3), dtype=np.float32) weight = np.zeros((h, w, 1), dtype=np.float32) for y in range(0, h - crop_size + 1, stride): for x in range(0, w - crop_size + 1, stride): patch = image[y:y+crop_size, x:x+crop_size] patch_tensor = torch.from_numpy(patch.transpose(2, 0, 1)).unsqueeze(0).float() / 255.0 with torch.no_grad(): pred = model(patch_tensor.cuda()).cpu().squeeze(0).permute(1, 2, 0).numpy() output[y:y+crop_size, x:x+crop_size] += pred weight[y:y+crop_size, x:x+crop_size] += 1.0 # 重叠区域加权平均,消除边界割裂感 return output / np.maximum(weight, 1.0)5.5 训练集与测试集雾浓度分布差异:评估指标虚高
现象:在RESIDE合成测试集上PSNR能到30以上,换到真实雾图上一测,视觉质量和指标双双拉胯。
原因:这几乎是所有去雾工作都会遇到的老大难。RESIDE的合成雾图用的是均匀大气光值和全局透射率,真实世界的雾浓度随距离连续变化,还有非均匀的散射介质。模型学到的是合成雾的分布,而不是物理雾的分布。
解决:想提升真实场景效果,至少做两件事。第一,训练时使用domain randomization,随机调整雾图合成参数,扩大数据分布覆盖。第二,用去雾结果做自监督微调,把真实雾图输入模型、模型输出再合成伪雾图,计算一致性损失。行业里叫cycle-consistency,实现起来不复杂,但能明显缓解域偏移。
6. 评估指标与模型融合技巧:PSNR之外的第二只眼
6.1 PSNR与SSIM的局限:补一个无参考指标
训练时盯着PSNR提升没有错,但要清醒地认识到PSNR对空间结构不敏感。一张图整体平移几个像素,PSNR会掉很多,但人眼看起来几乎没差别。SSIM对局部结构更敏感,但它的全局池化方式会掩盖局部伪影。所以我在模型选型时必看三样东西:PSNR、SSIM和一个无参考指标。
6.2 可视化对比图:暗部细节与边缘锐度
指标只是筛子,最终要让眼睛说话。我的习惯是把三组图并排放:雾图、模型输出、清晰图(如果有),然后裁三个局部区域放大——天空区域看有没有banding色带,暗部区域看细节是否糊死,边缘区域看是否有halo伪影。这些区域恰好是去雾模型最容易出问题的地方,指标再高,这几个区域翻车也不能上线。
6.3 一个推理优化的技巧:半精度推理
模型训练用FP32,推理阶段完全可以换成FP16,在RTX系列GPU上显存占用直接减半,速度能提升30%到50%。
model = model.half() # 先把模型转成half model.eval() hazy = hazy.half().cuda() with torch.no_grad(): pred = model(hazy).float() # 输出转回float32再做后处理需要注意的是,归一化和后处理阶段还是用float32,不然OpenCV和numpy的数值操作会出现精度损失。另外通道数少于3的中间层在FP16下偶尔有溢出风险,但去雾模型输出是RGB三通道,实际使用中没碰到过问题。
去雾这件事我做了不少项目,踩得最多的坑不是模型结构而是数据分布。合成雾与真实雾之间的距离,决定了模型上线之后的真实效果。在评估时只看PSNR不看可视化图,是被指标骗过之后总结的血泪经验。现在我的流程里,评估阶段PSNR和SSIM只作为门槛,视觉对照才是最终审批人。希望上面这些思路和踩坑记录对你有用,帮你在做基于Vision Transformer的图像去雾时少走弯路。
提示:训练好模型后用一张完全没见过的真实雾图做冒烟测试,而不是先跑测试集。如果真实图效果还行,再去调指标;如果翻车了,先回头检查数据分布再做模型改动,这个顺序能省很多调试时间。
本文还有配套的精品资源,点击获取