news 2026/10/5 10:51:35

Vision Transformer图像去雾实战:Patch尺寸与Decoder设计关键

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
Vision Transformer图像去雾实战:Patch尺寸与Decoder设计关键

简介:本资源是一套基于Vision Transformer(ViT)的图像去雾算法完整实现方案,面向计算机视觉方向的研究生、算法工程师及深度学习实践者,聚焦于恶劣天气下图像质量退化问题的端到端建模与复现。压缩包共340个文件,涵盖204个Python核心训练/推理脚本(含option.py参数配置、模型定义与数据加载模块)、39张效果对比图与可视化结果(png/gif)、16个YAML配置文件(支持不同数据集与模型变体)、9个Jupyter Notebook实验记录及8份Markdown使用说明,整体体积156.34MB,结构清晰、开箱即用。已有470人学习下载,提供从环境配置、预训练权重加载(支持My_best_model路径自定义)、补丁尺寸(--train_ps=128)等关键参数调优到Loss landscape分析(含cifar100_vit_ti等多组CSV实验数据)的全流程支撑,特别适合开展ViT在低层视觉任务中的迁移研究与工程落地验证。

1. Vision Transformer 做图像去雾,真不是套个 ViT 头就完事:它能实测提升 PSNR 2.3dB,但 patch 尺寸设错直接让模型学成“雾里看花”

你手头有一张被浓雾笼罩的监控截图,想恢复出车牌号——传统暗通道先验(DCP)跑出来边缘发虚、颜色偏青,DehazeNet 的 CNN 特征容易过平滑,而这份基于 Vision Transformer 的去雾源码,用的是真正的 ViT 架构(不是 ViT backbone + CNN decoder 的缝合怪),在自建雾图数据集上实测 PSNR 较 ResNet-50 基线高 2.3dB,SSIM 提升 0.041。它不依赖物理模型,靠全局注意力建模雾浓度空间分布,尤其擅长处理远距离大雾区域的纹理重建。适合两类人:一是做低光照/恶劣天气图像增强的算法工程师,需要可复现、可微调的端到端去雾 baseline;二是计算机视觉方向研究生,想拿 Vision Transformer 做图像复原类课题,但苦于找不到真正用 ViT 做 encoder-decoder 全架构设计的开源实现。注意:这不是一个 pip install 就能跑的玩具项目,它要求你理解 patch embedding 的尺寸约束、位置编码与雾图分辨率的耦合关系,以及 loss landscape 文件(如cifar100_vit_ti_losslandscape.csv)背后隐藏的训练稳定性线索——这些文件不是冗余,而是作者调试时记录的梯度曲率变化,是判断模型是否陷入局部极小的“黑匣子日志”。


2. 从源码结构到核心模块:为什么这个 ViT 去雾模型不用 CNN 做 decoder?

2.1 源码包解压后的真实目录结构与关键文件定位

解压python源码+使用说明.zip后,你会看到如下主干结构:

├── My_best_model/ # 预训练权重存放目录(含多个 .pth 文件,按数据集划分) ├── datasets/ # 数据集加载逻辑,支持自定义雾图路径 │ ├── __init__.py │ └── dehaze_dataset.py # 核心 Dataset 类,支持 paired/unpaired 模式 ├── models/ # 模型定义 │ ├── __init__.py │ ├── vit_dehaze.py # 主模型:ViT encoder + transformer-based decoder │ └── blocks.py # 自定义 Attention Block(含雾感知门控机制) ├── option.py # 全局参数配置(训练/测试/数据路径全在这里) ├── train.py # 训练入口,含 loss 定义(L1 + perceptual + edge-aware) ├── test.py # 推理脚本,支持单图/批量处理 └── utils/ # 工具函数:patch 拆分、雾浓度估计、PSNR/SSIM 计算

提示:cifar100_vit_ti_losslandscape.csv等文件并非训练必需,而是作者在不同超参组合下记录的 loss 曲线采样点(横轴为 step,纵轴为 loss 值),用于分析优化过程是否震荡、收敛是否平滑。它们的存在说明该项目经过了系统性调参,不是随手训出来的。

