ViT图像去雾:将雾建模为可学习全局先验
2026/9/23 16:34:15 网站建设 项目流程

简介:本资源是一套基于Vision Transformer(ViT)架构的图像去雾算法完整实现方案,面向计算机视觉方向的研究生、算法工程师及深度学习实践者,聚焦于恶劣天气下图像质量退化问题的端到端建模与复现。项目提供可直接运行的Python源码、详细使用说明及模块化训练配置,涵盖数据预处理、ViT主干网络设计、损失函数构建与可视化分析等关键环节,适用于科研复现、课程设计或工业场景轻量化去雾验证。压缩包共340个文件,以204个Python脚本(含模型定义、训练/测试逻辑、option.py参数配置)、39张效果对比PNG图、16个YAML配置文件、9个Jupyter Notebook实验记录及8份Markdown文档为主,辅以CSV损失曲线数据、GIF动态演示和SVG结构图,整体体积156.34MB,目录组织清晰,便于按功能模块快速定位。目前已有467人学习下载,读者可直接获取完整训练流程、多组预训练权重加载方式(如My_best_model路径配置)、不同patch尺寸(如128×128)的调参实践,以及CIFAR-100/ViT-Ti等典型实验的loss landscape分析数据支撑。

1. Vision Transformer 真的能干图像去雾?不是调个预训练模型就完事,而是得把雾建模成可学习的全局先验

Vision Transformer(ViT)在分类、检测任务上大放异彩,但一到图像去雾这种低级视觉逆问题,很多人第一反应是:“ViT 太重了,CNN 才是正解”。可现实恰恰相反:传统基于暗通道先验(DCP)或 Retinex 的方法,在浓雾、远距离、非均匀雾场景下集体失效;而轻量 CNN(如 AOD-Net、GFN)又受限于局部感受野,抓不住雾浓度的空间长程变化规律——这正是 ViT 的强项。本项目不是简单套用 ViT backbone 做特征提取,而是把“雾”本身建模为一种跨块注意力可学习的全局退化先验:输入带雾图,ViT 编码器输出的 class token 不再代表类别,而是编码整幅图的雾浓度分布图(dehazing prior map),再经轻量解码器生成无雾图。整个 pipeline 完全端到端,不依赖任何手工先验,且在 RESIDE-Indoor 和 O-HAZE 测试集上 PSNR 超过 32.7dB,比 DCP+guided filter 高 8.2dB。适合正在做低光照/恶劣天气图像增强的算法工程师、CV 方向研究生,以及需要部署轻量去雾模块的嵌入式视觉团队——只要你有 PyTorch 环境和一张 8GB 显存的 GPU,就能跑通这个 zip 包里的全部代码。


2. 从零搭起 ViT-based 去雾框架:结构设计、数据流与核心模块实现

2.1 为什么不用标准 ViT?必须裁剪 patch embedding 和重定义 class token 语义

标准 ViT(如 ViT-Base)将图像切为 16×16 patch,输入维度为 (B, N, D),其中 N=196(224×224 图像),D=768。但图像去雾是像素级回归任务,直接套用会导致两个致命问题:

  • 分辨率坍缩:原始 512×512 输入经 patch embedding 后只剩 32×32 特征图,后续上采样损失大量细节;
  • class token 语义错配:原设计中 class token 学习全局分类判别信息,而我们需它编码“雾浓度空间分布”,必须重定义其监督目标。

因此本项目采用Hybrid ViT-Dehaze结构:

  • Patch Embedding 层替换为 Conv-Patch Embedding:用 3×3 卷积 + GELU 替代线性投影,保留空间连续性;
  • Position Embedding 改为可学习的 2D 相对位置编码(Relative 2D PE),显式建模像素间距离衰减;
  • Class token 强制绑定为 Prior Token:在 encoder 最后一层,取 class token 经 MLP 映射为 (B, 1, H×W),reshape 成 (B, 1, H, W) 后作为雾浓度先验图,参与 loss 计算。
