1. WTConv 想解决的到底是什么问题

很多视觉模型都会遇到一个老问题:

卷积的感受野不够大。

最直接的想法当然是把卷积核做大,比如从 3x35x5 一路往上加。但问题也很明显:

  • 参数会涨
  • 计算会变重
  • 大核卷积不一定真的稳定好训

WTConv 的思路比较巧,它不是硬把卷积核继续做大,而是先把特征图做一次小波分解,拆成不同频段,再在这些子带上做小卷积,最后再还原回来。

这件事的意义在于:

  • 低频部分能更早看到更大范围的结构
  • 高频部分还能保留细节
  • 你不需要直接上一个非常夸张的大核

论文里有一句很关键的话,可以直接概括它的价值:

WTConv 是给 depth-wise convolution 做的 drop-in replacement。

也就是说,如果你的网络里本来就有 depthwise conv,它是有机会直接接进去的,而不是必须重写整套 backbone。

2. 官方 GitHub 是怎么给你用的

官方仓库 README 里最核心的用法其实非常短:

from wtconv import WTConv2d

conv_dw = WTConv2d(32, 32, kernel_size=5, wt_levels=3)

这段代码好就好在它非常像普通卷积层的调用方式。你只需要多理解一个参数:

  • wt_levels

它表示做几层 wavelet decomposition,也就是小波分解的层数。

可以先这么粗暴理解:

  • wt_levels=1:看一层
  • wt_levels=2:看得更远
  • wt_levels=3:看得更远,但结构也更复杂一点

如果你原来网络里已经有类似这样的代码:

self.dwconv = nn.Conv2d(dim, dim, kernel_size=5, padding=2, groups=dim)

那 WTConv 的改法通常就是:

self.dwconv = WTConv2d(dim, dim, kernel_size=5, wt_levels=3)

这就是它最吸引人的地方:
不是“全网重构”,而是“局部替换”。

3. 先看图:WTConv 到底在模块内部做了什么

在这里插入图片描述

图里可以把 WTConv 的内部逻辑理解成 4 步:

  1. 输入特征图
  2. 做小波分解,拆成 LL / LH / HL / HH
  3. 在这些不同频段上做小卷积
  4. 再做逆小波变换,把结果拼回空间域

如果你以前没接触过 wavelet,也没关系。这里你只需要先记住一件事:

  • LL 更像大结构、低频信息
  • LH / HL / HH 更像边缘、纹理、方向性细节

4. 为什么这个模块会让感受野变大

这其实是 WTConv 最值得写的一点。

论文里强调的是:
通过 wavelet decomposition,卷积是在不同分辨率、不同频率带上做的。这样即便你用的仍然是小卷积核,它实际能“看到”的范围会更大。

我专门画了一张很直观的图:
在这里插入图片描述

如果拿一个 5x5 的 depthwise conv 当参照:

  • 普通 DWConv 5x5,感受野大致还是 5
  • WTConv L1 可以到 10
  • WTConv L2 可以到 20
  • WTConv L3 可以到 40

这也是 WTConv 文章好写的原因之一:

它不是单纯换一个 block 名字,而是有非常明确的结构收益。

5. 再看一张图:wavelet 分解后到底拆出了什么

在这里插入图片描述

这张图最值得记住的是:

  • LL 保留了大轮廓和主要结构
  • LH 更像水平方向的细节变化
  • HL 更像垂直方向的细节变化
  • HH 更像对角、高频、纹理性残差信息

这就意味着 WTConv 不是把所有信息混着卷,而是先分频,再处理,再还原。

这比很多“注意力分数一乘完就结束”的模块,更有一种信号处理上的解释感。

6. 我自己写了一个教学版最小例子

为了把这个模块讲清楚,我没有直接把官方仓库整份源码贴过来,而是自己写了一个教学版 MiniWTConv2d

这个版本不追求和官方实现逐行一致,但非常适合解释 WTConv 的核心逻辑:

class MiniWTConv2d(nn.Module):
    def __init__(self, channels: int, kernel_size: int = 5):
        super().__init__()
        padding = kernel_size // 2
        self.base_dw = nn.Conv2d(channels, channels, kernel_size, padding=padding, groups=channels)
        self.subband_dw = nn.Conv2d(channels * 4, channels * 4, kernel_size, padding=padding, groups=channels * 4)

    def forward(self, x: torch.Tensor):
        b, c, _, _ = x.shape
        filters = haar_analysis_filters(x.device, x.dtype).unsqueeze(1).repeat(c, 1, 1, 1)
        wave = F.conv2d(x, filters, stride=2, groups=c)
        subbands = wave.view(b, c, 4, wave.shape[-2], wave.shape[-1])
        merged = subbands.view(b, c * 4, wave.shape[-2], wave.shape[-1])
        merged = self.subband_dw(merged)
        merged = merged.view(b, c, 4, wave.shape[-2], wave.shape[-1])
        recon = F.conv_transpose2d(merged.view(b, c * 4, wave.shape[-2], wave.shape[-1]), filters, stride=2, groups=c)
        out = self.base_dw(x) + recon
        return out, subbands

