尧图精选

基于Vision Transformer的图像去雾:全局自注意力与工程实践详解

🕒 发布时间:2026/10/2 14:03:58 📁 来源:尧图网络
简介面向深度学习研究与计算机视觉开发本资源提供基于Vision Transformer的图像去雾算法Python实现包含完整源码、预训练权重与使用说明。适合具备模型训练基础的学生、算法工程师和科研人员可用于去雾实验复现、网络结构改进或迁移至图像恢复相关任务。压缩包共340个文件整体约156MB核心代码以204个Python脚本为主辅以16个YAML配置文件、9个Jupyter Notebook和8个Markdown文档覆盖模型构建、参数配置、训练评估与结果展示。另有10个GIF效果演示、12个CSV实验记录等其中CSV数据包含CIFAR-10/100上多种网络模型的损失景观与鲁棒性测试结果便于分析模型性能。训练参数统一集中在option.py中--train_ps可设置输入patch大小默认128预训练权重按不同数据集划分存放方便针对性微调。已有1444人学习适合希望系统掌握Transformer图像去雾实现细节并在此基础上开展深入研究的读者。1. 基于Vision Transformer的图像去雾先把任务拆给自注意力拍摄场景受到雾、霾影响时图像对比度和色彩保真度同步下降。传统暗通道先验对天空这类明亮区域常会失效CNN模型通过大量参数拟合退化过程却又容易被有限的局部感受野限制无法感知整幅图的雾浓度分布。Vision TransformerViT通过自注意力把整张图拉进同一个计算图patch化的做法又压缩了空间冗余这让图像去雾从“猜透射率”变成真正的全局回归任务也因此成为图像去雾算法研究里的主流选择。标题对应的工程代码是Python实现覆盖数据生成、模型训练、指标评估与可视化对比再附上项目介绍和使用说明适合图像复原方向的学生与算法工程师作为起点复现。2. 雾天成像模型与Vision Transformer的适配逻辑你先别急着搭网络。如果不清楚去雾模型为什么非要对整幅图做回归后面调参只会变成玄学。图像去雾和超分、去噪的差别在于退化是全局性的任何一个像素的清晰值都依赖大气光和该点的透射率。这两样东西分别描述全局光照和空间传输恰好一长一短ViT的结构在这里天然就有优势。2.1 大气散射模型两个未知量的病态反问题去雾常用的前提是大气散射模型形式为I(x) J(x) * t(x) A * (1 - t(x))I 是观测到的雾图J 是真实无雾图t 表示透射率A 是全局大气光。t 随距离指数衰减t exp(-βd)β 是散射系数。去雾问题就是已知 I同时求 J、t、A这是病态的一个方程里两个未知量如果不加约束解的数量是无穷的。雾的物理性质能提供一些约束相近深度区域的 t 变化平滑离镜头越远雾越浓。基于这些约束暗通道先验只能吃掉一部分场景遇到大片天空、白色墙面和光源区域先验假设直接崩了这也是暗通道方法在近白区域常出现色偏的原因。用深度学习做这件事时比较稳妥的路线是直接回归清晰图 J避开显式求 t 和 A辅助监督透射率图 t让网络学到的中间表示符合成像模型混合路线先估计 t 和 A再用大气散射模型重建 J最后对 J 做轻量精修。我在自己的试验里最常用的是第一条加第二条的组合直接回归清晰图作为主输出训练时加一个低权重的透射率图监督。这样模型在推理阶段只需要一个前向过程不需要复杂的后处理同时中间监督让梯度更稳定不会一路黑盒到底。这一节不放代码因为建模逻辑主要在数据合成和训练部分见后面的 4.1。2.2 Patch Embedding 和自注意力全局感受野如何匹配全局退化CNN 的感受野靠堆叠卷积来扩大底层卷积看到的是 3x3、7x7 这种局部窗口。对于雾这种大面积慢变化信号深层网络必须堆到很深才能“看到”整幅图而且底层局部操作很容易把边缘纹理和雾的梯度混在一起处理。ViT 的做法不同。首先把图像切成固定大小的小块每个 patch 拉平后经过线性投影映射成 tokentoken 之间用自注意力做交互。关键代码如下import torch import torch.nn as nn # 把一个批次的 RGB 图切成固定 patch并做线性映射成 token class PatchEmbed(nn.Module): def __init__(self, in_chans3, embed_dim768, patch_size16): super().__init__() self.patch_size patch_size # 卷积核等于 patch 大小步长也等于 patch 大小天然完成切块投影 self.proj nn.Conv2d(in_chans, embed_dim, kernel_sizepatch_size, stridepatch_size) def forward(self, x): # x: (B, 3, H, W) B, C, H, W x.shape x self.proj(x) # (B, embed_dim, H/patch_size, W/patch_size) x x.flatten(2) # (B, embed_dim, N) x x.transpose(1, 2) # (B, N, embed_dim) return x参数说明patch_size16 时一张 256x256 的输入被压成 16x16256 个 token序列长度大幅压缩后续自注意力的计算量也合适。embed_dim 是每个 token 的特征维度常见取值 384 或 768。小数据集或者显存受限时用 384 起步更稳。token 本身不带位置信息训练时还需要加可学习的位置编码否则模型不知道 patch 之间的空间顺序。这段代码在初始化 ViT 的输入处理时用到逻辑就是“切块 线性映射”没有多余算子。对图像去雾任务来说patch 后空间分辨率下降自注意力再在这些 token 间做全连接每个 token 的计算都会顾及全图。雾的区域分布、大气光的整体亮度这些东西天然属于全局信息和自注意力的匹配度比单纯卷积高得多。2.3 全局注意力的显存代价窗口注意力是妥协也是捷径直接做全局自注意力序列长度为 N 时复杂度是 O(N^2)。256x256 输入、patch16 时 N256尚可接受到了高清图像N 超过 1024 后显存压力明显。常见做法是采用 Swin Transformer 的移位窗口注意力先在一组窗口内做局部自注意力再通过窗口移位让相邻窗口交换信息。我在小分辨率训练时直接用全局自注意力推理分辨率提到 2K 以上就切窗口注意力或者干脆在低分辨率推理后再做一次轻量上采样。这里需要强调去雾最终效果并不完全由全局注意力决定窗口设置合理、多层特征融合到位效果差距通常不大但显存差距可以差出一个数量级。3. 实现 ViT 去雾网络主干结构选型与损失函数调参3.1 主干选型固定 ViT 还是 U 型多尺度如果直接把原版 ViT 拿来做去雾感受野有了但分辨率损失过大输出细节出不来。原因是去雾需要像素级输出而 ViT 的 patch 嵌入天然把空间分辨率缩减到 1/16直接上采样回来的边缘都很糊。常见做法是参照 U-Net 结构设计编码器-解码器内部用 Transformer 作为特征提取单元编码器第一层保持高分辨率第二层起用 patch 或卷积做降采样感受野逐层变大瓶颈层用完整的 ViT 层或者 Swin Block 提炼全局依赖解码器逐层上采样并跳连拼接编码器的同尺度特征恢复细节。这种结构下ViT 的全局建模主要承担“看懂全局光照”的任务卷积跳连承担边缘保真。我在项目里常用的实现是输入经过一条 3x3 卷积极联的 CNN 骨干生成多尺度特征最深一层用几个 Swin Block 做全局交互再上采样逐级合并。选择这种混合结构而不是纯 ViT原因有二。第一纯 ViT 在有限数据集下收敛慢预训练也不一定能在小数据集上发挥优势。第二去雾任务对边缘纹理的保留要求高CNN 跳跃连接比单纯 attention 更适合像素级细节。3.2 损失函数三项组合怎么定权重去雾损失最常用的是这三项加权求和L λ1·L1 λ2·L_SSIM λ3·L_perceptualL1 收敛快但容易模糊SSIM 保结构保对比感知损失约束视觉纹理。具体实现import torch import torch.nn.functional as F from torchvision.models import vgg16 # 感知损失用到 vgg16 的特征层外部需提前加载权重 class DehazeLoss(torch.nn.Module): def __init__(self, lambda_l11.0, lambda_ssim0.2, lambda_perc0.5): super().__init__() self.lambda_l1 lambda_l1 self.lambda_ssim lambda_ssim self.lambda_perc lambda_perc self.vgg vgg16(pretrainedTrue).features[:16].eval() def forward(self, pred, target): # L1 逐像素平均主心骨损失 l1 F.l1_loss(pred, target) # 用 pytorch_msssim 库计算 SSIM这里省略导入 ssim_value ssim(pred, target) # 感知损失对比 vgg 特征图的 L1 距离 perc F.l1_loss(self.vgg(pred), self.vgg(target)) return (self.lambda_l1 * l1 self.lambda_ssim * (1 - ssim_value) self.lambda_perc * perc)参数说明lambda_l1 通常大于 lambda_ssim 和 lambda_percL1 是地基保证像素回归不掉。lambda_ssim 太高会让颜色发灰参考值区间 0.1-0.3超过 0.5 后结构权重过强亮度动态范围会被压缩。lambda_perc 需要模型先加载预训练权重。在小数据集上建议先关掉感知损失训练十几轮再打开它做 finetune否则一开始梯度方向容易被感知特征带偏输出出现奇怪的纹理重复。训练时要把三项损失分开记录。如果只记总 loss后期很难判断是 L1 饱和还是感知损失主导。下面 4.2 节的日志会给出分开打点的写法。4. 用 Python 跑通去雾训练全流程数据合成、训练配置与指标记录4.1 用大气散射模型在线合成雾图图像去雾的数据获取比超分更麻烦真实配对的清晰/雾图数据量太少。大多数人选择用大气散射模型在无雾图上合成带雾样本甚至直接在 dataloader 里在线生成每个 epoch 都在变相当于免费的数据增强。下面的函数用随机 t 和 A 生成雾图import torch def random_haze(img, t_min0.4, t_max1.0): img: (B, 3, H, W) 归一化到 [0, 1] 返回退化图 haze 和透射率 t B, C, H, W img.shape # 透射率每张图一个随机标量模拟整体雾浓度差异 t torch.rand(B, 1, 1, 1) * (t_max - t_min) t_min # 大气光随机取 0.7-1.0对应归一化空间里的天空亮度 A torch.rand(B, 3, 1, 1) * 0.3 0.7 haze img * t A * (1 - t) return haze, t逻辑说明t 是逐图一个标量模拟均匀雾浓度。真实场景透射率是渐变的可以在 H 和 W 两个维度上用线性插值或者高斯噪声生成渐变 t效果会更真实。A 的范围取 0.7-1.0这是天空亮度折算到归一化空间后的典型区间。A 取值超过 1.0 时整张图会被压到发白训练出来的模型会倾向于把暗部提得太亮。在线生成的优点在于样本无穷无尽模型见过更多样的 t/A 组合。缺点是这种 t 与真实场景尺度并不完全一致所以训练后期要加入真实雾图做验证避免模型只会处理均匀雾。4.2 训练脚本参数与注意事项从 dataloader 到日志下面这段训练循环可以直接作为起点关键配置在参数说明里import torch from torch.utils.data import DataLoader, Dataset from torchvision import transforms from torch.optim import AdamW from torch.optim.lr_scheduler import CosineAnnealingLR import random import numpy as np # 固定随机种子方便复现 def seed_everything(seed42): random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) torch.cuda.manual_seed_all(seed) seed_everything(42) # 训练增强随机裁剪到 256水平/垂直翻转标准化到 [-1, 1] transform transforms.Compose([ transforms.RandomCrop(256), transforms.RandomHorizontalFlip(0.5), transforms.RandomVerticalFlip(0.5), transforms.ToTensor(), transforms.Normalize([0.5, 0.5, 0.5], [0.5, 0.5, 0.5]), ]) ds DehazeDataset(rootdata/clear, transformtransform) loader DataLoader(ds, batch_size8, shuffleTrue, num_workers4, pin_memoryTrue, drop_lastTrue) model DehazeViT(embed_dim384, depth4, num_heads6).cuda() optimizer AdamW(model.parameters(), lr2e-4, weight_decay0.05) scheduler CosineAnnealingLR(optimizer, T_max50) for epoch in range(50): model.train() total_loss 0.0 for imgs, _ in loader: imgs imgs.cuda() haze, t_target random_haze(imgs, t_min0.4, t_max1.0) pred, t_pred model(haze) loss dehaze_loss(pred, imgs) 0.1 * F.l1_loss(t_pred, t_target) optimizer.zero_grad() loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm5.0) optimizer.step() total_loss loss.item() scheduler.step() print(fepoch {epoch} loss {total_loss / len(loader):.4f})参数说明batch_size 8、分辨率 256、embed_dim 384 是一个较稳起步配置显存 8-10G 可以跑。显存不足时先把 batch 降到 4再考虑缩小 embed_dim不要一上来就砍 patch_size。AdamW 配合 weight_decay 0.05 是 Transformer 训练的标准组合。很多去雾模型翻车不是结构问题而是没配 AdamW 和足够长的 warmup。正式训练前最好用几个 epoch 做 warmup把学习率从 0 线性拉到 2e-4。在线随机裁剪和翻转是必须的。固定分辨率下训练出来的模型对位置非常敏感不做增强时测试集换个裁剪位置PSNR 能掉 1-2dB。透射率分支的监督权重我取 0.1太高会让主任务被带偏太低起不到稳定梯度作用。判断训练有没有问题要看两个趋势L1 loss 平稳下降同时验证集 PSNR 同步上升。如果训练 loss 一直降而验证集 PSNR 不动那就不是拟合不了而是泛化问题要去查数据分布和正则化而不是继续堆训练轮数。5. 图像去雾训练与推理的五个避坑记录这一章写我复现和调参时最常遇到的 5 类问题按“现象 → 原因 → 解决”整理全是可直接落地的操作。5.1 PSNR 涨了但输出图像灰蒙蒙的现象训练几十轮后 PSNR 还不错但输出贴出来以后对比度明显低于原图整张图蒙了一层灰。原因L1 损失对全局偏差不敏感模型输出的均值接近正确但方差不足。SSIM 权重不够或者 lambda_ssim 只作用于结构而没拉动全局对比度。解决调高 SSIM 分量权重更直接的做法是在输出后面追加一个 1x1 卷积层做逐通道色彩映射初始化成恒等逼近让网络自己学全局亮度补偿。另一个常用技巧是在数据增强里按概率把输入亮度减掉一些强迫网络恢复动态范围。5.2 全局自注意力显存爆掉现象256x256 正常换到 512x512 直接 OOM。原因全局自注意力复杂度是 O(N^2)分辨率翻倍后序列长度翻 4 倍注意力矩阵增长 16 倍。解决显存不够时改用窗口注意力或者把 patch_size 从 16 改成 32。patch 变大后 token 数减少单 token 覆盖的空间区域变大对全局任务反而有增益只是细节丢失明显。另一个技巧先用小 patch 训练推理时切大 patch注意位置编码要做双线性插值否则会出现周期性的块状伪影。5.3 合成雾图效果好真实雾图一塌糊涂现象在合成测试集上 PSNR 30拿真实雾图一跑颜色偏紫或者局部区域过暗。原因合成数据里 t 是全局均匀或只做了简单渐变真实雾的透射率场和亮度分布与这个差距很大。真实雾的散射系数随波长变化所以偏色问题更严重。解决把 t 从全局标量换成和输入同分辨的随机噪声场再做高斯平滑。这样每个局部透射率不同网络被迫学习空间变化的去雾而不是只做整体亮度减淡。另一个思路是加入少量真实雾图和对应的伪目标做混合训练伪目标用预训练模型生成人工校一遍再进训练集。5.4 训练前期 loss 不降模型输出全黑现象前几个 epoch 输出接近黑色梯度大小异常loss 剧烈波动。原因学习率过高、没有 warmup或者 PatchEmbed 输出缺少 pre-norm导致 self-attention 里 softmax 饱和注意力分布退化到均匀。解决检查学习率是不是超过 1e-3AdamW 下 1e-4 到 2e-4 极少出现这个现象。还不行就在 PatchEmbed 后加 LayerNorm。也可以把最后一个线性层的 bias 初始化成输出均值对应的常数让初始输出不远离目标分布。5.5 同一份代码重跑结果不一致指标差 0.5dB现象固定了随机种子但重跑后 PSNR 仍然有波动。原因PyTorch 默认非确定性卷积实现DataLoader 的 num_workers 多进程采样也会改变数据顺序。严格说不算 bug但对研究复现很致命。解决在训练入口固定 Python、NumPy、PyTorch 三套随机种子DataLoader 设 drop_lastTrue必要时把 num_workers 设为 0 或调低。对完全可复现有要求时再加 torch.use_deterministic_algorithms(True)。6. 验证与部署让去雾模型真正落地训练完模型别只看 PSNR/SSIM 两个数字。我的做法是同时导出输入雾图、模型输出、清晰真值三张图把输出和真值的差图放大显示能一眼看出颜色偏和边缘过火的位置。PSNR 高也可能因为输出平滑人眼看着就是“糊了”。评测时按场景拆开算常规均匀雾、浓雾、真实雾图。浓雾场景 PSNR 低一些是正常现象如果常规场景和真实场景的差距能控制在 2dB 以内模型泛化基本合格。推理部署阶段值得做两件事导出 ONNX 和位置编码插值。6.1 把模型导出成 ONNX脱离 PyTorch 跑推理import torch import torch.onnx model.eval() dummy torch.randn(1, 3, 512, 512) torch.onnx.export( model, dummy, dehaze.onnx, input_names[input], output_names[output], dynamic_axes{input: {0: batch, 2: height, 3: width}} )参数说明dynamic_axes 让模型在 batch 和分辨率变化时都能导出成功避免固定尺寸导致部署环节被迫重新训练。导出后可以用 onnxruntime 或 TensorRT 加速速度和显存占用都会明显改善。6.2 推理分辨率变化时的位置编码处理训练用 256x256推理经常要处理 512 或更大的图。ViT 的位置编码是训练时固定尺寸的直接切大图会出现网格状伪影。常见做法是把位置编码用双线性插值到目标序列长度再喂给后续 transformer block。最后留一句我的习惯不要迷信某个注意力模块一定优于卷积。去雾的最终判断标准永远是人在屏幕上看到的对比度恢复和颜色准确度。指标是参考不是终点。把合成数据、真实数据、人眼验证三条线都走通这个方案才算真正落地。希望帮到你。本文还有配套的精品资源点击获取
上一篇/下一篇内容由系统自动关联 返回资讯列表 →