2.2 模型架构本质:ViT encoder + cross-attention decoder,不是“ViT + U-Net”

该模型的 decoder 并非简单堆叠卷积上采样层,而是采用cross-attention based decoder:encoder 输出的 token 序列(shape:[B, N, C])作为 key/value,decoder 自身 learnable query(shape:[B, M, C])通过 cross-attention 聚焦于 encoder 的全局上下文。这种设计让 decoder 能显式建模“哪里雾重、哪里需强重建”,比 CNN decoder 更适应雾浓度空间异质性。

关键代码片段(models/vit_dehaze.py):

# ViT encoder 输出:x_enc shape = [B, N, C] x_enc = self.vit_encoder(x) # N = (H//patch_size) * (W//patch_size) # Decoder query 初始化(learnable positional embedding) query_pos = self.query_embed.weight.unsqueeze(0) # [1, M, C] query = self.query_feat.weight.unsqueeze(0) # [1, M, C] # Cross-attention:query 对 encoder tokens 做 attention attn_out = self.cross_attn(query, x_enc, x_enc) # [B, M, C] # 后续 MLP head 生成去雾图(reshape 回 H, W) out = self.mlp_head(attn_out) # [B, M, 3*patch_size**2] out = rearrange(out, 'b m (c p1 p2) -> b c (m1 p1) (m2 p2)', p1=self.patch_size, p2=self.patch_size, m1=self.img_h//self.patch_size, m2=self.img_w//self.patch_size)

参数说明:

  • patch_size:决定 encoder 输入 patch 大小,直接影响N(token 数量)和M(decoder query 数量);
  • img_h,img_w:必须能被patch_size整除,否则rearrange会报错;
  • query_embed和query_feat是 learnable 参数,不是固定位置编码,允许 decoder 动态学习关注重点。

2.3 数据加载逻辑:支持真实雾图 + 合成雾图混合训练

datasets/dehaze_dataset.py中的__getitem__方法做了三件事:

  1. 双路径加载:若opt.unpaired == False(默认),加载 clean 图 + 对应合成雾图(paired);若True,则从 clean 雾图池随机采样配对(unpaired);
  2. 雾浓度自适应裁剪:根据输入图雾浓度估计值(用utils/fog_estimation.py的快速方差法),动态调整train_ps(训练 patch 大小),浓雾区域优先取大 patch;
  3. 在线雾化增强:对 clean 图用utils/synthetic_fog.py实时添加 multi-scale 雾层(非简单高斯模糊),模拟真实雾散射特性。

注意:option.py中--train_ps 128是默认值,但实际训练中建议根据你的 GPU 显存和输入图分辨率动态调整。例如:输入图 1024×768,patch_size=16→N=4800tokens,显存占用约 12GB(V100);若设patch_size=32,N=1200,显存降至 5GB,但可能丢失细粒度雾结构。


3. 训练全流程实操:从环境配置到权重加载,每一步都踩过坑

3.1 环境依赖与 Python 版本硬性要求

该项目基于 PyTorch 1.12 + TorchVision 0.13 构建,不兼容 PyTorch 2.x(因torch.nn.MultiheadAttention在 2.0+ 中默认启用enable_nested_tensor=False,而本项目依赖 nested tensor 的自动 padding)。Python 版本必须为3.8 或 3.9(3.10+ 会导致einops的rearrange在某些 GPU 上报CUDA error: device-side assert triggered)。

安装命令(严格按顺序):

# 创建干净环境 conda create -n vit-dehaze python=3.8 conda activate vit-dehaze # 安装指定版本 PyTorch(CUDA 11.3) pip install torch==1.12.1+cu113 torchvision==0.13.1+cu113 torchaudio==0.12.1 --extra-index-url https://download.pytorch.org/whl/cu113 # 安装其余依赖(requirements.txt 中缺失 einops,必须手动补) pip install numpy opencv-python scikit-image matplotlib tqdm einops timm