这段代码怎么读

先别急着全读,抓 4 行就够了。

第 1 段:先准备基础 depthwise conv
self.base_dw = nn.Conv2d(channels, channels, kernel_size, padding=padding, groups=channels)

这一层的作用很像“保底分支”。

  • 它保留了原始空间域上的 depthwise conv
  • 不让整个模块只依赖 wavelet 分支
第 2 段:把输入拆成 4 个子带
wave = F.conv2d(x, filters, stride=2, groups=c)
subbands = wave.view(b, c, 4, wave.shape[-2], wave.shape[-1])

这里做的事情就是:

  • 用固定 Haar filter 做分解
  • 把每个通道拆成 4
  • 得到 LL / LH / HL / HH

如果输入是:

(B, C, 32, 32)

那拆完以后会变成:

(B, C, 4, 16, 16)

也就是:

  • 通道数没丢
  • 多了一个 4,代表四个子带
  • 空间尺寸减半
第 3 段:对子带做 depthwise conv
merged = subbands.view(b, c * 4, wave.shape[-2], wave.shape[-1])
merged = self.subband_dw(merged)

这一段可以把它理解成:

  • 先把四个子带展开成 4C
  • 然后在 wavelet 域上做 depthwise conv

这一步是 WTConv 的核心,因为卷积不再只发生在原始空间图上,而是发生在不同频率成分上。

第 4 段:再还原回去
recon = F.conv_transpose2d(..., stride=2, groups=c)
out = self.base_dw(x) + recon

这里做两件事:

  • 用逆变换把 wavelet 域结果还原回空间域
  • 再和原始 depthwise conv 分支做融合

所以 WTConv 不是“替代空间卷积”,更像是:

空间卷积 + wavelet 域卷积 的结合。

7. 最小例子跑出来是什么样

我把这个教学版模块实际跑了一次,输入输出形状如下:

Input shape : (2, 8, 32, 32)
Subband shape: (2, 8, 4, 16, 16) -> (B, C, 4, H/2, W/2)
Output shape: (2, 8, 32, 32)
Output mean : 0.0448
Output std  : 0.7732

这个输出至少说明两件事:

  • 它确实做了“拆分 -> 处理 -> 还原”
  • 最终输出 shape 还是能无缝接回主干网络的

这也是为什么我会把它归到“即插即用模块”这一类,而不是“只能在论文结构里存在”的定制层。

8. 你在自己模型里该怎么插

如果你手上是 ConvNeXt、轻量 CNN、或者一些本身就大量使用 depthwise conv 的 backbone,WTConv 的接法通常最自然。

最粗暴也最常见的改法就是:

原始版本

self.dwconv = nn.Conv2d(dim, dim, kernel_size=7, padding=3, groups=dim)

替换后

from wtconv import WTConv2d

self.dwconv = WTConv2d(dim, dim, kernel_size=5, wt_levels=3)

这里有一个实践建议:

  • 不要一上来全网乱换
  • 先只替换 backbone 里一类典型 depthwise conv
  • 先固定 wt_levels=12
  • 跑通之后再试 3

因为 WTConv 的价值主要在“扩大感受野”,所以它更适合那些:

  • 本来就依赖局部卷积堆叠
  • 想保留 CNN 结构
  • 但又想补一点大范围建模能力

9. 这个模块适合写成 CSDN 文章的原因

说实话,GitHub 上很多“新模块”不适合写文章,原因很简单:

  • 仓库太乱
  • 原理太绕
  • 用法不直接
  • 读者看完也不知道该往哪塞

WTConv 比较例外,它有几个很明显的内容优势:

  • 来源清楚:官方 GitHub + ECCV 2024
  • 主题明确:大感受野卷积
  • 接法明确:替换 depthwise conv
  • 图很好画:流程图、子带图、感受野图都能讲明白
  • 代码好写:哪怕不照搬官方实现,也能写出教学版例子

所以它特别适合写成这种结构:

  1. GitHub 上发现了一个新模块
  2. 它不是注意力,而是 wavelet + conv
  3. 为什么它能扩大感受野
  4. 代码该怎么接
  5. 我自己写了个最小例子验证逻辑

这就不只是“介绍论文”,而是更像一篇真正能让别人上手改模型的技术文章。

10. 我觉得 WTConv 的优势和问题都很明显

优势

  • 概念新,但不悬浮
  • 真正可以替换现有 depthwise conv
  • 对“大感受野”这个问题给出了一个很清楚的结构答案
  • 相比很多注意力模块,它更容易讲出信号处理层面的动机

问题

  • 它不是零成本模块,结构上还是比普通卷积复杂
  • 对不熟悉 wavelet 的读者,第一次看会有门槛
  • 真正接进项目时,训练稳定性、速度和收益还得自己测
  • 不是所有 backbone 都像 ConvNeXt 那样适合直接替换

所以我的建议不是“看到新模块就全量替换”,而是:

把它当成一个值得试的结构增强件,而不是万能药。

附:我这次实际参考的 GitHub / 论文

Logo

汇聚全球AI编程工具,助力开发者即刻编程。

更多推荐