
Spherical CNN 这个方向光看论文很容易一头雾水。球面谐波、SO(3)群表示、Wigner-D矩阵、旋转等变这些名词堆在一起数学功底再好的人也会被绕晕。我的经验是想真正掌握它必须把源代码打开一行一行搞清楚数据是怎么流转、矩阵是怎么变换、梯度是怎么反传的。这篇文章我就顺着自己的调试经历把 Spherical CNN 的源码从整体架构到核心模块拆开聊希望能给正在啃这个方向的朋友省点时间。1. 先搞懂它在解决什么问题平面卷积在球面上为什么失效1.1 等矩形投影里的“极点灾难”做全景图或者球面数据处理的人对 equirectangular 投影一定不陌生。就是把球面按经纬度展开成一张矩形图横轴是经度纵轴是纬度。这种表示方式很直观存储也简单一个二维数组就能装下整个球面信号。但问题在于这种投影在靠近两极的区域会产生严重的畸变。赤道附近一个 1 度的格子是差不多正方形的到了纬度 80 度的地方同一个经度跨度对应的实际弧长已经缩小到赤道处的六分之一不到到了极点整个纬线圈成了一个点。如果用普通卷积在这个图上滑窗卷积核在赤道看到的是一片局部的真实内容在极区看到的却是被横向拉伸了几倍的扭曲内容。同一个卷积核在不同位置提取的特征物理含义完全对不上。更麻烦的是旋转等变性在这里彻底丢失了。你在球面上把一个物体旋转 45 度反映到等矩形投影图上内容不是简单地平移 45 个像素而是一种非常复杂的非线性变形。普通 CNN 对平移是等变的对旋转完全不买账。所以处理全景图、卫星遥感数据、分子表面电势这类球面数据时直接套用平面 CNN效果会打折扣而且对输入朝向非常敏感。1.2 旋转等变是怎么被引入的群卷积的思想Spherical CNN 的核心思路是把“对球面做卷积”这件事本身重新定义。它要保证的是这么一条性质如果输入信号被旋转了某个角度那么网络中间层的特征图也应该被旋转同样的角度而不是发生形变。这个性质在数学上叫旋转等变性。为了实现它论文里引入的是群卷积的思路。普通卷积的核心是在图像上滑动一个卷积核计算局部加权和群卷积则更进一步不仅要在球面的每个位置上计算卷积响应还要在所有的旋转方向上计算响应。也就是说卷积结果不再是一个球面信号而是在 SO(3) 旋转群上的信号。这样做的代价是特征图的维度变大了好处是网络对输入朝向完全不敏感分类和识别任务能获得很强的泛化能力。1.3 源码里绕不开的三座大山SH、SO(3)、Wigner-D真正打开源码之后你会发现所有复杂的数学都浓缩在了三样东西里。第一样是球面谐波Spherical Harmonics简称 SH。它是定义在球面上的一组正交基函数作用类似傅里叶变换里的正弦余弦波。任何一个球面信号都可以分解成一系列 SH 系数的加权和。源码里会有一个模块专门做球面谐波变换SHT把空间域的球面信号变换到频域。第二样是 SO(3) 群本身。球面上的旋转操作构成的数学群体源码里会生成大量旋转矩阵用来实现对特征图的旋转变换。第三样是 Wigner-D 矩阵。这是 SO(3) 群上的“傅里叶基底”用于处理旋转群上的信号变换。球面卷积在频域实现时依赖的就是 Wigner-D 矩阵的分块对角结构。理解这一点是读懂所有核心代码的钥匙。一句话总结球面卷积在空间域做很麻烦所以源码几乎都在频域里操作——先 SHT 到 SH 系数空间然后利用 Wigner-D 矩阵的分块结构做卷积计算最后再逆变换回来。2. 源码结构总览一个典型的 Spherical CNN 仓库里都有什么2.1 目录结构与核心文件分工我读过几个 Spherical CNN 的开源实现包括论文作者的原始版本和一些社区复刻版。虽然框架不同有的是 TensorFlow 1.x有的是 PyTorch但文件组织方式高度相似。一个典型的仓库结构大致是这样的spherical_cnn/ ├── sh_transform/ # 球面谐波变换相关 │ ├── sht.py # 前向/逆向 SHT 核心实现 │ ├── legendre.py # 连带勒让德多项式计算 │ └── sampling.py # 球面采样网格生成 ├── layers/ │ ├── conv.py # 球面卷积层 / SO(3) 卷积层 │ ├── nonlinearity.py # 等变非线性层 │ ├── pooling.py # 球面池化层 │ └── normalization.py # 归一化层 ├── models/ │ └── sp_cnn.py # 网络组装 ├── data/ │ └── spherical_dataset.py # 球面数据加载与预处理 ├── train.py # 训练脚本 └── utils/ ├── rotation.py # 旋转矩阵 / Wigner-D 生成 └── io_utils.py这个分工很有讲究。SH 变换和旋转矩阵生成属于底层基础模块它们不依赖任何深度学习框架纯粹的 NumPy/SciPy 就能搞定。卷积层、非线性层、池化层是网络的核心组件负责把频域操作封装成可训练的层。模型文件把各层串起来训练脚本负责数据加载和反向传播。2.2 数据流与张量尺寸变化一张图看懂代码走向我强烈建议你拿到源码后第一步不是读代码而是把张量的形状变化画出来。我在最初读代码时没做这一步结果在 SH 系数维度和旋转维度上来回迷路白白浪费了好几天。在一个典型的图像分类任务里输入是[Batch, Height, Width, Channels]其中 Height 和 Width 对应球面的纬度采样数和经度采样数。进入网络后张量经历的变化大致是空间域 → 频域[B, H, W, C]通过 SHT 变成[B, Bandwidth, Bandwidth, C]。这里的 Bandwidth通常用 L 表示是截断的球面谐波频带数。注意这里的维度是复数域实际存储时可能拆成实部和虚部或者用复数张量表示。频域卷积在 SH 系数空间做卷积操作特征图的尺寸不会改变仍然是[B, L, L, C_out]。这是频域卷积的优势——不需要 padding也没有 stride 的概念。非线性激活频域不能直接做 ReLU因为频域里的逐点 ReLU 会破坏等变性。源码的做法一般是逆变换回空间域做 ReLU再变换回频域。这一步会产生额外的 SHT 开销但数学上是必须的。池化通过降低带宽来实现。[B, L, L, C]变成[B, L, L, C]其中 L L。相当于丢弃高频的 SH 系数保留低频成分。分类层最后通常是 global pooling在频域上对所有的系数做某种聚合得到一个固定长度的向量接全连接层和 softmax。![数据流不可见请以文字描述为准]这个流程中的数据形状变化是整个源码的主心骨。你把这个理清了后面读任何一段代码都能快速定位它处于整个流水线的哪个环节。2.3 框架选择与版本兼容劝退无数人的隐形坑Spherical CNN 这篇 ICLR 2018 论文的原始代码是 TensorFlow 1.x 写的。我只能说如果你想直接跑原始仓库先把心理预期放低。TF 1.x 的 session 机制、tf.contrib模块、各种隐式全局变量在今天的深度学习中已经非常不顺手。我试过在 Python 3.8 和 TF 2.10 里跑原版代码光是把tf.contrib替换掉就花了一个晚上后面还有一堆tf.compat.v1的兼容补丁。如果只是学习源码逻辑我建议看社区里用 PyTorch 重写的实现。PyTorch 对复数的支持相对清爽自动求导机制也比 TF 1.x 的静态图直观很多。但要注意PyTorch 1.8 之后的复数张量 API 有过一次不小的调整老代码里用torch.complex64创建张量的写法在新版本上可能会警告甚至报错。另外SH 变换的实现通常依赖scipy.special.lpmn或者numpy.polynomial.legendre计算连带勒让德多项式。这几个函数的参数含义略有不同而且在高阶数比如带宽 64 以上时数值稳定性会下降调试时容易在隐蔽的地方翻车。3. 核心代码模块逐段拆解从采样到等变卷积3.1 球面采样网格等矩形 vs HEALPix代码里是怎么选的球面信号在计算机里必须以离散形式存储。最直接的方式是均匀经度/纬度网格也就是等矩形网格。在源码里通常会用两个一维数组theta np.linspace(0, np.pi, num_theta)和phi np.linspace(0, 2*np.pi, num_phi)来生成网格坐标。这个方式简单但存在两个问题。一是极点处过采样严重二是从经纬网格构造球面谐波系数时数值积分的权重不均匀。许多源码会在积分时用sin(theta)作为权重来补偿但这只能缓解问题不能根治。另外一个常见的采样方案是 HEALPix。它的核心思想是把球面分成面积相等的一系列像素每个像素的形状接近正方形极点区域的畸变被有效控制。Healpix 在天体物理领域用得非常多所以有成熟的开源库。如果你的网络要处理的是全天空微波背景辐射这类数据强烈建议直接用 HEALPix而不是等矩形网格。在源码层面HEALPix 的坐标生成通常会调用healpy库的接口。但要注意HEALPix 网格与 SH 变换的配合没有等矩形网格那么直接需要额外的索引映射和重采样操作。原作者代码里没有用 HEALPix而是用了等矩形网格加加权积分我觉得纯粹是出于实现便利的考虑并非数学上最优。3.2 球面谐波变换SHT的实现代码里的分步计算SHT 是整个 Spherical CNN 地基。看不懂它后面所有层都是空中楼阁。从数学定义说球面信号 ( f(\theta, \phi) ) 的 SH 系数 ( \hat{f}_{l}^{m} ) 等于信号在某个球面谐波基函数上的投影公式是[ \hat{f}_l^m \int_0^{2\pi} \int_0^\pi f(\theta, \phi) , Y_l^{m*}(\theta, \phi) , \sin\theta , d\theta , d\phi ]其中 ( Y_l^m ) 是球面谐波函数它本身可以拆成两部分关于经度的复指数函数 ( e^{im\phi} ) 和关于纬度的连带勒让德多项式 ( P_l^m(\cos\theta) )。这个可分离结构是源码实现的关键。在实际代码里SHT 通常分两步完成第一步沿经度方向做 FFT。因为 ( e^{imphi} ) 本质上就是傅里叶基底所以对球面信号的每一根纬线固定 theta变化 phi 的一维数组做 FFT就能得到中间结果。这一步用numpy.fft.fft就行速度很快。第二步沿纬度方向做勒让德变换。把第一步的结果和连带勒让德多项式做内积得到最终的 SH 系数。这一步源码里通常用scipy.special.lpmn生成勒让德多项式再通过矩阵乘法完成投影。从源码的角度看SHT 的实现通常被封装成一个函数输入是空间域的采样数据输出是 SH 系数张量。需要注意几个细节带宽截断SH 系数的数量不是无限的源码会让你指定一个bandwidth通常记为 ( L )只有 ( l L ) 的系数被保留下来。( L ) 决定了频率分辨率的最高限度也直接决定了系数的数量。系数顺序不同源码对 SH 系数的存储顺序约定不同。有的是[l, m]顺序有的是[m, l]有的把实数形式和复数形式混用。我踩过的坑就是两个模块之间系数顺序不一致逆变换出来的结果完全错乱。读源码时第一时间确认这个约定。前向和逆向的归一化有些实现把归一化因子放在前向变换里有些放在逆向里有些用 Parseval 定理来校准。这直接关系到网络训练时能量守恒和梯度幅值的正确性。3.3 球面卷积层的频域实现从代码看卷积为什么变成了“乘法”如果是第一次读球面卷积层可能会有个疑问为什么代码里没有卷积操作全是矩阵乘法和张量重塑这其实是频域卷积的典型特征。根据卷积定理空间域的卷积等价于频域的逐元素乘法。球面卷积也不例外。在源码里球面卷积层做的事情本质上是把输入信号通过 SHT 变到频域得到系数 ( \hat{f}_l^m )。把卷积核也变到频域得到 ( \hat{h}_l )注意球面卷积核通常只依赖于 ( l )不依赖于 ( m )这是旋转等变带来的约束。在频域做乘法( \hat{g}_l^m \hat{f}_l^m \cdot \hat{h}_l )。可选地做逆 SHT回到空间域进行后续处理。在代码层面conv_sphere的核心逻辑通常就几行。关键在第三步那个乘法因为系数的维度是 ( [L, L] )对应 ( l ) 和 ( m )而卷积核是 ( [L] ) 的向量源码会利用 NumPy 或 PyTorch 的 broadcasting 机制让每个 ( \hat{h}_l ) 自动乘到所有 ( m ) 系数上。而到了 SO(3) 卷积层逻辑更复杂一些。它处理的不再是球面信号而是定义在旋转群上的信号。每个特征图都有多个 ( l ) 分量每个分量是一个 ( (2l1) \times (2l1) ) 的矩阵。卷积操作变成了矩阵乘法涉及 Wigner-D 矩阵的分块结构。源码里这个操作往往是通过einsum来实现的例如torch.einsum(bcij,ijco-bco, feature, kernel)这种形式。einsum的好处是简洁坏处是阅读门槛高调试时很难看清楚维度是怎么对齐的。3.4 等变非线性层ReLU 为什么要在空间域做大多数人刚开始读源码时会忽略非线性层觉得不就是个 ReLU 吗但在 Spherical CNN 里非线性是最容易出错的一个环节。原因很简单频域里的逐点 ReLU 会破坏旋转等变性。因为 ReLU 是逐点运算在频域的“点”对应的是系数而不是球面上的空间位置。对系数做 ReLU 没有明确的几何意义而且会破坏系数之间的线性关系导致旋转后结果不一致。源码里的正确做法是对频域特征做逆 SHT回到空间域。在空间域做 ReLU。再做前向 SHT回到频域。这个操作的代价很高因为 SHT 本身就是最耗时的计算。所以源码里大部分实现会在内存中缓存正向和逆向变换的矩阵避免重复计算。值得一提的是有些进阶源码会使用norm ReLU或者gated nonlinearity来减少性能损耗。前者对每个频带系数的范数做激活后者额外学一个门控系数。这些方案在空间域和频域之间只需部分切换能显著提升训练速度但源码复杂度会上升一个量级。3.5 旋转等变的实现Wigner-D 矩阵在代码里的真实面目很多人读到这里就卡住了。Wigner-D 矩阵是 ( (2l1) \times (2l1) ) 的不可约表示矩阵它描述了在某个旋转操作下频域系数如何变换。源码里Wigner-D 矩阵的生成通常用的是scipy.special里的sph_harm或者专用库lie_learn人话就是对每个频带 ( l )给定一组旋转参数例如欧拉角生成一个矩阵 ( D^l(\alpha,\beta,\gamma) )。在代码使用层面最关键的是矩阵乘法怎么和特征张量的维度对齐。假设特征张量在频域的形状是[B, L, L, C]其中第一维的 L 对应 m或 l第二维的 L 对应 index旋转操作实际上是选取一个旋转矩阵然后对每个通道分别做矩阵乘法。我在源码里见过两种主要实现方式显式矩阵乘法每次旋转时构造一个[B, (2l1), (2l1)]的 Wigner-D 矩阵然后用matmul作用在特征上。这种方式适合小规模数据容易理解。对旋转集合做批处理在等变网络中要同时计算多个旋转下的响应就把所有旋转矩阵拼成一个更大的张量用一次批量矩阵乘法完成。这种方式节省循环开销但内存占用大而且代码可读性差。调试时的经验是先用一个很小的网络、一个固定的旋转角度手算或手推一遍特征经过旋转之后的变化再看看代码输出是否吻合。如果吻合再相信代码否则别急着继续往下走。3.6 池化层与降带宽操作View 一下就把频率砍了Spherical CNN 里的池化层实现极其简单简单到你以为自己看错了。它做的只是把 SH 系数的后部分截掉只保留低频分量。假设当前带宽是 ( L32 )特征维度是[B, 32, 32, C]池化到 ( L16 ) 的话代码里会这样处理pooled feature[..., :16, :, :]就这。把 m 维度超过 16 的系数全部扔掉。这样操作在频域完全合理因为低频分量保留的是信号的主体结构高频分量对应细节。这种池化方式没有参数不会产生额外开销还天然保证了等变性。不过有个坑要注意降带宽之后后续层的卷积核尺寸、归一化参数、Wigner-D 矩阵的维度都要跟着调整否则张量维度不匹配。源码里通常会把带宽作为网络结构的一个全局参数传入确保各个层步调一致。你在修改或移植代码时务必检查每一处用到带宽的地方。4. 复现与调试中的踩坑记录4.1 带宽选择小了丢信息大了炸显存我之前有一次用 Spherical CNN 做全景场景分类数据集是 512x256 的等矩形图。最初我按照论文推荐的带宽 ( L64 ) 来跑结果显存直接爆掉。后来我把带宽降到 ( L32 )模型精度掉了不少因为高频纹理信息对场景分类很重要。我自己实际测试下来给出一个经验范围输入分辨率推荐带宽 L说明64x3216够用主要保住低频结构128x6416~24平衡精度与算力256x12824~32常用场景效果不错512x25632~48高精度要求才考虑显存压力大带宽 ( L ) 近似对应空间分辨率分辨率越高能支撑的最高频率也越高。但带宽增大会导致 SH 系数数量以 ( O(L^2) ) 速度增长中间层特征图的体积增长更快。源码在初始化自定义网络时建议先用小带宽跑通流程再用大带宽做最终实验。4.2 数值精度单精度还是双精度这是一个非常细节但极其致命的问题。SHT 涉及大量的三角学和勒让德多项式计算在高频段大 l 值会出现严重的数值不稳定。我自己调试时发现在带宽 ( L48 ) 的情况下如果全程用 float32逆变换回来的空间域信号会有明显的噪声把关键计算部分改成 float64 后误差小了将近 3 个数量级。但 full precision 带来的问题是显存翻倍、训练速度下降。源码里常见的折中方案是在 SHT 模块用 float64 计算积分矩阵或者说变换矩阵计算完成后将其缓存转成 float32 供训练使用。隐藏的好处是梯度反传时使用缓存的矩阵不再涉及勒让德多项式的重新计算速度快很多。4.3 等变性测试怎么验证源码真的“等变”我在调试 Spherical CNN 的过程中一半时间都花在自检上。最有效的自检方法是随机选一个旋转角度然后做两步操作先旋转输入信号再通过网络前向传播先通过网络前向传播再旋转输出的特征图。然后对比两个结果是否一致。如果网络是完全等变的两个结果应该完全一样或者数值上误差在阈值内。不一致的地方在哪一段出现就说明问题出在哪一段。我亲测这个方法能快速定位 bug。比如之前我在测试时发现第二阶段和第一阶段的结果总是差一点查了半天发现是 Wigner-D 矩阵的实现把共轭转置搞反了。4.4 从 TF1 代码迁移到 PyTorch 时的移植要点如果你拿到的是原版 TF1 代码想迁到 PyTorch几个地方需要注意复数张量的处理TF1 时代常用两个实张量拼在一起表示复数PyTorch 从 1.8 开始原生支持复数张量。重建代码时可以直接用torch.complex64但要注意旧代码里的real和imag调用方式需要全部替换。自定义梯度SHT 模块如果用scipy计算变换矩阵在 PyTorch 里需要注册torch.autograd.Function在forward里执行矩阵乘法在backward里利用 SHT 矩阵的性质逆变换矩阵是前向变换矩阵的伪逆或转置计算梯度。随机旋转数据增强原版代码里生成随机旋转矩阵用的是lie_learn库的接口迁移时可以直接用scipy.spatial.transform.Rotation.random()然后转成欧拉角再喂给 Wigner-D 生成函数。4.5 高频噪声与能量泄漏做 SH 变换时如果输入信号在空间域不是严格带限的那么高频成分会在重建时泄漏到低频系数中产生类似频谱混叠的问题。源码层面没有自动处理这个需要你自己在信号进入网络前做低通滤波。我之前在实验里发现训练损失不稳定从第四五个 epoch 开始震荡。排查后发现是我的输入信号在切到等矩形网格时产生了很陡的边缘导致高频能量极强SH 系数里出现了明显的吉布斯现象。后来我在预处理环节加入高斯滤波在极点区域做平滑过渡问题马上解决了。这说明输入预处理对 Spherical CNN 的影响有时比网络结构本身还大。5. 源码改造与扩展从分类到更复杂的任务读懂了基础源码之后很有必要尝试按自己的任务需求改造它。这里分享几个常见的方向和源码层面的改动技巧。5.1 把球面卷积用到非图像数据上点云和分子势能面Spherical CNN 不只能处理全景图。源代码里的 SHT 模块、卷积层、等变层是完全独立于具体数据的。我近期在一个项目里处理点云法向量方向的概率分布用的就是把点云法向量统计成球面信号再输入到这个网络里。这个场景下唯一需要改的是数据预处理部分把点云法向量方向通过某种核密度估计离散化到球面网格上剩下的一切都不用动。分子势能面也是经典应用。分子结构可以用一组原子坐标表示配分函数或者能量在旋转下不变这就天然适合用等变网络处理。把原子坐标投影到径向函数加球面谐波展开作为网络输入是目前一个非常火的方向。5.2 换成非等变的后续任务把 Spherical CNN 当作特征提取器很多情况下我们的任务端点不是旋转等变的。比如全景图里的物体检测需要知道物体在哪个位置这本质上要求网络输出的特征对旋转保持某种“位置敏感”的信息。这时可以把 Spherical CNN 前几层作为特征提取器把频域系数作为特征向量然后在最后的任务头里引入一些非等变操作。常见的做法是把频域特征逆变换回空间域在空间域做检测。这样既享受了前面等变层的稳定性又不会受限于等变约束导致的空间定位能力不足。源码改造的核心通常就是调整网络最后一层的输出形式把频域张量换成空间域张量。5.3 新的旋转群与多尺度框架等变卷积的思想完全可以沿用到其他旋转群上。比如在二维平面上旋转对应的是 SE(2) 群上的等变卷积。源码层面可以复用的东西很多Wigner-D 矩阵的生成逻辑、频域乘法的实现框架、等变非线性的处理方式都是抽象出来的公共模块。多尺度框架也是可行的扩展方向。目前 SH 变换返回的是全频带的系数而不同频带其实对应不同尺度的几何信息。在代码中你可以在频带维度上做切片把低频与高频分别输入到不同分支的网络然后把它们在后续层拼接或相加。这种操作由于 SH 系数的可分割性实施起来比普通 CNN 的特征金字塔还要简单——你不需要任何下采样操作直接切维度就行。5.4 一个提高训练稳定性的小技巧最后说一个小技巧是我在调参时发现的。在频域卷积层之后对系数做 LayerNorm 或 InstanceNorm 的时候要注意归一化所在维度。因为频域系数的每个通道有实部和虚部而 Wigner-D 的不变性要求归一化的操作不能破坏频带之间的相对关系。所以通常做法是每个样本、每个频带独立做归一化均值在带宽内计算。这在 PyTorch 里用一次torch.view加torch.mean(dim..., keepdimTrue)就能搞定但方向搞错的话等变性就悄悄没了。我的实际经验是在调试任何等变网络时都保持一个最小复现脚本在手边这个脚本里有一个固定的随机种子、一个固定的旋转角度和一个能肉眼可判断的空间信号。每次改动代码之后就跑一遍检查等变性是否保持。这个脚本能救你无数次。6. 一些值得看的参考实现如果你打算亲自上手除了读论文原文下面几个仓库值得优先看论文作者的原始实现在 GitHub 上搜Spherical CNN能找到 ICLR 2018 作者的官方代码。虽然是 TF1 版但数学注释非常详细特别适合理解核心思想。我个人不推荐直接拿它训练但读代码非常友好。PyTorch 复刻版社区里有一个比较完整的 PyTorch 版文件组织清晰适合在此基础上改自己的代码。它把 SHT 和 Wigner-D 分别封装成了独立的autograd.Function这一点对我启发很大——把底层变换和上层网络解耦调试时能迅速隔离问题。lie_learn库这是一个专门做李群学习相关的工具库里面包含了大量球面谐波、Wigner-D 矩阵、旋转群采样的实现。很多 Spherical CNN 的源码都直接依赖它。如果你要实现自己的等变层这个库的文档值得仔细看。再提醒一点任何仓库跑通之后第一时间做的事情是“破坏性测试”改一个维度看报错信息是否准确指出问题所在。等变网络的张量维度复杂如果你不能快速从报错中定位问题后面改代码的效率会非常低。我自己啃 Spherical CNN 源码最大的体会是数学公式和代码实现之间存在一大片“隐含知识”。公式写的是连续的积分变换代码里全是一堆矩阵乘法和维度变换。只有把每一步操作对应到公式的哪一个部分才能真正弄懂这个网络在做什么。希望这篇解析能帮你少走一些弯路。