Vision Transformer图像去雾实战:从网络搭建到训练避坑
简介本资源面向计算机视觉方向的学生、研究人员与算法工程师提供一套基于Vision Transformer的图像去雾算法完整研究与实现方案适合具备一定深度学习基础、希望深入理解Transformer在底层视觉任务中应用的读者。压缩包共340个文件约156.35MB以204个Python源码为核心辅以39张png与10个gif结果图、16个yaml配置、12个csv实验记录、9个ipynb笔记本及8个md说明文档覆盖模型定义、训练脚本、数据处理与实验分析等模块。内容围绕ViT去雾网络搭建、训练流程与指标评估展开csv与ipynb可用于复现损失曲线与对比实验yaml便于调整超参数md文档则梳理项目结构与使用方式。目前已有97人学习下载可作为课程设计、毕业设计或科研复现的参考帮助读者快速掌握从数据准备到模型验证的完整链路。1. 从一张灰蒙蒙的照片说起Vision Transformer 图像去雾到底在做什么手里有一批户外监控抓拍图雾霾天拍出来的画面像蒙了一层灰纱远处车牌、行人轮廓全糊在一起。传统做法是上暗通道先验或者 Retinex 那一套调参调到怀疑人生换一批数据又得重来。这两年 Vision TransformerViT在去雾任务上逐渐成为主力方案原因很直接自注意力机制能建模长距离像素依赖雾的分布本身是全局性的局部卷积核感受野不够用ViT 天然适合处理这种全局退化。这个标题对应的是一套完整的工程落地包Python 源码、配套数据集、项目说明文档。它解决的核心问题是——给定一张有雾图像输出一张清晰图像并且这套流程能在本地跑通、能换数据、能改结构。适合谁有一定 Python 和 PyTorch 基础、想复现去雾算法或者拿去做二次开发的工程师也适合研究生做课题时找一个能跑通的基线。下面从原理选型一路讲到训练排错中间该抄的代码、该调的参数、该避的坑都写清楚。2. Vision Transformer 去雾的网络结构怎么搭从 Patch Embedding 到重建头2.1 为什么去雾任务适合用 ViT 而不是纯 CNN卷积网络在去雾里不是不能用早期 AOD-Net、GridDehazeNet 都是 CNN 路线效果也不差。但 CNN 的局部感受野决定了它要堆很多层才能覆盖整张图的雾分布参数量和显存开销上去了而且对非均匀雾的建模能力有限。ViT 把图像切成固定大小的 patch每个 patch 展平后加位置编码直接送进 Transformer Encoder自注意力在每一层都能让任意两个 patch 交互。雾的浓度在空间上是渐变的这种全局交互恰好对上。代价也很明显ViT 需要更多数据才能训得动patch 边界容易产生块状伪影位置编码对分辨率变化敏感。所以实际去雾方案里纯 ViT 很少见常见做法是混合结构——浅层用卷积做局部特征提取和降采样深层用 Transformer 做全局建模最后用卷积重建头恢复细节。这套思路在 Restormer、Uformer 这类工作上已经被验证过下面给的代码也是按这个混合范式来搭的。2.2 搭建一个可跑通的混合 ViT 去雾网络先看整体结构输入一张 256×256 的 RGB 有雾图经过一个 3×3 卷积做浅层特征提取然后进入四级编码器每级由 Patch Embedding步长卷积和若干 Transformer Block 组成解码器对称上采样最后卷积输出残差加上输入得到去雾结果。代码用 PyTorch 写关键部分都加了注释。import torch import torch.nn as nn import torch.nn.functional as F class PatchEmbed(nn.Module): 用步长卷积把特征图切分成 patch 并投影到 embed_dim def __init__(self, in_ch, embed_dim, patch_size2): super().__init__() self.proj nn.Conv2d(in_ch, embed_dim, kernel_sizepatch_size, stridepatch_size) self.norm nn.LayerNorm(embed_dim) def forward(self, x): x self.proj(x) # [B, C, H/p, W/p] B, C, H, W x.shape x x.flatten(2).transpose(1, 2) # [B, N, C] x self.norm(x) return x, H, W class TransformerBlock(nn.Module): 标准多头自注意力 前馈网络带残差和 LayerNorm def __init__(self, dim, num_heads8, mlp_ratio4.0): super().__init__() self.norm1 nn.LayerNorm(dim) self.attn nn.MultiheadAttention(dim, num_heads, batch_firstTrue) self.norm2 nn.LayerNorm(dim) self.mlp nn.Sequential( nn.Linear(dim, int(dim * mlp_ratio)), nn.GELU(), nn.Linear(int(dim * mlp_ratio), dim) ) def forward(self, x): # 自注意力分支 h self.norm1(x) h, _ self.attn(h, h, h) x x h # 前馈分支 x x self.mlp(self.norm2(x)) return x class DehazeViT(nn.Module): 四级编码器-解码器混合去雾网络 def __init__(self, in_ch3, base_dim48, num_blocks(2, 2, 4, 2)): super().__init__() self.shallow nn.Conv2d(in_ch, base_dim, 3, padding1) dims [base_dim, base_dim * 2, base_dim * 4, base_dim * 8] # 编码器 self.enc_embed nn.ModuleList() self.enc_blocks nn.ModuleList() for i in range(4): in_d dims[i - 1] if i 0 else base_dim self.enc_embed.append(PatchEmbed(in_d, dims[i], patch_size2)) self.enc_blocks.append(nn.Sequential( *[TransformerBlock(dims[i]) for _ in range(num_blocks[i])])) # 解码器 self.dec_blocks nn.ModuleList() self.up_convs nn.ModuleList() for i in range(3, 0, -1): self.dec_blocks.append(nn.Sequential( *[TransformerBlock(dims[i]) for _ in range(num_blocks[i - 1])])) self.up_convs.append(nn.ConvTranspose2d(dims[i], dims[i - 1], 2, stride2)) self.out_conv nn.Conv2d(base_dim, in_ch, 3, padding1) def forward(self, x): identity x feat self.shallow(x) skips [] for i in range(4): feat, H, W self.enc_embed[i](feat) feat self.enc_blocks[i](feat) B, N, C feat.shape feat feat.transpose(1, 2).reshape(B, C, H, W) skips.append(feat) # 解码 feat skips[-1] for idx, i in enumerate(range(3, 0, -1)): feat self.dec_blocks[idx](feat.flatten(2).transpose(1, 2)) B, N, C feat.shape H W int(N ** 0.5) feat feat.transpose(1, 2).reshape(B, C, H, W) feat self.up_convs[idx](feat) feat feat skips[i - 1] # 跳跃连接 out self.out_conv(feat) return torch.clamp(out identity, 0, 1) # 残差学习输出干净图逻辑说明PatchEmbed用步长卷积代替传统 ViT 的 unfold 切块好处是卷积本身带局部归纳偏置小数据集上收敛更快。TransformerBlock是标准结构batch_firstTrue让张量维度保持[B, N, C]避免反复 permute。DehazeViT的编码器逐级降分辨率、升通道解码器对称恢复跳跃连接把浅层细节直接送到对应层级缓解 Transformer 上采样后的模糊。参数说明base_dim48是显存和精度的折中8G 显存跑 256×256 的 batch size 8 没问题num_blocks(2,2,4,2)表示四个 stage 的 Transformer Block 数量深层多放几个是因为低分辨率下注意力计算量小num_heads8对dim48来说 head_dim 只有 6偏小实际部署时建议base_dim至少 64或者把第一级 head 数降到 4。patch_size2是编码器每级降采样倍率四级下来总降采样 16 倍256 输入对应最深层 16×16 的特征图。2.3 损失函数选型L1 打底感知和边缘做补充去雾不是超分像素级 L1 损失能保证整体色调不跑偏但容易让高频细节变糊。常见做法是 L1 为主加一个 VGG 感知损失和 Sobel 边缘损失做约束。权重上 L1 给 1.0感知给 0.04边缘给 0.1这个比例在多个去雾数据集上比较稳。class DehazeLoss(nn.Module): def __init__(self, vggNone): super().__init__() self.l1 nn.L1Loss() self.vgg vgg # 预训练 VGG16 的 features 部分冻结参数 self.sobel_x torch.tensor([[-1,0,1],[-2,0,2],[-1,0,1]], dtypetorch.float32).view(1,1,3,3) self.sobel_y self.sobel_x.transpose(2, 3) def edge_loss(self, pred, gt): # 对灰度图做 Sobel 卷积比较边缘响应 p pred.mean(dim1, keepdimTrue) g gt.mean(dim1, keepdimTrue) sx self.sobel_x.to(pred.device) sy self.sobel_y.to(pred.device) ex_p torch.abs(F.conv2d(p, sx, padding1)) torch.abs(F.conv2d(p, sy, padding1)) ex_g torch.abs(F.conv2d(g, sx, padding1)) torch.abs(F.conv2d(g, sy, padding1)) return self.l1(ex_p, ex_g) def forward(self, pred, gt): loss self.l1(pred, gt) if self.vgg is not None: f_pred self.vgg(pred) f_gt self.vgg(gt) loss loss 0.04 * self.l1(f_pred, f_gt) loss loss 0.1 * self.edge_loss(pred, gt) return loss逻辑说明edge_loss把 RGB 转灰度后做 Sobel比较预测图和真值图的边缘强度逼网络保留物体轮廓。感知损失用 VGG 中间层特征让纹理更自然。注意 VGG 参数要冻结否则训练初期会把感知网络也带偏。参数说明感知损失权重 0.04 是经验值太大画面会出现 VGG 特有的纹理伪影边缘损失 0.1 对去雾够用如果数据集本身分辨率低、边缘不明显可以降到 0.05。Sobel 核在 forward 里每次.to(device)有开销实际训练时建议在__init__里注册成 buffer。3. 数据集准备与训练流程从原始图像对到可收敛的模型3.1 去雾数据集的常见组织方式与预处理去雾训练需要成对数据有雾图 对应清晰图。常见来源有 RESIDE 系列ITS、OTS、Haze4K、Dense-Haze 等。标题里提到的数据集大概率是这类公开集的子集或整理版。目录结构一般是train/hazy、train/clear、test/hazy、test/clear文件名一一对应。拿到手第一件事是检查配对完整性用脚本扫一遍有没有缺图或者尺寸对不上的。import os from PIL import Image def check_pairs(hazy_dir, clear_dir): hazy_files sorted(os.listdir(hazy_dir)) clear_files sorted(os.listdir(clear_dir)) assert len(hazy_files) len(clear_files), \ f数量不一致: hazy{len(hazy_files)}, clear{len(clear_files)} bad [] for hf, cf in zip(hazy_files, clear_files): hp os.path.join(hazy_dir, hf) cp os.path.join(clear_dir, cf) with Image.open(hp) as im1, Image.open(cp) as im2: if im1.size ! im2.size: bad.append((hf, im1.size, cf, im2.size)) if bad: print(f发现 {len(bad)} 对尺寸不匹配:) for item in bad[:10]: print(item) else: print(f全部 {len(hazy_files)} 对配对正常) return bad check_pairs(./data/train/hazy, ./data/train/clear)逻辑说明这个脚本先比数量再比尺寸尺寸不一致的样本在训练时会触发广播错误或者被错误 resize提前查出来能省很多调试时间。实际项目里还应该检查图像模式RGB 还是灰度、是否有损坏文件。参数说明sorted保证两边顺序一致前提是命名规则相同。如果命名规则不同需要改成按 ID 匹配。Image.open用上下文管理器避免文件句柄泄漏数据量大时这点很重要。预处理方面训练时随机裁剪 256×256随机水平翻转不做垂直翻转去雾任务里天空和地面的先验不同垂直翻转会破坏这个分布。归一化用 ImageNet 的均值和方差就行因为浅层卷积和 VGG 感知损失都是基于这个分布预训练的。3.2 训练脚本的关键参数与显存控制训练循环本身不复杂关键是几个参数要设对。优化器用 AdamW初始学习率 2e-4余弦退火到 1e-6权重衰减 1e-4。batch size 根据显存来8G 卡跑上面那个base_dim48的模型256×256 输入大概能到 8。不够就上梯度累积别硬撑大 batch去雾任务对 batch size 没那么敏感。import torch from torch.utils.data import DataLoader from torch.optim.lr_scheduler import CosineAnnealingLR def train_one_epoch(model, loader, optimizer, criterion, device, accum_steps1): model.train() total_loss 0.0 optimizer.zero_grad() for step, (hazy, clear) in enumerate(loader): hazy, clear hazy.to(device), clear.to(device) pred model(hazy) loss criterion(pred, clear) / accum_steps loss.backward() if (step 1) % accum_steps 0: torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0) optimizer.step() optimizer.zero_grad() total_loss loss.item() * accum_steps return total_loss / len(loader) # 主训练配置 device torch.device(cuda if torch.cuda.is_available() else cpu) model DehazeViT(base_dim48).to(device) criterion DehazeLoss(vggNone).to(device) # 先用纯 L1边缘跑通 optimizer torch.optim.AdamW(model.parameters(), lr2e-4, weight_decay1e-4) scheduler CosineAnnealingLR(optimizer, T_max100, eta_min1e-6) for epoch in range(100): loss train_one_epoch(model, train_loader, optimizer, criterion, device, accum_steps2) scheduler.step() if epoch % 10 0: print(fEpoch {epoch}, loss{loss:.4f}, lr{scheduler.get_last_lr()[0]:.2e}) if epoch % 20 0: torch.save(model.state_dict(), fdehaze_epoch{epoch}.pth)逻辑说明梯度累积用accum_steps控制等效 batch size 等于实际 batch 乘以累积步数。梯度裁剪max_norm1.0对 Transformer 类模型几乎是必须的注意力层的梯度容易爆。学习率调度用余弦退火比 StepLR 更平滑去雾这种回归任务上收敛曲线更稳。参数说明lr2e-4是 AdamW 在 ViT 类模型上的常用起点如果 loss 前几个 epoch 就震荡降到 1e-4。weight_decay1e-4对 Transformer 参数做正则别设太大1e-2 级别会把注意力权重压死。T_max100要和总 epoch 数一致中途改总轮数记得同步改这个值。保存策略上每 20 epoch 存一次去雾任务后期提升慢存太密浪费磁盘。3.3 验证指标PSNR 和 SSIM 怎么算才不出错训练时看 loss验证时看 PSNR 和 SSIM。这两个指标计算有几个坑PSNR 要在 [0,1] 或 [0,255] 同一量纲下算SSIM 要用高斯窗口而不是均匀窗口彩色图要在 YCbCr 的 Y 通道上算。下面给一个不依赖第三方库的实现。import torch import torch.nn.functional as F import math def compute_psnr(pred, gt): pred, gt 均为 [0,1] 范围的 tensor mse F.mse_loss(pred, gt) if mse 0: return float(inf) return 10 * math.log10(1.0 / mse.item()) def compute_ssim(pred, gt, window_size11, sigma1.5): 在 Y 通道上计算 SSIM输入 [B,3,H,W] # RGB 转 Y def to_y(x): return 0.299 * x[:,0] 0.587 * x[:,1] 0.114 * x[:,2] p, g to_y(pred).unsqueeze(1), to_y(gt).unsqueeze(1) # 高斯窗口 coords torch.arange(window_size, dtypetorch.float32) - window_size // 2 gauss torch.exp(-(coords ** 2) / (2 * sigma ** 2)) gauss (gauss / gauss.sum()).unsqueeze(0) kernel gauss.t() gauss kernel kernel.unsqueeze(0).unsqueeze(0).to(pred.device) # 局部均值、方差、协方差 mu_p F.conv2d(p, kernel, paddingwindow_size//2) mu_g F.conv2d(g, kernel, paddingwindow_size//2) sigma_p F.conv2d(p * p, kernel, paddingwindow_size//2) - mu_p ** 2 sigma_g F.conv2d(g * g, kernel, paddingwindow_size//2) - mu_g ** 2 sigma_pg F.conv2d(p * g, kernel, paddingwindow_size//2) - mu_p * mu_g C1, C2 0.01 ** 2, 0.03 ** 2 ssim_map ((2 * mu_p * mu_g C1) * (2 * sigma_pg C2)) / \ ((mu_p ** 2 mu_g ** 2 C1) * (sigma_p sigma_g C2)) return ssim_map.mean().item()逻辑说明PSNR 直接由 MSE 推导注意输入范围。SSIM 先转 Y 通道再用高斯核做局部统计C1、C2是稳定常数对应动态范围 1.0。sigma_p理论上应该非负浮点误差可能出小负数实际影响可忽略。参数说明window_size11、sigma1.5是 SSIM 原论文的推荐值别乱改改了和别人报的指标没法比。验证时要把模型设成eval()并且torch.no_grad()否则 BatchNorm 和 Dropout 会让指标偏低。如果数据集测试集有 GT直接算没有 GT 就只能靠视觉对比这时候可以拿 FADE 这类无参考指标做辅助但别当唯一标准。4. 避坑与排查去雾训练里最容易翻车的五个地方4.1 训练 loss 正常下降但输出全灰或全黑现象训练日志里 loss 从 0.3 降到 0.05看起来收敛很好但推理时输出一张接近纯灰或者纯黑的图。原因通常是残差学习的输出被 clamp 截断后梯度消失或者最后一层初始化让输出偏向 0。解决检查out_conv的初始化用nn.init.zeros_把最后一层权重和偏置清零让网络从恒等映射开始学同时确认torch.clamp(out identity, 0, 1)里identity是有雾输入而不是清晰图。如果还是灰把 L1 换成 Charbonnier 损失试试对低误差区域梯度更友好。4.2 PSNR 很高但视觉上雾没去干净现象测试集 PSNR 到 30dB 以上但肉眼看远处还是灰蒙蒙。原因是 L1 损失对大面积的均匀雾区优化得好对局部浓雾区不够敏感而 PSNR 是全局平均浓雾区的小误差被大片干净区稀释了。解决在损失里加一个基于暗通道先验的辅助项或者对雾浓度高的区域加权。简单做法是用有雾图的暗通道值当权重图暗通道值高的地方雾浓给更大权重。另外验证时别只看平均 PSNR分区域统计或者看最差的那几张。4.3 显存溢出OOM但 batch size 已经降到 1现象batch size 设成 1 还是 OOM报错指向注意力层。原因是自注意力的显存占用和 patch 数量的平方成正比256×256 输入经过四级降采样后最深层是 16×16256 个 token注意力矩阵 256×256 不大但第一级降采样后是 128×12816384 个 token注意力矩阵 16384×16384 直接爆。解决第一级不要用全局自注意力改成窗口注意力或者用卷积替代或者把patch_size从 2 改成 4第一级 token 数降到 4096。另一个办法是用torch.utils.checkpoint对 Transformer Block 做梯度检查点显存换时间。4.4 换数据集后效果断崖式下跌现象在 RESIDE 上训到 32dB换到自采的雾天数据上输出偏色严重。原因是不同数据集的雾生成模型不同RESIDE 是合成雾自采是真实雾域 gap 很大。解决先做域适应用自采数据做无监督微调损失用对比学习或者一致性正则或者简单点把自采数据里清晰的部分挑出来当 GT有雾的部分当输入做小规模有监督微调。别指望一个合成数据训出来的模型直接上真实场景这是去雾领域的老大难问题。4.5 推理速度太慢单张 1080P 要好几秒现象模型在 256×256 上跑得挺快部署到 1080P 输入时慢到不可用。原因是 ViT 的计算量随分辨率平方增长1080P 直接送进去 token 数爆炸。解决推理时用滑动窗口或者多尺度策略把大图切成有重叠的 256×256 块分别处理再拼接重叠区域做加权融合消除接缝。或者把模型导出成 ONNX 再用 TensorRT 加速Transformer 类模型在 TensorRT 上一般能有 2 到 3 倍提升。如果还不行考虑知识蒸馏用大模型教一个小 CNN 学生模型推理时只跑学生模型。5. 进阶技巧用预训练权重和混合精度把训练成本压下来如果从零训上面那个网络单卡 8G 显存跑 100 epoch 大概要两三天。实际项目里没必要从零开始两个技巧能省一大半时间。第一个是加载 ImageNet 预训练的 ViT 权重把编码器前几层初始化成预训练参数解码器随机初始化。虽然去雾和分类任务差异大但浅层的边缘、纹理特征是可迁移的。加载时注意维度匹配base_dim和预训练模型的embed_dim不一致时用插值或者只加载部分层。def load_pretrained_encoder(model, ckpt_path, base_dim48): 加载预训练 ViT 权重到编码器维度不匹配的层跳过 ckpt torch.load(ckpt_path, map_locationcpu) state ckpt.get(model, ckpt) model_dict model.state_dict() loaded, skipped 0, 0 for k, v in state.items(): # 只加载编码器里的注意力和 MLP 权重 if enc_blocks in k and k in model_dict: if model_dict[k].shape v.shape: model_dict[k] v loaded 1 else: skipped 1 model.load_state_dict(model_dict) print(f加载 {loaded} 个参数跳过 {skipped} 个维度不匹配的) return model逻辑说明只匹配enc_blocks里的参数浅层卷积和 Patch Embedding 因为输入通道和维度不同一般不加载。shape检查是必须的否则load_state_dict直接报错。加载完打印数量方便确认是不是真的加载上了。第二个技巧是混合精度训练AMP用torch.cuda.amp把前向和反向的矩阵运算降到 fp16显存占用能降 30% 到 40%速度提升 20% 左右。注意损失计算和梯度累积要在 fp32 下做否则数值不稳定。scaler torch.cuda.amp.GradScaler() def train_amp(model, loader, optimizer, criterion, device): model.train() for hazy, clear in loader: hazy, clear hazy.to(device), clear.to(device) optimizer.zero_grad() with torch.cuda.amp.autocast(): pred model(hazy) loss criterion(pred, clear) scaler.scale(loss).backward() scaler.unscale_(optimizer) torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) scaler.step(optimizer) scaler.update()逻辑说明autocast上下文里前向用 fp16GradScaler负责放大损失避免 fp16 下溢。unscale_之后再裁剪梯度顺序不能反否则裁剪的是放大后的梯度。scaler.update()根据是否出现 inf/nan 动态调整缩放因子。参数说明AMP 对去雾任务基本无损PSNR 差异在 0.05dB 以内。如果训练中出现 loss 变 nan先检查是不是感知损失里的 VGG 在 fp16 下溢出把 VGG 部分强制 fp32 就行。另外GradScaler的初始缩放因子默认 65536一般不用改。最后说一个验证技巧训练完别只看测试集指标把模型输出和输入、GT 拼成三栏对比图随机抽 20 张看。去雾的很多问题——颜色偏移、局部残留雾、过度增强导致噪声放大——指标上看不出来肉眼一看就现形。我自己习惯是每个 epoch 存一张验证图的对比训练结束后翻一遍比看 loss 曲线有用得多。这套方案值不值得做如果你手头有去雾需求又不想从零推导公式拿这个结构改改输入输出就能用数据准备和训练流程都是现成的。希望帮到你。本文还有配套的精品资源点击获取
上一篇/下一篇内容由系统自动关联
返回资讯列表 →