尧图网络科技YAOTU DIGITAL 获取报价
获取报价
首页 / 资讯中心 / 文章详情

基于Vision Transformer的图像去雾:从注意力机制到Python工程实战

发布时间:2026/9/28 23:11:45

资讯中心
01
ARTICLE

基于Vision Transformer的图像去雾:从注意力机制到Python工程实战

基于Vision Transformer的图像去雾:从注意力机制到Python工程实战
简介这份资源是面向计算机相关专业学生与项目实战学习者的图像去雾算法实现项目基于VisionTransformer架构完成可作为毕业设计、课程设计或期末大作业的参考方案。项目经导师指导并通过评审源码均经本地编译调试确保可运行难度适中适合具备一定深度学习基础、希望掌握Transformer在底层视觉任务中应用的读者。压缩包共338个文件约156.35MB以204个Python源码为主体辅以yaml配置、csv实验记录、ipynb笔记本、png与gif可视化结果及md说明文档覆盖模型定义、训练配置、损失曲线与效果对比等环节。已有102人学习关注。读者可获得完整的去雾算法实现流程、可复现的训练与评估脚本、实验数据记录以及文档说明便于快速理解ViT在图像复原中的设计思路并在此基础上进行二次开发或论文写作。1. 基于 Vision Transformer 的图像去雾从注意力机制到可复现的 Python 工程雾霾天拍出来的照片远景发灰、对比度塌陷、颜色偏移这是图像去雾要解决的核心问题。传统方法靠暗通道先验在天空区域和高亮区域经常翻车CNN 方案感受野有限对长距离依赖建模不足。Vision Transformer 把图像切成 patch 序列用自注意力直接建模任意两个位置的关系天然适合去雾这种需要全局理解雾浓度分布的任务。这个方向适合有 Python 和 PyTorch 基础、想做一个完整深度学习项目的人也适合需要一份能跑通、能改、能写进简历的去雾方案。下面从原理到代码把这条路走一遍。2. Vision Transformer 去雾的网络结构怎么搭patch 嵌入与注意力恢复2.1 为什么去雾任务适合用 Transformer 而不是纯 CNN去雾的本质是估计每个像素的透射率或直接回归清晰图像。雾的分布不是局部的——一片天空的雾浓度会影响整片区域的亮度估计CNN 的卷积核再大感受野也是逐层堆出来的浅层看不到远处深层又丢了细节。Transformer 的自注意力是全局的每个 patch 都能直接和所有其他 patch 交互这对估计非均匀雾场很关键。但纯 Transformer 也有问题计算量随 patch 数量平方增长高分辨率去雾图直接切 patch 会爆显存。常见做法是先用卷积做浅层特征提取和降采样再送入 Transformer 块做全局建模最后用卷积上采样恢复分辨率。这种混合结构在去雾任务里比纯 ViT 更实用显存占用可控细节恢复也更好。另一个选型理由是去雾数据集通常不大RESIDE 是常用的但真实配对数据难获取纯 ViT 需要大量数据预训练。混合结构里卷积部分可以复用预训练权重Transformer 部分用较小参数量也能收敛。我一般会先用 ImageNet 预训练的 backbone 初始化再在去雾数据上微调。2.2 最小可运行的 ViT 去雾网络代码下面是一个可以直接跑通的简化版网络输入雾图输出清晰图。结构是浅层卷积 → patch 嵌入 → Transformer 编码器 → 解码上采样。import torch import torch.nn as nn import torch.nn.functional as F class PatchEmbed(nn.Module): 把特征图切成 patch 并线性嵌入 def __init__(self, in_ch64, embed_dim256, patch_size4): super().__init__() self.patch_size patch_size self.proj nn.Conv2d(in_ch, embed_dim, kernel_sizepatch_size, stridepatch_size) def forward(self, x): # x: [B, 64, H, W] - [B, embed_dim, H/p, W/p] x self.proj(x) B, C, H, W x.shape # 展平成序列: [B, N, C] x x.flatten(2).transpose(1, 2) return x, (H, W) class TransformerBlock(nn.Module): 标准 Transformer 编码器块 def __init__(self, dim256, num_heads8, mlp_ratio4.0, dropout0.1): super().__init__() self.norm1 nn.LayerNorm(dim) self.attn nn.MultiheadAttention(dim, num_heads, dropoutdropout, batch_firstTrue) self.norm2 nn.LayerNorm(dim) self.mlp nn.Sequential( nn.Linear(dim, int(dim * mlp_ratio)), nn.GELU(), nn.Dropout(dropout), nn.Linear(int(dim * mlp_ratio), dim), nn.Dropout(dropout) ) def forward(self, x): # 自注意力 残差 x x self.attn(self.norm1(x), self.norm1(x), self.norm1(x))[0] # FFN 残差 x x self.mlp(self.norm2(x)) return x class ViTDehaze(nn.Module): 混合 ViT 去雾网络 def __init__(self, base_ch64, embed_dim256, depth6, num_heads8): super().__init__() # 浅层特征提取 self.shallow nn.Sequential( nn.Conv2d(3, base_ch, 3, padding1), nn.ReLU(inplaceTrue), nn.Conv2d(base_ch, base_ch, 3, padding1), nn.ReLU(inplaceTrue) ) # patch 嵌入 self.patch_embed PatchEmbed(base_ch, embed_dim, patch_size4) # Transformer 编码器 self.blocks nn.ModuleList([ TransformerBlock(embed_dim, num_heads) for _ in range(depth) ]) self.norm nn.LayerNorm(embed_dim) # 解码上采样回原分辨率 self.decoder nn.Sequential( nn.Conv2d(embed_dim, base_ch * 4, 3, padding1), nn.PixelShuffle(2), # 2倍上采样 nn.ReLU(inplaceTrue), nn.Conv2d(base_ch, base_ch * 4, 3, padding1), nn.PixelShuffle(2), # 再2倍共4倍 nn.ReLU(inplaceTrue), nn.Conv2d(base_ch, 3, 3, padding1), nn.Sigmoid() # 输出归一化到 [0,1] ) def forward(self, x): # 浅层特征 feat self.shallow(x) # [B, 64, H, W] # patch 嵌入 tokens, (H, W) self.patch_embed(feat) # [B, N, 256] # Transformer 编码 for blk in self.blocks: tokens blk(tokens) tokens self.norm(tokens) # 恢复空间维度 B, N, C tokens.shape feat tokens.transpose(1, 2).reshape(B, C, H, W) # 解码上采样 out self.decoder(feat) # 残差连接学习雾的残差比直接回归清晰图更稳 return torch.clamp(x - out 0.5, 0, 1)这段代码的关键设计点PatchEmbed用 stride 等于 patch_size 的卷积实现切块和嵌入比手动 unfold 更高效TransformerBlock是标准的 pre-norm 结构训练更稳定解码器用PixelShuffle做上采样比转置卷积更不容易产生棋盘伪影最后输出用残差形式网络学的是「雾的残差」而不是直接回归清晰图收敛更快。参数说明embed_dim256是 token 维度显存不够可以降到 128depth6是 Transformer 块数去雾任务一般 4 到 8 层够用num_heads8是注意力头数通常设为 embed_dim 的 1/32patch_size4意味着 256x256 输入会产生 64x644096 个 token如果显存吃紧可以改成 8。2.3 损失函数怎么选L1、感知损失与 SSIM 的组合去雾如果只用 L1 损失结果会偏平滑细节丢失。常见做法是 L1 感知损失 SSIM 损失加权组合。感知损失用预训练 VGG 提取特征让输出在语义层面接近清晰图SSIM 损失直接优化结构相似度对去雾这种结构恢复任务很有效。import torch import torch.nn as nn import torchvision.models as models class DehazeLoss(nn.Module): def __init__(self, w_l11.0, w_per0.1, w_ssim0.5): super().__init__() self.w_l1 w_l1 self.w_per w_per self.w_ssim w_ssim # 感知损失用 VGG16 的前几层 vgg models.vgg16(pretrainedTrue).features[:16].eval() for p in vgg.parameters(): p.requires_grad False self.vgg vgg def forward(self, pred, target): # L1 损失 l1 F.l1_loss(pred, target) # 感知损失 feat_pred self.vgg(pred) feat_target self.vgg(target) per F.l1_loss(feat_pred, feat_target) # SSIM 损失简化版用 1 - SSIM ssim_val self.ssim(pred, target) ssim_loss 1 - ssim_val return self.w_l1 * l1 self.w_per * per self.w_ssim * ssim_loss def ssim(self, x, y): # 简化 SSIM 计算实际可用 pytorch-ssim 库 C1, C2 0.01**2, 0.03**2 mu_x F.avg_pool2d(x, 3, 1, 1) mu_y F.avg_pool2d(y, 3, 1, 1) sigma_x F.avg_pool2d(x**2, 3, 1, 1) - mu_x**2 sigma_y F.avg_pool2d(y**2, 3, 1, 1) - mu_y**2 sigma_xy F.avg_pool2d(x*y, 3, 1, 1) - mu_x*mu_y ssim_map ((2*mu_x*mu_y C1)*(2*sigma_xy C2)) / \ ((mu_x**2 mu_y**2 C1)*(sigma_x sigma_y C2)) return ssim_map.mean()权重建议w_l11.0是基础w_per0.1不要太大否则颜色会偏w_ssim0.5对结构恢复帮助明显。如果训练时发现输出偏暗可以把w_l1降到 0.8让感知损失占比相对提高。3. 训练流程与数据准备从 RESIDE 数据集到可收敛的配置3.1 数据集怎么组织和加载去雾常用 RESIDE 数据集包含室内和室外子集。常见做法是用 ITS室内或 OTS室外做训练SOTS 做测试。数据组织成配对形式clear/放清晰图hazy/放对应雾图文件名一一对应。import os from torch.utils.data import Dataset, DataLoader from PIL import Image import torchvision.transforms as T class DehazeDataset(Dataset): def __init__(self, clear_dir, hazy_dir, size256, trainTrue): self.clear_dir clear_dir self.hazy_dir hazy_dir self.names sorted(os.listdir(clear_dir)) self.size size self.train train self.transform T.Compose([ T.Resize((size, size)), T.ToTensor() ]) def __len__(self): return len(self.names) def __getitem__(self, idx): name self.names[idx] clear Image.open(os.path.join(self.clear_dir, name)).convert(RGB) hazy Image.open(os.path.join(self.hazy_dir, name)).convert(RGB) clear self.transform(clear) hazy self.transform(hazy) # 训练时随机翻转增强 if self.train: if torch.rand(1) 0.5: clear T.functional.hflip(clear) hazy T.functional.hflip(hazy) return hazy, clear # 使用示例 train_set DehazeDataset(data/RESIDE/clear, data/RESIDE/hazy, size256) train_loader DataLoader(train_set, batch_size8, shuffleTrue, num_workers4)注意RESIDE 的 ITS 子集有 13990 对训练图OTS 有 72 万对左右显存和时间有限的话先用 ITS。测试用 SOTS 的 500 张室内图。如果找不到 RESIDE也可以用自己合成的配对数据拿清晰图按大气散射模型加雾。3.2 训练循环与关键超参import torch from torch.optim import Adam from torch.optim.lr_scheduler import CosineAnnealingLR device torch.device(cuda if torch.cuda.is_available() else cpu) model ViTDehaze(base_ch64, embed_dim256, depth6).to(device) criterion DehazeLoss().to(device) optimizer Adam(model.parameters(), lr2e-4, betas(0.9, 0.999)) scheduler CosineAnnealingLR(optimizer, T_max100, eta_min1e-6) best_psnr 0 for epoch in range(100): model.train() total_loss 0 for hazy, clear in train_loader: hazy, clear hazy.to(device), clear.to(device) pred model(hazy) loss criterion(pred, clear) optimizer.zero_grad() loss.backward() # 梯度裁剪防止 Transformer 训练不稳定 torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) optimizer.step() total_loss loss.item() scheduler.step() avg_loss total_loss / len(train_loader) print(fEpoch {epoch1}, Loss: {avg_loss:.4f}, LR: {scheduler.get_last_lr()[0]:.6f}) # 每 10 轮存一次权重 if (epoch 1) % 10 0: torch.save(model.state_dict(), fvit_dehaze_epoch{epoch1}.pth)关键超参lr2e-4是 Transformer 微调的常用起点太大容易震荡太小收敛慢batch_size8在 256x256 输入下大约占 8G 显存不够就降到 4 并累积梯度clip_grad_norm_的 1.0 是防止梯度爆炸的后悔药Transformer 训练初期梯度经常很大。CosineAnnealingLR比 StepLR 更平滑去雾任务一般 100 到 200 轮收敛。3.3 评估指标PSNR 和 SSIM 怎么算才不骗自己import numpy as np from skimage.metrics import peak_signal_noise_ratio as psnr from skimage.metrics import structural_similarity as ssim def evaluate(model, test_loader, device): model.eval() psnr_list, ssim_list [], [] with torch.no_grad(): for hazy, clear in test_loader: hazy hazy.to(device) pred model(hazy).cpu().numpy() clear clear.numpy() for i in range(pred.shape[0]): p np.transpose(pred[i], (1, 2, 0)) c np.transpose(clear[i], (1, 2, 0)) psnr_list.append(psnr(c, p, data_range1.0)) ssim_list.append(ssim(c, p, data_range1.0, channel_axis2)) return np.mean(psnr_list), np.mean(ssim_list)注意PSNR 要在 RGB 空间算data_range1.0因为输出是 [0,1]。SSIM 的channel_axis2表示按通道算。如果 PSNR 很高但视觉上还是有雾说明模型可能过平滑了这时候要看感知损失权重是不是太低。4. 避坑与排查ViT 去雾训练中最容易翻车的 5 个地方4.1 损失不下降输出全是灰色现象训练几十轮后输出图几乎全是灰色PSNR 卡在 12 左右不动。原因残差连接写错了x - out 0.5里的out如果初始化偏大输出直接饱和。解决把最后一层卷积的权重初始化改小或者去掉0.5改成x - out并确保out范围在 [-0.5, 0.5]。更稳的做法是输出层不加 Sigmoid用torch.tanh然后乘 0.5。4.2 显存爆炸batch_size 降到 1 还 OOM现象256x256 输入batch_size1 仍然 OOM。原因patch_size4 时 token 数是 64x644096注意力矩阵是 4096x4096显存占用是 token 数的平方。解决把 patch_size 改成 8token 数降到 1024注意力矩阵缩小 16 倍或者用梯度检查点torch.utils.checkpoint换显存。4.3 训练集 PSNR 很高测试集一塌糊涂现象训练集 PSNR 到 35测试集只有 18。原因过拟合RESIDE ITS 只有 13990 对模型参数量太大。解决加数据增强随机裁剪、颜色抖动加 dropoutTransformer 块里设 0.1或者用 OTS 的 72 万对数据做预训练再在 ITS 上微调。4.4 输出有棋盘伪影现象放大看输出图有网格状伪影。原因解码器用了转置卷积ConvTranspose2dstride 和 kernel_size 不匹配时会产生棋盘效应。解决换成PixelShuffle或者F.interpolate 卷积上面代码里用的就是PixelShuffle如果还有伪影就检查上采样倍数和卷积核是否对齐。4.5 颜色偏暗或偏蓝现象去雾后整体偏暗或者天空区域偏蓝。原因L1 损失对亮度敏感训练数据里暗图多的话模型会偏向输出暗图感知损失用 VGG 在 ImageNet 上预训练对蓝色通道响应强。解决调整损失权重把w_l1降到 0.8w_per降到 0.05或者在数据增强里加随机亮度调整让模型见过各种亮度。5. 进阶技巧用注意力图可视化验证模型到底学到了什么训练完一个去雾模型怎么判断它是真的在去雾还是只是学了颜色映射我一般会做两件事一是把 Transformer 最后一层的注意力图可视化看模型关注哪些区域二是用合成雾图做控制变量测试固定清晰图改变雾浓度看输出是否跟着变。注意力图可视化的做法在TransformerBlock的forward里把attn的权重存下来取最后一个 block 的平均注意力reshape 成 H×W 的热力图叠加在原图上。如果模型真的在去雾注意力应该集中在雾浓度高的区域比如远景和天空而不是均匀分布。# 在 TransformerBlock.forward 里加 attn_weights, _ self.attn(self.norm1(x), self.norm1(x), self.norm1(x)) self.last_attn attn_weights.detach() # 存下来供可视化 # 可视化时 attn model.blocks[-1].last_attn # [B, N, N] attn_map attn.mean(dim1) # 平均所有 query 的注意力 attn_map attn_map.reshape(B, H, W) # 恢复空间维度 # 归一化后叠加到原图另一个验证方法是做「雾浓度扫描」拿一张清晰图按大气散射模型合成不同浓度的雾图beta 从 0.5 到 3.0送入模型看输出的 PSNR 是否随浓度增加而单调下降。如果浓度很高时 PSNR 反而上升说明模型在「猜」而不是在「去雾」。我自己的习惯是每次改完网络结构或损失函数先跑 10 个 epoch 看损失曲线和一张验证图的输出确认没有明显翻车再跑完整训练。这个习惯帮我省了很多次白跑 100 轮的时间。希望帮到你。本文还有配套的精品资源点击获取
02
RELATED NEWS

相关资讯

更多网站建设与数字化升级内容

03
WHY YAOTU

想打造同款高转化官网?

懂行业、懂生意,从建站到增长一站式陪跑

◈

场景化定制

不做模板站,围绕你的业务场景量身设计,小众不撞款。

◐

营销型架构

以转化目标组织内容与路径,让官网真正带来询盘。

▲

全周期服务

设计、开发、运营、运维一体,上线只是开始。

免费获取你的建站方案

留下需求,专属顾问 24 小时内为你输出方案建议。