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 ) 分量,每个分量是一个 ( (2l+1) \times (2l+1) ) 的矩阵。卷积操作变成了矩阵乘法,涉及 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 矩阵是 ( (2l+1) \times (2l+1) ) 的不可约表示矩阵,它描述了在某个旋转操作下,频域系数如何变换。源码里,Wigner-D 矩阵的生成通常用的是scipy.special里的sph_harm或者专用库lie_learn,人话就是:对每个频带 ( l ),给定一组旋转参数(例如欧拉角),生成一个矩阵 ( D^l(\alpha,\beta,\gamma) )。
在代码使用层面,最关键的是矩阵乘法怎么和特征张量的维度对齐。假设特征张量在频域的形状是[B, L, L, C],其中第一维的 L 对应 m(或 l),第二维的 L 对应 index,旋转操作实际上是选取一个旋转矩阵,然后对每个通道分别做矩阵乘法。
我在源码里见过两种主要实现方式:
显式矩阵乘法:每次旋转时构造一个
[B, (2l+1), (2l+1)]的 Wigner-D 矩阵,然后用matmul作用在特征上。这种方式适合小规模数据,容易理解。对旋转集合做批处理:在等变网络中,要同时计算多个旋转下的响应,就把所有旋转矩阵拼成一个更大的张量,用一次批量矩阵乘法完成。这种方式节省循环开销,但内存占用大,而且代码可读性差。
调试时的经验是:先用一个很小的网络、一个固定的旋转角度,手算或手推一遍特征经过旋转之后的变化,再看看代码输出是否吻合。如果吻合,再相信代码,否则别急着继续往下走。
3.6 池化层与降带宽操作:View 一下就把频率砍了
Spherical CNN 里的池化层,实现极其简单,简单到你以为自己看错了。它做的只是把 SH 系数的后部分截掉,只保留低频分量。
假设当前带宽是 ( L=32 ),特征维度是[B, 32, 32, C],池化到 ( L=16 ) 的话,代码里会这样处理:
pooled = feature[..., :16, :, :]就这。把 m 维度超过 16 的系数全部扔掉。这样操作在频域完全合理,因为低频分量保留的是信号的主体结构,高频分量对应细节。这种池化方式没有参数,不会产生额外开销,还天然保证了等变性。
不过有个坑要注意:降带宽之后,后续层的卷积核尺寸、归一化参数、Wigner-D 矩阵的维度都要跟着调整,否则张量维度不匹配。源码里通常会把带宽作为网络结构的一个全局参数传入,确保各个层步调一致。你在修改或移植代码时,务必检查每一处用到带宽的地方。
4. 复现与调试中的踩坑记录
4.1 带宽选择:小了丢信息,大了炸显存
我之前有一次用 Spherical CNN 做全景场景分类,数据集是 512x256 的等矩形图。最初我按照论文推荐的带宽 ( L=64 ) 来跑,结果显存直接爆掉。后来我把带宽降到 ( L=32 ),模型精度掉了不少,因为高频纹理信息对场景分类很重要。
我自己实际测试下来,给出一个经验范围:
| 输入分辨率 | 推荐带宽 L | 说明 |
|---|---|---|
| 64x32 | 16 | 够用,主要保住低频结构 |
| 128x64 | 16~24 | 平衡精度与算力 |
| 256x128 | 24~32 | 常用场景,效果不错 |
| 512x256 | 32~48 | 高精度要求才考虑,显存压力大 |
带宽 ( L ) 近似对应空间分辨率,分辨率越高,能支撑的最高频率也越高。但带宽增大会导致 SH 系数数量以 ( O(L^2) ) 速度增长,中间层特征图的体积增长更快。源码在初始化自定义网络时,建议先用小带宽跑通流程,再用大带宽做最终实验。
4.2 数值精度:单精度还是双精度?
这是一个非常细节但极其致命的问题。SHT 涉及大量的三角学和勒让德多项式计算,在高频段(大 l 值)会出现严重的数值不稳定。我自己调试时发现,在带宽 ( L=48 ) 的情况下,如果全程用 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=..., keepdim=True)就能搞定,但方向搞错的话,等变性就悄悄没了。
我的实际经验是,在调试任何等变网络时,都保持一个最小复现脚本在手边,这个脚本里有一个固定的随机种子、一个固定的旋转角度,和一个能肉眼可判断的空间信号。每次改动代码之后,就跑一遍检查等变性是否保持。这个脚本能救你无数次。
6. 一些值得看的参考实现
如果你打算亲自上手,除了读论文原文,下面几个仓库值得优先看:
论文作者的原始实现:在 GitHub 上搜
Spherical CNN,能找到 ICLR 2018 作者的官方代码。虽然是 TF1 版,但数学注释非常详细,特别适合理解核心思想。我个人不推荐直接拿它训练,但读代码非常友好。PyTorch 复刻版:社区里有一个比较完整的 PyTorch 版,文件组织清晰,适合在此基础上改自己的代码。它把 SHT 和 Wigner-D 分别封装成了独立的
autograd.Function,这一点对我启发很大——把底层变换和上层网络解耦,调试时能迅速隔离问题。lie_learn库:这是一个专门做李群学习相关的工具库,里面包含了大量球面谐波、Wigner-D 矩阵、旋转群采样的实现。很多 Spherical CNN 的源码都直接依赖它。如果你要实现自己的等变层,这个库的文档值得仔细看。
再提醒一点,任何仓库跑通之后,第一时间做的事情是“破坏性测试”:改一个维度,看报错信息是否准确指出问题所在。等变网络的张量维度复杂,如果你不能快速从报错中定位问题,后面改代码的效率会非常低。
我自己啃 Spherical CNN 源码,最大的体会是:数学公式和代码实现之间,存在一大片“隐含知识”。公式写的是连续的积分变换,代码里全是一堆矩阵乘法和维度变换。只有把每一步操作对应到公式的哪一个部分,才能真正弄懂这个网络在做什么。希望这篇解析能帮你少走一些弯路。