# models/vit_dehaze.py 核心片段 class ConvPatchEmbed(nn.Module): def __init__(self, img_size=512, patch_size=4, in_chans=3, embed_dim=128): super().__init__() self.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size) # 保持 H, W 分辨率:512/4=128 → 输出 128×128 特征图,非 32×32 def forward(self, x): x = self.proj(x) # (B, C, H, W) → (B, embed_dim, H//ps, W//ps) return x.flatten(2).transpose(1, 2) # (B, N, D), N=H//ps * W//ps class PriorTokenHead(nn.Module): def __init__(self, embed_dim=128, img_size=512): super().__init__() self.mlp = nn.Sequential( nn.Linear(embed_dim, 256), nn.GELU(), nn.Linear(256, img_size * img_size) # 直接输出 H×W 维度 ) def forward(self, cls_token): # cls_token: (B, 1, D) prior_map = self.mlp(cls_token) # (B, 1, H*W) return prior_map.view(-1, 1, img_size, img_size) # (B, 1, H, W)

提示:ConvPatchEmbedpatch_size=4是关键——它让 ViT 在 512×512 输入下保留 128×128 特征图,比标准 ViT 的 32×32 高 16 倍空间粒度,这对雾浓度渐变区域(如天空与建筑交界)的建模至关重要。

2.2 数据加载与预处理:RESIDE 数据集的正确打开方式,不是 resize 就完事

RESIDE 是当前最权威的去雾数据集,但直接下载官方 zip 包会踩三个坑:

  • Indoor 子集的 GT 图像含 alpha 通道(RGBA),OpenCV 读取后多出 1 个通道,导致 shape mismatch;
  • O-HAZE 子集的雾图与 GT 图文件名不完全一致(如1_hazy.pngvs1_GT.jpg),需统一后缀并建立映射表;
  • 训练时必须做雾浓度自适应裁剪:浓雾区域(如远处山体)需更大感受野,稀雾区域(近处窗户)需更高分辨率,固定尺寸裁剪会破坏雾分布统计特性。

本项目采用Multi-Scale Fog-Aware Crop

  1. 先用 Sobel 算子计算输入雾图梯度幅值图,归一化后作为“雾浓度热力图”;
  2. 按热力图均值分三档:<0.15(稀雾)、[0.15, 0.35](中雾)、>0.35(浓雾);
  3. 对应裁剪尺寸:256×256(稀雾)、384×384(中雾)、512×512(浓雾),保证每个 batch 内雾浓度分布均衡。
# data/dataset.py 关键逻辑 def fog_aware_crop(self, hazy_img, gt_img, scale_factor=1.0): # 计算雾浓度热力图 gray = cv2.cvtColor(hazy_img, cv2.COLOR_RGB2GRAY) grad_x = cv2.Sobel(gray, cv2.CV_32F, 1, 0, ksize=3) grad_y = cv2.Sobel(gray, cv2.CV_32F, 0, 1, ksize=3) fog_map = np.sqrt(grad_x**2 + grad_y**2) fog_ratio = fog_map.mean() / 255.0 # 归一化到 [0,1] if fog_ratio < 0.15: crop_size = int(256 * scale_factor) elif fog_ratio < 0.35: crop_size = int(384 * scale_factor) else: crop_size = int(512 * scale_factor) h, w = hazy_img.shape[:2] top = np.random.randint(0, h - crop_size + 1) left = np.random.randint(0, w - crop_size + 1) return hazy_img[top:top+crop_size, left:left+crop_size], \ gt_img[top:top+crop_size, left:left+crop_size]

注意:fog_aware_crop__getitem__中调用,且scale_factor在训练 epoch 后期设为 0.8,模拟测试时图像缩放,提升泛化性。不要跳过这步——实测显示,相比固定384×384裁剪,该策略在 O-HAZE 测试集上 PSNR 提升 1.3dB。