提示:timm库用于加载预训练 ViT 权重(如vit_tiny_patch16_224),但本项目实际使用的是自定义 ViT encoder,timm仅作参考实现,非运行必需。

3.2 修改option.py的 5 个关键参数

打开option.py,以下参数必须按你的硬件和数据修改:

参数名默认值必改原因推荐值(RTX 3090)
--train_ps128决定 patch 大小,影响显存和感受野96(平衡显存与细节)
--batch_size8与train_ps强耦合若train_ps=96,设为 12
--pretrain_weights"My_best_model/vit_ti_cifar100.pth"指向你自己的预训练权重路径"My_best_model/my_fog_pretrain.pth"
--data_dir"./datasets/RESIDE/"数据集根目录,需包含train/clean和train/hazy子目录"./my_fog_data/"
--save_dir"./checkpoints/"权重保存路径,确保有写入权限"./checkpoints_myrun/"

修改后验证方式:运行python train.py --help,确认参数已生效。

3.3 启动训练:一条命令背后的三个隐式检查

执行训练命令:

python train.py --train_ps 96 --batch_size 12 --pretrain_weights "My_best_model/vit_ti_cifar100.pth" --data_dir "./my_fog_data/"

该命令启动前会自动执行三个检查:

  1. Patch 尺寸校验:检查data_dir下任意一张图的宽高是否能被train_ps整除,不能则报错Image size must be divisible by patch_size;
  2. 权重兼容性检查:加载pretrain_weights后,比对state_dict的 key 名与模型定义是否一致,若 encoder 层名不匹配(如blocks.0.attn.qkv.weightvsblocks.0.attn.wqkv.weight),报错Key mismatch in pretrained weights;
  3. Loss landscape 文件关联:若option.py中--use_loss_landscape True(默认 False),则尝试读取cifar100_vit_ti_losslandscape.csv作为 early stopping 的参考曲线,文件不存在则跳过。

4. 避坑指南:这 4 个血泪问题让我重训了 7 次

4.1 现象:训练 loss 初期剧烈震荡(±5.0),100 epoch 后仍 > 0.8

原因:option.py中--lr 2e-4对 ViT encoder 过大,导致 attention weight 更新失稳;同时--weight_decay 1e-4未对 decoder query 参数单独设置,造成 query embed 过拟合。
解决:在train.py的 optimizer 构建处,为 decoder query 添加独立学习率:

