DehazeFormer 原理详解:Transformer 如何重塑图像去雾技术
一、引言
图像去雾是计算机视觉中的经典难题。传统方法依赖大气散射模型估计透射率和大气光,而深度学习方法(如 AOD-Net、FFA-Net)基于 CNN,感受野有限。DehazeFormer 创新地将 Vision Transformer 引入去雾任务,利用自注意力机制建模全局依赖关系,在多个基准数据集上达到 SOTA。
本文将深入解析 DehazeFormer 的架构设计和核心原理。
二、大气散射模型回顾
在理解 DehazeFormer 之前,先回顾雾天成像的物理模型:
I(x) = J(x) × t(x) + A × (1 – t(x))
其中:
- I(x): 观测到的有雾图像
- J(x): 需要恢复的清晰图像
- t(x): 透射率图(transmission map),t(x) = e^(-β·d(x))
- A: 全局大气光(global atmospheric light)
- β: 大气散射系数
- d(x): 场景深度
去雾的目标是从 I(x) 恢复 J(x),需要估计 t(x) 和 A。这是一个欠定问题。
三、DehazeFormer 架构详解
3.1 整体架构
输入: 有雾图像 I ∈ ℝ^(3×H×W)
│
▼
┌─────────────────────┐
│ Patch Embedding │ Conv 3→C, k=3, 产生 patch 特征
│ (无重叠分块) │
└─────────┬───────────┘
│
▼
┌─────────────────────┐
│ SK Fusion Layer 0 │ SKFF: 选择性核特征融合
│ (多尺度特征融合) │ ┌─ 3×3 卷积分支 ─┐
│ │ ├─ 5×5 卷积分支 ─┤→ 自适应融合
│ │ └─ 7×7 卷积分支 ─┘
└─────────┬───────────┘
│
▼
┌─────────────────────┐
│ DehazeFormer Block │ × N (通常 N=4~8)
│ × 4 │
│ │ ┌── LayerNorm ──┐
│ │ │ │
│ ┌─────────────────┐ │ │ ┌──────────┐ │
│ │ LayerNorm │ │ ├──│ W-MSA │─┤ → +
│ └────────┬────────┘ │ │ └──────────┘ │
│ │ │ │ │
│ ┌────────▼────────┐ │ │ ┌──────────┐ │
│ │ W-MSA / SW-MSA │ │ ├──│ MLP │─┤ → +
│ │ (窗口/移位窗口) │ │ │ └──────────┘ │
│ └────────┬────────┘ │ │ │
│ │ │ └───────────────┘
│ ┌────────▼────────┐ │
│ │ LayerNorm │ │
│ └────────┬────────┘ │
│ │ │
│ ┌────────▼────────┐ │
│ │ FeedForward │ │ 深度可分离卷积 MLP
│ │ (DWConv + GELU) │ │
│ └────────┬────────┘ │
│ │ │
│ ▼ │
│ + 残差 │
└─────────┬───────────┘
│
▼
┌─────────────────────┐
│ 重建头 (Refinement) │ Conv → PixelShuffle → Conv
│ │ 或直接 Conv 输出
└─────────┬───────────┘
│
▼
输出: 去雾图像 J ∈ ℝ^(3×H×W)
3.2 核心组件实现
3.2.1 窗口多头自注意力(W-MSA)
import torch
import torch.nn as nn
import torch.nn.functional as F
class WindowAttention(nn.Module):
"""基于窗口的多头自注意力"""
def __init__(self, dim, window_size, num_heads):
super().__init__()
self.dim = dim
self.window_size = window_size
self.num_heads = num_heads
head_dim = dim // num_heads
self.scale = head_dim ** –0.5
self.qkv = nn.Linear(dim, dim * 3, bias=True)
self.proj = nn.Linear(dim, dim)
# 相对位置偏置(可学习)
self.relative_position_bias_table = nn.Parameter(
torch.zeros((2 * window_size – 1) * (2 * window_size – 1), num_heads)
)
def forward(self, x):
B_, N, C = x.shape # [B*num_windows, window_size², C]
# QKV 投影
qkv = self.qkv(x).reshape(B_, N, 3, self.num_heads, C // self.num_heads)
qkv = qkv.permute(2, 0, 3, 1, 4)
q, k, v = qkv[0], qkv[1], qkv[2]
# 注意力
attn = (q @ k.transpose(–2, –1)) * self.scale
# 添加相对位置偏置
relative_position_bias = self.relative_position_bias_table[
self.relative_position_index.view(–1)
].view(self.window_size**2, self.window_size**2, –1)
relative_position_bias = relative_position_bias.permute(2, 0, 1)
attn = attn + relative_position_bias.unsqueeze(0)
attn = F.softmax(attn, dim=–1)
x = (attn @ v).transpose(1, 2).reshape(B_, N, C)
x = self.proj(x)
return x
3.2.2 移位窗口机制(SW-MSA)
W-MSA 只在窗口内计算注意力,缺少窗口间的信息交互。SW-MSA 通过将窗口移位来解决:
def window_partition(x, window_size):
"""将特征图划分为窗口"""
B, H, W, C = x.shape
x = x.view(B, H // window_size, window_size,
W // window_size, window_size, C)
windows = x.permute(0, 1, 3, 2, 4, 5)
windows = windows.contiguous().view(–1, window_size, window_size, C)
return windows
def window_reverse(windows, window_size, H, W):
"""窗口恢复为特征图"""
B = int(windows.shape[0] / (H // window_size * W // window_size))
x = windows.view(B, H // window_size, W // window_size,
window_size, window_size, –1)
x = x.permute(0, 1, 3, 2, 4, 5)
x = x.contiguous().view(B, H, W, –1)
return x
3.2.3 SKFF(Selective Kernel Feature Fusion)
SKFF 模块融合多尺度特征,是 DehazeFormer 的另一个关键创新:
class SKFF(nn.Module):
"""Selective Kernel Feature Fusion"""
def __init__(self, channels, num_features=32):
super().__init__()
self.conv3 = nn.Conv2d(channels, channels, 3, 1, 1, groups=channels)
self.conv5 = nn.Conv2d(channels, channels, 5, 1, 2, groups=channels)
self.conv7 = nn.Conv2d(channels, channels, 7, 1, 3, groups=channels)
# 特征选择(Squeeze → Excitation)
self.fc_reduce = nn.Conv2d(channels * 3, num_features, 1)
self.fc_expand3 = nn.Conv2d(num_features, channels, 1)
self.fc_expand5 = nn.Conv2d(num_features, channels, 1)
self.fc_expand7 = nn.Conv2d(num_features, channels, 1)
def forward(self, x):
f3 = self.conv3(x)
f5 = self.conv5(x)
f7 = self.conv7(x)
# 全局平均池化 → 通道选择
u = torch.cat([f3, f5, f7], dim=1) # [B, C×3, H, W]
u_gap = F.adaptive_avg_pool2d(u, 1) # [B, C×3, 1, 1]
u_gap = self.fc_reduce(u_gap)
# 软注意力
a3 = self.fc_expand3(u_gap).sigmoid()
a5 = self.fc_expand5(u_gap).sigmoid()
a7 = self.fc_expand7(u_gap).sigmoid()
# 自适应融合
out = f3 * a3 + f5 * a5 + f7 * a7
return out + x
四、与 CNN 方法的对比
| 感受野 | 局部(受卷积核大小限制) | 全局(自注意力) |
| 特征交互 | 逐层渐进 | 自由长程交互 |
| 多尺度处理 | 多分支/特征金字塔 | SKFF 自适应融合 |
| 参数量 | ~4.5M | ~2.5M(轻量版本) |
| PSNR (SOTS-indoor) | 36.39 | 36.82 |
五、训练要点
5.1 损失函数
DehazeFormer 使用组合损失:
class DehazeLoss(nn.Module):
def __init__(self):
super().__init__()
self.l1 = nn.L1Loss()
self.perceptual = PerceptualLoss() # VGG 感知损失
def forward(self, pred, target):
l1_loss = self.l1(pred, target)
perc_loss = self.perceptual(pred, target)
return l1_loss + 0.04 * perc_loss
5.2 数据增强
- 随机裁剪 256×256
- 水平/垂直翻转
- 旋转(90°、180°、270°)
- 色彩抖动
六、总结
DehazeFormer 将 Swin Transformer 架构成功应用于图像去雾,核心创新包括:W-MSA/SW-MSA 的窗口注意力机制降低了计算复杂度,SKFF 多尺度自适应融合增强了细节恢复能力。相比传统 CNN 方法,在全局信息建模和细节恢复上均有显著提升。