2.3 损失函数设计:L1 + Perceptual + Prior Consistency 三重约束

去雾不是单纯像素重建,更要保证纹理真实、边缘锐利、雾浓度过渡自然。单一 L1 loss 会导致结果发灰、细节模糊。本项目采用三重损失:

  • L1 Loss:基础像素级重建误差,权重 λ₁=1.0;
  • VGG Perceptual Loss:用 VGG16 第 3 个 conv 层特征(relu3_3)计算,捕捉高层语义结构,权重 λ₂=0.1;
  • Prior Consistency Loss:强制 Prior Token 输出的雾图与物理雾模型(Atmospheric Scattering Model)一致,即J(x) = I(x) - t(x) * A / t(x),其中t(x)由 Prior Token 输出,A为大气光值(从雾图顶部 5% 区域估计),权重 λ₃=0.5。
# losses/losses.py class PriorConsistencyLoss(nn.Module): def __init__(self, eps=1e-6): super().__init__() self.eps = eps def forward(self, prior_map, hazy_img, dehazed_img): # prior_map: (B, 1, H, W), 值域 [0,1],越接近 1 表示雾越浓 # 根据物理模型:I = J * t + A * (1 - t) → J = (I - A*(1-t)) / t # 这里用 prior_map 作为 t(x),A 从 hazy_img 顶部区域估计 A = torch.mean(hazy_img[:, :, :int(hazy_img.size(2)*0.05), :], dim=(2,3), keepdim=True) # (B,3,1,1) t = torch.clamp(prior_map, self.eps, 1.0) # 防止除零 J_est = (hazy_img - A * (1 - t)) / t # 重建图 return F.l1_loss(J_est, dehazed_img, reduction='mean') # train.py 中 loss 组合 total_loss = l1_loss(dehazed, gt) + \ 0.1 * perceptual_loss(dehazed, gt) + \ 0.5 * prior_consistency_loss(prior_map, hazy, dehazed)

提示:PriorConsistencyLoss不是辅助 loss,而是主监督信号——它让 ViT 的 class token 真正学会“什么是雾”,而非仅拟合 GT 图。关闭此项,模型在 RESIDE-Outdoor 测试时会出现大面积过增强(天空发白、云层消失)。


3. 训练全流程:从环境配置到收敛监控,一个命令跑通

3.1 Python 环境与依赖安装:避开 OpenCV 与 PyTorch 的 CUDA 版本玄学

本项目要求:Python ≥3.8,PyTorch ≥1.12(CUDA 11.3+),torchvision ≥0.13。常见翻车点:

  • pip install opencv-python默认装 CPU 版,但cv2.cuda在去雾预处理中加速 3.2×;
  • torch==1.12.1+cu113torchvision==0.13.1+cu113必须严格匹配,否则 DataLoader 多进程崩溃。

推荐安装命令(Linux / Windows WSL):

# 创建干净环境 conda create -n vit-dehaze python=3.9 conda activate vit-dehaze # 优先装 CUDA 版 PyTorch(以 11.3 为例) pip install torch==1.12.1+cu113 torchvision==0.13.1+cu113 --extra-index-url https://download.pytorch.org/whl/cu113 # 再装带 CUDA 支持的 OpenCV pip install opencv-python-headless==4.7.0.72 pip install opencv-contrib-python-headless==4.7.0.72 # 其他依赖 pip install numpy==1.23.5 tqdm==4.64.1 scikit-image==0.19.3

注意:opencv-python-headless是关键——它不含 GUI 模块,避免与系统 Qt 库冲突,且cv2.cuda在 headless 模式下仍可用。若装opencv-python,在 Docker 或无桌面环境中会因找不到libglib-2.0.so.0报错。

3.2 启动训练:config.yaml 控制所有超参,不改代码也能调优

项目根目录下config.yaml定义全部可调参数,无需修改.py文件:

