WTFD:基于小波变换与Transformer的多尺度特征提取技术
1. 项目概述
今天要跟大家分享的是我们团队最新发表在TGRS 2024上的工作——WTFD(Wavelet-based Transformer for Feature Distillation),一个基于小波变换和Transformer的多尺度特征提取模块。这个模块最大的特点就是能够同时捕捉低频全局信息和高频细节特征,在各种视觉任务上都实现了显著的性能提升。
在实际应用中我们发现,传统的特征提取方法往往存在一个两难选择:要么过于关注全局特征而丢失细节,要么陷入局部细节而忽略整体结构。WTFD通过创新的多尺度特征提取和增强机制,完美解决了这个问题。经过在多个公开数据集上的测试,我们的方法在分类、检测、分割等任务上都能带来1.5%-3.2%的准确率提升,而且计算开销增加非常有限。
2. 核心设计思路
2.1 多尺度特征提取的必要性
在计算机视觉任务中,不同层次的特征对最终性能的影响是不同的。低频分量通常包含图像的全局结构和主体信息,而高频分量则记录了边缘、纹理等细节特征。传统CNN通过堆叠卷积层来隐式地学习这些特征,但这种方式的特征提取是黑箱的,缺乏明确的控制机制。
我们通过大量实验发现,在复杂场景下(如遥感图像分析、医学图像处理等),单纯依赖CNN提取的特征往往会出现以下问题:
- 对小物体或细节特征捕捉不足
- 对光照、尺度变化敏感
- 特征表示缺乏明确的物理意义
2.2 小波变换的优势
WTFD选择小波变换作为基础工具,主要基于以下几个考量:
- 时频局部化特性:可以同时分析信号的时域和频域特征
- 多分辨率分析:通过不同尺度的小波基函数,可以自然地提取多尺度特征
- 计算效率:离散小波变换(DWT)的计算复杂度仅为O(n),非常适合嵌入到深度学习模型中
我们采用了Haar小波作为基础变换核,因为它的计算最简单,而且已经证明在深度学习模型中表现良好。具体实现时,我们对输入特征图进行二维DWT分解,得到LL(低频)、LH(水平高频)、HL(垂直高频)和HH(对角高频)四个子带。
2.3 Transformer的引入
单纯使用小波变换虽然可以分离不同频段的特征,但如何有效利用这些特征仍然是个挑战。WTFD创新性地引入了Transformer机制来处理多尺度特征:
- 低频通路:使用轻量级Transformer处理LL子带,捕捉全局依赖关系
- 高频通路:设计了一个交叉注意力模块,让三个高频子带(LH、HL、HH)可以互相增强
- 特征融合:最后通过逆小波变换(IWT)将处理后的各子带特征重新组合
这种设计有以下几个优势:
- 明确区分了不同频段特征的处理方式
- 通过注意力机制实现了跨尺度的特征交互
- 保持了特征的物理可解释性
3. 模块实现细节
3.1 整体架构
WTFD模块的完整处理流程如下:
- 输入特征图X ∈ R^(H×W×C)
- 进行DWT分解,得到四个子带:
- LL ∈ R^(H/2×W/2×C)
- LH ∈ R^(H/2×W/2×C)
- HL ∈ R^(H/2×W/2×C)
- HH ∈ R^(H/2×W/2×C)
- 低频处理通路:
- LL通过一个轻量Transformer块
- 输出增强后的LL'
- 高频处理通路:
- LH、HL、HH通过交叉注意力模块
- 输出增强后的LH'、HL'、HH'
- 进行IWT重构,得到最终输出特征图Y ∈ R^(H×W×C)
3.2 关键组件实现
3.2.1 轻量Transformer设计
为了降低计算成本,我们对标准Transformer做了以下优化:
- 使用分组自注意力:将通道分成4组分别计算注意力,然后拼接
- 采用跨步卷积进行token混合,替代标准的全连接层
- 位置编码使用可学习的相对位置偏置
具体实现代码如下(PyTorch版本):
class LightweightTransformer(nn.Module): def __init__(self, dim, num_heads=4, groups=4): super().__init__() self.norm = nn.LayerNorm(dim) self.attn = GroupedSelfAttention(dim, num_heads, groups) self.conv = nn.Conv2d(dim, dim, kernel_size=3, stride=1, padding=1) def forward(self, x): B, C, H, W = x.shape x = x.flatten(2).transpose(1, 2) # B, N, C x = x + self.attn(self.norm(x)) x = x.transpose(1, 2).view(B, C, H, W) x = x + self.conv(x) return x3.2.2 高频交叉注意力模块
高频特征的处理需要特别关注不同方向特征的交互。我们设计了一个三路交叉注意力机制:
- 每个高频子带先通过一个1×1卷积进行特征变换
- 然后计算三个子带之间的交叉注意力权重
- 最后根据注意力权重进行特征融合
具体计算过程:
- 对于LH子带,它的输出是:LH' = α₁₁LH + α₁₂HL + α₁₃*HH
- 类似地计算HL'和HH'
- 其中注意力权重α通过三个子带的特征相似度计算得到
这种设计使得不同方向的高频特征可以互相增强,特别是对于那些在单一方向上不明显的边缘特征。
3.3 逆变换与特征融合
经过Transformer增强后的各子带特征需要通过逆小波变换重新组合。这里有一个关键细节:我们不是简单地进行IWT,而是引入了一个可学习的融合权重:
Y = IWT(LL' + λ₁*(LH' + HL' + HH'))
其中λ₁是一个可学习的标量参数,初始值为0.5。这种设计使得模型可以自适应地调整高频特征的贡献度。
4. 实验与性能分析
4.1 实验设置
我们在多个标准数据集上评估了WTFD的性能:
- 分类任务:ImageNet-1K
- 检测任务:COCO
- 分割任务:ADE20K
- 遥感图像分类:NWPU-RESISC45
基线模型选择了ResNet、Swin Transformer等主流架构。WTFD作为一个即插即用模块,被添加到这些模型的各个阶段之间。
4.2 主要结果
在ImageNet-1K分类任务上,WTFD带来了显著的性能提升:
| 骨干网络 | 原始top-1 | +WTFD | 提升 |
|---|---|---|---|
| ResNet-50 | 76.3% | 78.1% | +1.8% |
| Swin-T | 81.2% | 82.7% | +1.5% |
| ConvNeXt-T | 82.1% | 84.3% | +2.2% |
在COCO检测任务上,以RetinaNet为检测器:
| 骨干网络 | mAP | +WTFD | 提升 |
|---|---|---|---|
| ResNet-50 | 36.4 | 38.9 | +2.5 |
| Swin-T | 42.1 | 44.3 | +2.2 |
4.3 计算开销分析
虽然WTFD引入了额外计算,但通过精心设计,开销增加非常有限:
| 模型 | Params(M) | FLOPs(G) | +WTFD Params | +WTFD FLOPs |
|---|---|---|---|---|
| ResNet-50 | 25.5 | 4.1 | +1.2M | +0.3G |
| Swin-T | 28.3 | 4.5 | +1.5M | +0.4G |
5. 实际应用技巧
5.1 部署建议
- 插入位置:建议在网络的每个下采样阶段前插入WTFD模块
- 通道数设置:WTFD内部通道数可以设为输入通道数的1/4到1/2
- 训练策略:初始学习率可以设为骨干网络的1/2
5.2 常见问题解决
训练不稳定:
- 先固定WTFD的参数,训练骨干网络几个epoch
- 然后解冻WTFD一起训练
内存占用过高:
- 可以减少WTFD中Transformer的头数
- 或者使用梯度检查点技术
在某些数据集上效果不明显:
- 尝试调整高频特征的融合权重λ₁
- 可以增加高频通路的注意力头数
5.3 扩展应用
除了标准的视觉任务,WTFD还可以应用于:
- 医学图像分析:对CT/MRI图像的多尺度特征提取特别有效
- 视频理解:处理时空特征时,可以用时间维度的WTFD变体
- 图像生成:作为GAN中的特征提取模块,可以生成更清晰的细节
6. 模块变体与改进方向
6.1 小波基选择
除了Haar小波���我们还尝试了其他小波基:
- Daubechies小波:更平滑,但计算量稍大
- Biorthogonal小波:对称性好,适合图像处理
- 可学习小波基:端到端训练小波滤波器
实验表明,对于大多数任务,Haar小波已经足够好,且计算效率最高。
6.2 注意力机制改进
- 空间受限注意力:只计算局部窗口内的注意力,减少计算量
- 通道注意力:在频域子带之间也引入通道注意力
- 动态头数:根据输入特征复杂度自适应调整注意力头数
6.3 与其他模块的结合
- 与CNN结合:在WTFD前后加入卷积层增强局部特征
- 与MLP-Mixer结合:用WTFD替代部分MLP层
- 与知识蒸馏结合:用WTFD作为教师模型的特征提取器
在实际项目中,我们发现将WTFD插入到现有模型的浅层和中间层效果最好,既能提取多尺度特征,又不会引入过多计算负担。对于需要实时推理的场景,可以考虑使用分组数更多的轻量版WTFD。