optimizer = torch.optim.AdamW([ {'params': model.encoder.parameters(), 'lr': 1e-4}, {'params': model.decoder.query_embed.parameters(), 'lr': 5e-5}, # 单独调低 {'params': model.decoder.query_feat.parameters(), 'lr': 5e-5}, {'params': model.decoder.cross_attn.parameters(), 'lr': 2e-4}, ], weight_decay=1e-4)

4.2 现象:推理结果全图泛白,PSNR 反而比输入雾图低

原因:test.py中--save_images True时,输出图未做torch.clamp(0, 1)截断,且utils/postprocess.py的denormalize函数误用了 ImageNet 均值标准差([0.485,0.456,0.406]),而本项目训练用的是[-1,1]归一化。
解决:修改utils/postprocess.py第 23 行:

# 错误写法(用 ImageNet stats) # img = img * torch.tensor([0.229, 0.224, 0.225]) + torch.tensor([0.485, 0.456, 0.406]) # 正确写法(本项目用 [-1,1],需转回 [0,1]) img = torch.clamp((img + 1) / 2, 0, 1) # 直接反归一化

4.3 现象:cifar100_vit_ti_losslandscape.csv读取失败,报UnicodeDecodeError: 'utf-8' codec can't decode byte 0xff

原因:该 CSV 文件实际是二进制保存的 numpy array(.npy伪装成.csv),作者用np.savetxt时指定了fmt='%f'但未加encoding='utf-8',Windows 系统默认用 GBK 编码打开。
解决:用 numpy 直接加载:

import numpy as np loss_landscape = np.loadtxt('cifar100_vit_ti_losslandscape.csv', delimiter=',') # 会失败 # 改为: loss_landscape = np.load('cifar100_vit_ti_losslandscape.csv'.replace('.csv', '.npy')) # 实际文件名是 .npy

玄学提示:作者把.npy文件后缀硬改成.csv,是为了让 GitHub 直接预览(CSV 可视化),但实际内容是二进制。解压后检查文件大小:真正的 CSV 应 >1MB,若只有 12KB,大概率是.npy。

4.4 现象:多卡训练时报错RuntimeError: Expected all tensors to be on the same device

原因:models/blocks.py中LayerNorm层的weight和bias参数未随 model 移动到 GPU,因其在__init__中用nn.Parameter(torch.zeros(...))初始化,但未显式.to(device)。
解决:在blocks.py的__init__末尾添加:

self.norm1.weight.data = self.norm1.weight.data.to(device) self.norm1.bias.data = self.norm1.bias.data.to(device) # 同理处理 norm2

或更规范的做法:在train.py的model.to(device)后,加一行model = torch.nn.DataParallel(model)(单机多卡)。


5. 进阶技巧:用 loss landscape 文件诊断过拟合,并定制你的雾浓度敏感 decoder

5.1 解析cifar100_vit_ti_losslandscape.csv:它不只是曲线图

该文件实际是 3D loss surface 的二维切片采样,共三列:step,lr,loss。其中lr列并非学习率,而是loss landscape 的横坐标扰动强度(即在当前权重附近加噪声ε ~ N(0, lr)后的 loss 值)。作者用此评估模型鲁棒性:若lr=0.01时 loss 波动 <0.05,说明模型处于平坦极小值区;若lr=0.001时 loss 已飙升,说明过拟合。

解析脚本(analyze_landscape.py):

import numpy as np import matplotlib.pyplot as plt # 加载真实 .npy 文件(别信 .csv 后缀) data = np.load('cifar100_vit_ti_losslandscape.npy') # shape: [N, 3] steps, epsilons, losses = data[:, 0], data[:, 1], data[:, 2] # 按 epsilon 分组,计算每个扰动强度下的 loss std eps_unique = np.unique(epsilons) std_per_eps = [losses[epsilons == e].std() for e in eps_unique] plt.plot(eps_unique, std_per_eps, 'o-') plt.xlabel('Perturbation Strength (ε)') plt.ylabel('Loss Std') plt.title('Loss Landscape Flatness') plt.grid(True) plt.show()

解读:若曲线呈“U型”(小 ε 和大 ε 时 std 都高),说明模型处于尖锐极小值,易过拟合;若整体平缓(std <0.02),则权重泛化性强。我实测发现,当--train_ps从 128 降到 64 时,std_per_eps在 ε=0.005 处从 0.018 升至 0.042,证实小 patch 加剧了 sharpness。

5.2 定制雾浓度感知 decoder:插入 fog-gating module

原始 decoder 对所有 query 一视同仁,但实际雾图中,天空区域雾浓度高、纹理少,道路区域雾浓度低、边缘多。我们可在 cross-attention 后插入 fog-gating:

# 在 vit_dehaze.py 的 decoder forward 中 # attn_out shape: [B, M, C] fog_map = self.fog_estimator(attn_out) # [B, M, 1], sigmoid 输出雾浓度 [0,1] gated_out = attn_out * fog_map # 强制模型在高雾区降低重建强度 out = self.mlp_head(gated_out)

fog_estimator实现(轻量级):

class FogGating(nn.Module): def __init__(self, dim): super().__init__() self.proj = nn.Sequential( nn.Linear(dim, dim//4), nn.GELU(), nn.Linear(dim//4, 1), nn.Sigmoid() ) def forward(self, x): # x: [B, M, C] return self.proj(x) # [B, M, 1]

效果:在 RESIDE-SOTS 测试集上,PSNR 提升 0.4dB,且天空区域伪影减少 37%(人工评测)。

5.3 一个硬核习惯:每次修改train_ps,必重跑patch_size_validator.py

我写了个校验脚本,放在utils/patch_size_validator.py:

import os from PIL import Image import argparse def validate_patch_size(data_dir, patch_size): for split in ['train', 'val']: img_dir = os.path.join(data_dir, split, 'hazy') for img_name in os.listdir(img_dir)[:10]: # 只查前10张 img = Image.open(os.path.join(img_dir, img_name)) if img.width % patch_size != 0 or img.height % patch_size != 0: print(f"❌ {img_name}: {img.size} not divisible by {patch_size}") return False print(f"✅ All images divisible by {patch_size}") return True if __name__ == "__main__": parser = argparse.ArgumentParser() parser.add_argument('--data_dir', type=str, required=True) parser.add_argument('--patch_size', type=int, required=True) args = parser.parse_args() validate_patch_size(args.data_dir, args.patch_size)

运行:python utils/patch_size_validator.py --data_dir "./my_fog_data/" --patch_size 96
从那以后我每次改train_ps,都强制走一遍这个脚本——它避免了 90% 的size mismatch报错,省下至少 3 小时重训时间。希望帮到你。

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

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/10/5 10:50:04

Delphi可视化开发中界面布局的设计思路与优化

在Delphi可视化开发中&#xff0c;界面布局直接决定了用户体验的好坏。一个清晰、美观、响应迅速的界面&#xff0c;不仅能提升软件的易用性&#xff0c;更能体现开发者的专业性。本文将深入探讨Delphi中界面布局的核心设计思路与优化技巧&#xff0c;旨在帮助开发者构建出更优…

作者头像 李华
网站建设 2026/10/5 10:49:46

虚拟机密码恢复与重置全指南:Linux与Windows平台实操

开头直接进入场景&#xff1a;我遇到过太多次这样的求助——“虚拟机密码忘了怎么办”。在这个问题上&#xff0c;搜索框里的热门词同样是“虚拟机怎么安装”“vmware虚拟机安装教程”这类入门操作&#xff0c;密码恢复反而成了没人系统讲过的东西。整篇文章聊的是干净、可落地…

作者头像 李华
网站建设 2026/10/5 10:48:47

分布式锁进阶:Redisson MultiLock 联锁原理与多实例实战

做分布式系统绕不开分布式锁&#xff0c;这篇文章我拿实际项目里的 Redisson MutiLock&#xff08;联锁&#xff09;来说事&#xff1a;它到底解决了什么问题、加锁解锁的过程是怎么设计的、基于什么原理&#xff0c;以及我如何在一个只有 Windows 的测试环境里&#xff0c;硬生…

作者头像 李华
网站建设 2026/10/5 10:48:36

Python卷积神经网络疲劳检测毕设包:从环境配置到预警系统实战

简介&#xff1a;这份毕业设计资源面向计算机相关专业学生与Python初学者&#xff0c;提供一套基于卷积神经网络的人脸识别驾驶员疲劳检测与预警系统完整源码&#xff0c;可用于课程设计、毕设答辩或深度学习入门实践。压缩包共15个文件&#xff0c;约2.8MB&#xff0c;包含3个…

作者头像 李华
网站建设 2026/10/5 10:48:25

网络安全工程师薪资与就业前景:入行路线、岗位路径与避坑指南

“网络安全就业前景怎么样&#xff1f;网络安全工程师多少钱一个月&#xff1f;”这个问题我几乎每周都会被问到一次&#xff0c;聊的人从刚毕业的应届生到大厂被优化的前端都有。我在这个行业从搞渗透测试起家&#xff0c;到后来带安全团队、参与过不少企业的安全体系建设&…

作者头像 李华
网站建设 2026/10/5 10:48:07

视频质量诊断与GB28181平台融合:EasyVQD+EasyGBS打造运维闭环

1. 监控运维里的"隐形杀手"&#xff1a;画质故障为什么总是最后被发现 干监控运维这行超过十年的人&#xff0c;基本都有过这种经历&#xff1a;某个客户的录像调出来了&#xff0c;结果一看画面&#xff0c;花屏花了一星期&#xff0c;蓝屏蓝了两天&#xff0c;偏偏…

作者头像 李华