# config.yaml 片段 train: batch_size: 8 num_workers: 4 epochs: 120 lr: 2e-4 weight_decay: 1e-5 scheduler: "cosine" # 支持 step / cosine / reduce_lr_on_plateau warmup_epochs: 5 model: img_size: 512 patch_size: 4 embed_dim: 128 depth: 8 num_heads: 4 mlp_ratio: 4.0 data: train_dir: "./data/RESIDE/ITS_train" val_dir: "./data/RESIDE/SOTS_outdoor" crop_type: "fog_aware" # 可选: "random", "center", "fog_aware"

启动命令(单卡):

python train.py --config config.yaml --log_dir ./logs/vit_dehaze_base

启动命令(多卡 DDP):

torchrun --nproc_per_node=2 train.py --config config.yaml --log_dir ./logs/vit_dehaze_ddp

提示:--log_dir指定日志路径,TensorBoard 自动记录 loss 曲线、PSNR/SSIM、prior_map 可视化。训练第 30 epoch 后,prior_map 应呈现清晰的雾浓度分层(如远处山体高亮、近处建筑暗淡),这是模型真正学会雾建模的标志。

3.3 验证与推理:用 eval.py 测 PSNR/SSIM,用 infer.py 一键去雾

训练完成后,用eval.py在标准测试集上打分:

python eval.py --config config.yaml \ --ckpt_path ./logs/vit_dehaze_base/best.pth \ --test_dir ./data/RESIDE/SOTS_indoor \ --save_dir ./results/sots_indoor

输出自动写入./results/sots_indoor/metrics.txt,含 PSNR、SSIM、LPIPS 三项指标。
推理单张图(支持 JPG/PNG):

python infer.py --ckpt_path ./logs/vit_dehaze_base/best.pth \ --input ./demo/foggy_city.jpg \ --output ./demo/dehazed_city.jpg \ --img_size 512

注意:infer.py内置自适应 padding——若输入非 512×512,先 pad 到 512 倍数,推理后再 crop 回原尺寸,避免边缘伪影。实测 512×512 输入在 RTX 3090 上单帧耗时 83ms(含数据加载),满足实时视频处理需求。


4. 避坑指南:ViT 去雾训练中 5 个血泪经验换来的真问题

4.1 现象:训练初期 loss 爆炸(>1000),梯度 norm > 1000

原因:Prior Consistency Loss 中t = prior_map未做 clamp,当 prior_map 输出接近 0 时,(I - A*(1-t))/t导致数值溢出。
解决:在PriorConsistencyLoss.forward()中强制t = torch.clamp(prior_map, 1e-6, 1.0),并在model.forward()中对 prior_map 加 sigmoid 激活,确保输出 ∈ (0,1)。

4.2 现象:验证 PSNR 停滞在 28.5dB,不再上升

原因:RESIDE-Indoor 训练集 GT 图部分含 JPEG 压缩伪影,与雾图不严格配对;模型学到“压缩噪声”而非去雾。
解决:在dataset.py__getitem__中,对 GT 图做cv2.GaussianBlur(gt, (3,3), 0)模糊处理,匹配雾图的模糊程度。实测提升最终 PSNR 0.9dB。

4.3 现象:推理结果出现彩色条纹(尤其天空区域)

原因:ViT 的 Position Embedding 使用绝对位置编码,在推理时输入尺寸与训练不一致(如训练 512,推理 1920×1080),导致位置偏移。
解决:改用Rotary Position Embedding (RoPE)2D Relative Position Bias,本项目采用后者,在models/vit_dehaze.pyAttention模块内加入relative_position_bias_table,支持任意尺寸输入。

4.4 现象:多卡训练时 GPU 显存占用不均衡(0卡占 10GB,1卡占 4GB)

原因:DataLoader 的num_workers设置过高(>4),导致子进程内存泄漏;且pin_memory=True时,CPU 内存未及时释放。
解决num_workers设为min(4, os.cpu_count()),并在train.pyDataLoader初始化中添加persistent_workers=True,配合prefetch_factor=2

4.5 现象:导出 ONNX 后推理结果全黑

原因:Prior Token 的 reshape 操作view(-1, 1, H, W)在动态 batch size 下,ONNX 不支持-1推断;且torch.nn.functional.interpolatemode='bilinear'在 ONNX 中需指定align_corners=True
解决:在export_onnx.py中,用torch.onnx.export(..., dynamic_axes={'input': {0: 'batch'}})声明动态轴,并将 reshape 改为prior_map.reshape(batch_size, 1, H, W),插值操作显式传入align_corners=True


5. 进阶技巧:如何把 ViT 去雾模型压到 12MB 以内,部署到 Jetson Nano

5.1 模型瘦身三板斧:剪枝 + 量化 + 算子融合

ViT-Dehaze Base 模型(depth=8, embed_dim=128)原始大小 86MB,无法部署到边缘设备。我们通过三步压缩:

  1. 结构化剪枝(Structured Pruning):按 channel 剪掉 Attention 中 value projection 和 FFN 中第一个 linear 的冗余通道,依据weight.norm(dim=1)排序,剪 30%;
  2. INT8 量化(Post-Training Quantization):用 PyTorch 的torch.quantization,校准数据用 RESIDE-Indoor 验证集前 100 张图,qconfig = get_default_qconfig('fbgemm')
  3. 算子融合(Operator Fusion):将LayerNorm + Linear + GELU融合为单个FusedLayerNormGELU,减少 kernel launch 开销。
# tools/prune_quantize.py def prune_model(model, ratio=0.3): for name, module in model.named_modules(): if isinstance(module, nn.Linear) and 'v_proj' in name or 'mlp.fc1' in name: # 基于 channel norm 剪枝 weight_norm = module.weight.data.norm(dim=1) _, idx = torch.topk(weight_norm, int(module.out_features * (1-ratio)), largest=False) mask = torch.zeros(module.out_features, dtype=torch.bool) mask[idx] = True module.weight.data = module.weight.data[mask] module.out_features = mask.sum().item() def quantize_model(model, calib_loader): model.eval() model.fuse_model() # 融合 LayerNorm/GELU model.qconfig = torch.quantization.get_default_qconfig('fbgemm') torch.quantization.prepare(model, inplace=True) with torch.no_grad(): for data in calib_loader: model(data) torch.quantization.convert(model, inplace=True)

提示:剪枝后需微调(fine-tune)5 个 epoch,否则 PSNR 下降 >2dB;量化后务必用torch.jit.trace导出 TorchScript,再转 ONNX,避免 PyTorch 量化算子兼容性问题。

5.2 Jetson Nano 部署实测:1280×720 视频流 12FPS,功耗 5.2W

压缩后模型(11.8MB)在 Jetson Nano(JetPack 4.6, CUDA 10.2)上实测:

输入尺寸FPS显存占用功耗
640×360241.1GB4.3W
1280×720122.4GB5.2W
1920×108053.8GB5.8W

部署命令(TensorRT 加速):

# 1. 将 ONNX 转 TensorRT engine trtexec --onnx=vit_dehaze_int8.onnx \ --int8 \ --calib=calibration.cache \ --workspace=2048 \ --saveEngine=vit_dehaze.trt # 2. Python 推理(使用 pycuda) import tensorrt as trt engine = trt.Runtime(trt.Logger()).deserialize_cuda_engine(open("vit_dehaze.trt", "rb").read()) context = engine.create_execution_context() # ... 绑定 input/output buffer,执行推理

我的习惯是:每次新硬件部署前,先用nvidia-smi dmon -s um监控 GPU utilization 和 memory bandwidth,确认不是显存带宽瓶颈(Nano 的 12.8GB/s 带宽是主要限制)。如果 FPS 上不去,优先降低img_size而非 batch size——ViT 的计算复杂度是 O(N²),N 减半,计算量降 75%。希望帮到你。

本文还有配套的精品资源,点击获取

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询