如果你做过图像分类或者目标检测,一定遇到过这种“诡异”现象:训练集里都是竖着摆放的物体,模型表现很好;一旦把测试图旋转 90° 或者 180°,准确率立刻下滑。只要训练数据足够“全”,模型确实能靠数据增强硬扛过去。但真正值得思考的问题是:模型自身的结构,有没有把旋转变换纳入计算规则?这正是本文主题“旋转等变性”(Rotational Equivariance)想回答的问题。
旋转等变性不是一个锦上添花的 trick,也不是“多转几张图训练”这种经验性补救。它是一种设计原则:当输入图像发生旋转时,模型的中间特征和最终输出是否按同样规则、可预测地发生变换。如果做到了,模型就不是“见过大量旋转图片所以猜对了”,而是从结构上保证了旋转后的结果和旋转前的结果保持一致的映射关系。两者的区别,是经验记忆和几何建模的区别。
这篇文章会从直观例子讲到群论基础,再用手写代码和 escnn 库演示一个完整的旋转等变 CNN 训练与验证流程。你不需要成为数学专家也能读懂核心思想。读完你会明白:旋转等变网络适合解决什么问题,不适合解决什么问题;它和数据增强到底差在哪里;以及在实际 PyTorch 项目里,怎样快速搭建一个可运行的等变模型。
1. 为什么说旋转等变性是几何先验,而不是数据增强
先看一个最简单的场景。我们要训练一个 CNN 做手写数字识别。原始训练集里,“6”基本都是正着写的,几乎没有倒过来的样本。测试时如果输入一张旋转 180° 的 6,普通 CNN 极可能把它错认为 9。要解决这个问题,最容易想到的方法是做旋转数据增强:把训练图片随机旋转一定角度,让模型“见过”各种姿势的 6。
数据增强确实有效,但它的本质是“堵漏洞”。模型依然不知道 6 旋转 180° 就是 6,它只是从扩充后的数据里额外学了一条经验规则。如果测试时出现训练阶段从未见过的角度,比如训练时只增强到 90°、180°、270°,测试却来一个 47° 的旋转,模型可能再次崩溃。换句话说,数据增强把“旋转不变性”当作统计规律去拟合,而不是当作物理规则去建模。
旋转等变性则要求从网络结构层面解决这个问题。一个对旋转等变的特征提取器可以表达为:
[ f(\text{Rot}\theta(x)) = \text{Rot}\theta(f(x)) ]
通俗解释就是:先旋转输入再送入网络,和先送入网络再旋转特征,得到的结果应当一致。如果做到这一点,网络虽然在数学上仍然是一个神经网络,但它的计算图内部已经显式编码了“旋转这种变换不会改变语义”的先验。
这种先验在很多视觉任务里非常合理。显微镜图像的方向、遥感影像的方向、医疗影像的方向、大幅面工业检测中工件的摆放方向,本质上都不是语义信息。模型把大量参数浪费在记忆方向特征上,是一种结构性的浪费。等变模型把方向变化“交给结构去处理”,让网络容量更集中在真正的判别内容上。
对于分类任务,我们最终往往希望输出是“旋转不变”的。对于检测、分割、姿态估计等任务,我们可能反而希望高层特征保留方向信息。等变网络的好处是它提供的是一个通用框架:你可以用等变卷积保留旋转信息,也可以再通过群池化把它变成不变特征。相比之下,普通 CNN 想保留或丢弃方向信息都没有显式手段,只能依赖隐式学习。
2. 旋转不变性、旋转等变性和数据增强的区别
很多初学者把这三个概念混在一起。先做一次严格区分。
旋转增强是一种训练策略,它不改变网络结构。模型仍然是普通卷积,只是训练样本变多了。从泛化角度看,它让模型在训练分布覆盖到的角度上表现得更好,但模型对旋转的适应没有数学保证。
旋转不变性是一种性质,指模型的输出在输入旋转后保持不变:
[ f(\text{Rot}_\theta(x)) = f(x) ]
这只适合分类、全局检索等任务。比如判断“这张图是不是包含猫”,猫头朝上还是朝下不影响结果。但对目标检测来说,模型需要输出物体的位置和姿态,如果特征表示完全不变,反而会丢失方向信息。
旋转等变性是更精细的要求:
[ f(\text{Rot}\theta(x)) = \text{Rot}\theta(f(x)) ]
输入转了 90°,输出特征图也转 90°;输入转了 45°,特征图也转 45°。这样,网络在“知道方向”的同时,又不被迫用大量卷积核去从零学习每种方向。
可以通过一个表格理解三种方案的差异:
| 方案 | 是否修改网络结构 | 是否对任意角度有理论保证 | 高层特征保留方向 | 计算开销 |
|---|---|---|---|---|
| 旋转数据增强 | 否 | 否,只覆盖训练过的角度 | 取决于训练情况 | 通常较低 |
| 旋转不变网络 | 是 | 是 | 不保留 | 中等 |
| 旋转等变网络 | 是 | 取决于设计的群与表示 | 保留 | 相对较高 |
真正容易踩坑的地方是:许多项目把“旋转增强后的模型”说成“旋转等变模型”。实际上增强后的普通 CNN 是一个对旋转“近似稳健”的模型,不是等变模型。等变性强调的不是“某种输出恰好不变”,而是“在结构中存在一个可验证的对称关系”。这也是等变模型可解释性更强的原因:它的行为不是靠运气,而是由群卷积的定义保证的。
所以,如果你的业务中旋转角度是可枚举的、固定的,比如 0°、90°、180°、270°,数据增强可能够用。但如果你处理的是连续角度、任意旋转的输入,或者你希望模型把学习容量从方向拟合中释放出来,那么旋转等变结构更值得研究。
3. 从群论到等变卷积:一个能用代码解释的数学框架
旋转等变性的数学工具是群论,但这里只需要理解几个最基础的概念。
群是一组变换的集合,并且满足封闭性、结合律、存在单位元和逆元。对二维旋转来说,所有角度旋转组成了一个连续群,记作 SO(2);如果我们只关心 N 个离散角度,比如 90° 的倍数,就得到一个有限旋转群,通常记作 C4 或 C8,取决于单位角是 90° 还是 45°。
“表示”是群论里另一个重要概念。我们可以把群元素从“抽象的旋转变换”对应为“对特征空间的具体操作”。比如一张 8 通道特征图,一个 45° 旋转既可以表现为“图片在平面上转了 45°”,也可以表现为“8 个通道按照某种顺序换了一下位置”。后者就是群在特征空间中的一个表示。
普通卷积的平移等变性是卷积天然具备的,原因是同一个卷积核会在空间不同位置滑动。但普通卷积没有内置旋转等变性。旋转一个输入图片,再走一遍普通卷积,得到的特征图不会等于先卷积再旋转的特征图。原因很直接:卷积核本身没有参与旋转,它在方向和位置上的响应模式并不对称。
群等变卷积(Group Equivariant Convolution)解决这个问题的方法是:把卷积作用域从平面上的 (x, y) 扩展为“平面位置 + 群元素 g”。每个特征不再只是 2D 平面上的一个通道,而是定义在平移旋转群上的一个场。卷积核在平面上滑动时,同时要沿着群元素的方向做平移。这样,输入旋转会表现为特征在群维度上的置换或变换,卷积操作对这个置换是兼容的。
这听起来抽象,但实现原理非常直观。假设我们只考虑 4 个旋转角度:0°、90°、180°、270°。普通卷积的输入是 H×W×C。群卷积可以把输入升维成 H×W×(C×4),其中 4 个方向的副本分别对应四种旋转。卷积核不再只沿空间位置滑动,还会沿 4 个“方向副本”的维度做共享权重滑动。旋转输入时,网络内部的工作只是把 4 个方向通道换一个顺序,后续卷积仍然以相同方式进行。
为了不用自己实现这类复杂卷积,社区通常使用封装好的库。常见的是 escnn,它支持二维平面上的旋转群、翻转群以及三维空间中的旋转群。escnn 的核心抽象包括:
- gspace:描述输入信号所在的空间和对称群;
- FieldType:描述一层特征的类型,例如每个位置使用 regular representation 还是 irreducible representation;
- GeometricTensor:把普通 PyTorch Tensor 包装成带群结构信息的张量;
- R2Conv:实现了二维旋转等变的卷积层。
理解这些抽象之前,不需要把数学推导全部搞懂。可以先把它当作一个“类型系统”:每一层输入输出都要声明自己属于哪种群表示,库会根据表示关系约束卷积核,使网络天然旋转等变。
4. 普通 CNN 为什么不具备旋转等变性
我们可以用一个极端例子来理解普通 CNN 的缺陷。假设输入一张 MNIST 数字 3。经过第一层卷积后,特征图会突出响应最强的纹理方向。把图旋转 180° 后再过同一层卷积,卷积核在空间上看到的图案相对位置发生了彻底改变。数字 3 的曲线方向和卷积核的纹理方向不再对齐,导致第一层输出的特征差异很大。
CNN 有平移等变性,是因为卷积核对于“相同的局部图案出现在不同位置”不敏感。卷积核在整张图上滑动,这种操作天然与平移操作可交换。旋转是另一种几何变换,它改变了局部图案与卷积核之间的相对方向,普通卷积核无法自动适应这种方向变化。
有人可能会问:CNN 里的最大池化难道不是为了提供一定平移不变性吗?对,最大池化只在空间局部窗口取最大值,对微小平移有一定稳健性,但它并不是旋转等变的。它不会把旋转后的特征“旋转回来”,只是压缩信息。更关键的是,普通池化会对特征做全局或局部聚合,一旦特征本身不是等变的,之后的分类头也只能在特定方向上表现好。
从另一个角度看,CNN 实际上把方向当作一种需要学习的特征。比如一个卷积核如果对横向边缘响应强,那它对旋转后的纵向边缘就不敏感。为了覆盖所有方向,网络只能增加卷积核数量,让一部分核学习横向、一部分核学习纵向。这就是为什么普通 CNN 通常需要大量参数才能逼近旋转稳健性,而等变卷积通过结构约束让同一个卷积核自动覆盖所有方向。
等变卷积之所以能减少参数,并不是因为它“魔改”了卷积,而是它把核空间做了约束。普通核空间是一个完整的 k×k 卷积核集合,等变卷积核则要求核在不同方向之间存在确定的对应关系。于是自由参数大幅减少,特征表达能力反而集中在有效模式上。
5. 主流实现路线:从简单包装到 Steerable CNN
想要在实际任务里获得旋转等变性,不是只有一种做法。不同路线有不同复杂度、精确度和适用范围。
5.1 旋转增强 + 预测时集成
最简单的一种“弱等变”做法,是在推理阶段对输入做 N 次旋转,得到 N 个预测,然后对概率取平均。这种做法在分类里很常见,工程上称为 Multi-View Inference。
def rotation_augmented_predict(model, x, num_rotations=8): preds = [] for k in range(num_rotations): xr = torch.rot90(x, k, dims=(-2, -1)) logits = model(xr) preds.append(logits) return torch.stack(preds, dim=0).mean(dim=0)它简单可靠,但它仍然是“集成”而不是“等变”。模型内部没有群结构,每个旋转方向都是独立做一次前向推理,计算量也要乘以 N。
5.2 群卷积网络
Cohen 和 Welling 提出的 Group Equivariant CNN 是更系统的方法。网络里第一层完成“ lifting”操作,把平面图像提升到群上的函数。后续卷积层定义在群上,因此对群的旋转自然等变。离散旋转群 C4、C8 的实现相对简单,适合教学和对性能要求高的场景。
5.3 Steerable CNN
如果想对任意连续角度的旋转都保持等变,需要引入 Steerable CNN。这类网络使用不可约表示分解特征空间,约束卷积核满足 steerability 条件。escnn 库可以方便地构建这种网络,并且可以根据任务选择离散群 N 或者连续旋转群。连续旋转群实现起来更复杂,但能避免离散群“对 47° 旋转不保证等变”的问题。
5.4 三维场景与 SE(3) 等变
在三维点云、分子建模、机器人操作中,需要处理的不只是二维旋转,而是三维旋转群 SO(3) 或三维旋转平移群 SE(3)。这类任务里,普通 3D 卷积也无法保证旋转等变。许多库提供了 SE(3) 等变卷积、等变图神经网络等算子,适合处理原子坐标、刚体变换等场景。
这篇文章后面的完整示例将使用离散旋转群 C8,实现一个图像分类网络。它是理解群卷积和 Steerable CNN 的好起点。
6. 环境准备与依赖安装
本文的代码示例使用 PyTorch 和 escnn。环境要求主要是 Python 3.9 以上,以及能安装 PyTorch 的机器。如果在 CPU 上跑 MNIST 小模型,没有 GPU 也能完成实验。
conda create -n equiv python=3.10 -y conda activate equiv pip install torch torchvision pip install escnn不同机器上的 CUDA 版本会直接影响 PyTorch 的安装方式,请以你的实际驱动为准。建议先安装 CPU 版验证代码,后续再切换到 GPU 版。
安装 escnn 后,可以快速验证是否导入成功:
python -c "from escnn import gspaces, nn; print('escnn ok')"如果安装过程报编译错误,多数情况是因为 PyTorch 或 Python 版本与 escnn 的要求不匹配。此时不要盲目升级,先去 escnn 官方文档确认版本对应关系。还有一种常见问题是机器上同时存在多个 conda 环境,导致pip install装到了一个torch版本不同的环境里。建议在一个干净虚拟环境中执行全部命令。
7. 完整示例:用 escnn 构建旋转等变 CNN
下面用一个“旋转 MNIST”场景演示旋转等变 CNN。模型结构基于 escnn 的R2Conv,旋转群选择 C8,也就是单位旋转角度为 45°。
import torch import torch.nn as nn from escnn import gspaces, nn as enn gspace = gspaces.Rotation2DOnR2(N=8) in_type = enn.FieldType(gspace, 1 * [gspace.regular_repr]) hidden1_type = enn.FieldType(gspace, 12 * [gspace.regular_repr]) hidden2_type = enn.FieldType(gspace, 24 * [gspace.regular_repr]) out_type = enn.FieldType(gspace, 10 * [gspace.trivial_repr]) class RotationEquivariantCNN(nn.Module): def __init__(self): super().__init__() self.features = enn.SequentialModule( enn.R2Conv(in_type, hidden1_type, kernel_size=5, padding=2), enn.InnerBatchNorm(hidden1_type), enn.ReLU(hidden1_type), enn.PointwiseAvgPool(hidden1_type, kernel_size=2, stride=2), enn.R2Conv(hidden1_type, hidden2_type, kernel_size=5, padding=2), enn.InnerBatchNorm(hidden2_type), enn.ReLU(hidden2_type), enn.PointwiseAvgPool(hidden2_type, kernel_size=2, stride=2), enn.R2Conv(hidden2_type, out_type, kernel_size=1, padding=0), ) def forward(self, x): x = enn.GeometricTensor(x, in_type) x = self.features(x) return x.tensor这段代码的关键点不是“写一个普通 CNN 再套壳”,而是每一层都基于FieldType声明了输入输出张量的群表示。regular_repr表示该层特征在旋转群下会随着旋转而等变地置换;trivial_repr表示最后一层输出在旋转群下保持不变。这样最后的分类 logits 就天然是旋转不变的。
需要特别注意的是:enn.InnerBatchNorm、enn.ReLU、enn.PointwiseAvgPool都是 escnn 自己提供的算子。你不能在这条链里随意插入torch.nn.BatchNorm2d或torch.nn.MaxPool2d,因为这些普通算子会破坏张量的群结构。虽然最终的 tensor 仍然是 4 维的形状,但它的通道维已经被表示为 multiple group orbits,普通 PyTorch 层不理解这种约束。
8. 训练代码与旋转等变性验证
下面用 MNIST 训练一个简单分类器。这里不刻意使用旋转增强,因为我们希望验证等变结构本身对旋转的适应能力。
import torch import torch.optim as optim import torchvision import torchvision.transforms as transforms from torch.utils.data import DataLoader transform = transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ]) train_set = torchvision.datasets.MNIST( root="./data", train=True, download=True, transform=transform ) train_loader = DataLoader(train_set, batch_size=64, shuffle=True) model = RotationEquivariantCNN() optimizer = optim.Adam(model.parameters(), lr=1e-3) criterion = nn.CrossEntropyLoss() def train_one_epoch(): model.train() total_loss, correct, total = 0.0, 0, 0 for images, labels in train_loader: optimizer.zero_grad() logits = model(images) loss = criterion(logits, labels) loss.backward() optimizer.step() total_loss += loss.item() * images.shape[0] correct += (logits.argmax(dim=1) == labels).sum().item() total += images.shape[0] return total_loss / total, correct / total for epoch in range(3): loss, acc = train_one_epoch() print(f"epoch {epoch + 1}, loss={loss:.4f}, acc={acc:.4f}")训练完成后,可以做一个简单的离散旋转等变性测试:把同一批图片分别旋转 0°、90°、180°、270°,再送入模型,检查预测标签是否保持一致。
def rotated_accuracy(model, loader): model.eval() correct, total = 0, 0 with torch.no_grad(): for images, labels in loader: base_logits = model(images) base_pred = base_logits.argmax(dim=1) for k in range(1, 8): rotated_images = torch.rot90(images, k, dims=(-2, -1)) rotated_logits = model(rotated_images) rotated_pred = rotated_logits.argmax(dim=1) correct += (rotated_pred == base_pred).sum().item() total += images.shape[0] return correct / total test_set = torchvision.datasets.MNIST( root="./data", train=False, download=True, transform=transform ) test_loader = DataLoader(test_set, batch_size=64, shuffle=False) print("rotation consistency:", rotated_accuracy(model, test_loader))如果模型构建正确,这个“旋转一致性”应该接近 1。这里我们比较的是旋转后的预测结果与原始角度的预测结果是否一致。需要注意,训练 3 个 epoch 后模型本身准确率可能没有达到最高,但等变性验证应当不受影响,因为等变性是结构保证,不是训练技巧。
如果发现旋转一致性明显低于 1,最可能的原因是模型某个层使用了破坏群结构的普通算子,或者gspace的参数与输入几何不匹配。还有一种情况是浮点误差被放大,不过对于分类 argmax 来说,这种误差通常只在极少数边界样本上出现。
9. 运行结果与效果验证
在没有 GPU 的普通 CPU 机器上,上述小模型大约几分钟就能跑完一个 epoch。第一个 epoch 结束后,MNIST 准确率通常可以达到 90% 以上。第三个 epoch 后,训练集准确率可以达到 97% 左右。由于示例只训练三个 epoch,最终测试集准确率不必追求极致。
执行旋转一致性测试时,预期输出是一个接近 1.0 的小数:
rotation consistency: 0.9989如果结果接近 1.0,说明模型对 8 个离散方向都能保持一致的预测。这个结果比普通 CNN 加旋转增强更“硬”的地方在于:普通 CNN 预测时,每转一个角度都要重新推理一次;这里的模型在结构上能做到单次推理即可适应旋转后的输入。
如果你想验证同一模型在普通 CNN 下的差异,可以做一个对照实验:把escnn相关模块替换成普通卷积、普通 BatchNorm 和普通 MaxPool,然后用完全相同的数据和训练轮数训练。在 MNIST 这种不考虑旋转的数据集上,两者差异不一定明显。但如果在测试阶段把所有图片旋转 90° 或随机旋转,普通 CNN 的准确率往往会明显下降,而等变网络几乎不受影响。
这也是解释旋转等变价值最有效的实验方式:训练阶段不加旋转增强,测试阶段只旋转输入,然后比较普通 CNN 与等变 CNN 的准确率差距。你会发现,普通 CNN 对旋转非常脆弱,等变 CNN 却因为几何结构保持稳定。
10. 常见问题与排查思路
在实际使用 escnn 或实现等变网络时,开发者最容易遇到下面几类问题。
| 问题现象 | 可能原因 | 排查方式 | 解决方案 |
|---|---|---|---|
FieldType构造报错 | gspace和FieldType不在同一个库版本下 | 检查导入来源是否一致 | 统一使用escnn.gspaces与escnn.nn |
| 前向传播报维度不匹配 | 输入张量不是 4 维或通道数不等于 FieldType 维度 | 打印张量形状和 FieldType size | 调整网络输入类型或预处理 |
| 模型输出错误 | 在等变模块链中混入了普通 PyTorch 层 | 查看模型结构定义 | 把普通 BatchNorm、MaxPool 换成 escnn 对应算子 |
| 预测在旋转后不稳定 | 选择的离散群 N 不能覆盖测试旋转角度 | 检查测试旋转是否为 45° 的整数倍 | 增大 N 或使用连续旋转群 |
| 训练速度明显变慢 | 群表示扩大了通道数或操作复杂度 | 记录每层耗时 | 减少每层 channels,或调整 N |
| 显存不足 | 中间特征通道数过大 | 监控显存占用 | 降低 hidden 层 channels 或 batch size |
第一类问题的根源是 e2cnn 和 escnn 两个库历史上有继承关系,很多网上教程还在使用老接口。如果你参考的是老代码,很可能出现r2_act和gspace混用的问题。稳妥做法是以官方文档为准,统一使用新库名和新接口。
还要特别提醒:等变卷积的“等变”范围是由你选择的群决定的。使用Rotation2DOnR2(N=8)时,模型只在旋转 45° 整数倍时严格等变。如果测试出现 30° 旋转,模型并不能从数学上保证结果不变。要做真正的连续旋转等变,需要换用连续旋转群的表示。这个细节在实际业务中非常重要,因为真实照片里的旋转角度往往是任意的。
另一个容易出错的地方是输入图像的几何约定。某些数据处理流程会把通道维放在最后,某些会提前做归一化,这些都不会影响群的等变性,因为等变性针对的是空间坐标的旋转。但如果你在数据预处理时不小心把图片做了非等比的 resize,旋转等变性就无从谈起了,因为输入本身已经被破坏。
11. 最佳实践与工程建议
旋转等变模型并不是在所有任务里都优于普通 CNN,也不是模型越大越好。下面这些建议来自实际工程中比较常见的取舍。
第一,先判断任务是否需要旋转等变。如果业务数据中物体方向原本就固定,比如车牌识别、文档 OCR,旋转等变帮助有限。如果物体方向随机或变化很大,比如遥感、病理切片、无人机航拍、工业零件检测,旋转等变性会带来显著收益。判断方法很简单:拿一批测试样本做随机旋转,看普通 CNN 的精度下降是否超过可接受范围。
第二,从离散旋转群开始。对一张二维图片,先用 C4 或 C8 建模,接口简单,训练速度也比连续群快得多。C8 已经覆盖 45° 间隔的旋转,在许多离散场景下足够。如果后续确实遇到任意角度的需求,再迁移到连续旋转群。每类群的计算复杂度不同,不要一上来就上最复杂的连续模型。
第三,注意“等变”并不等于“平移不变”。在分类任务里,你通常希望输出层是 trivial 表示,也就是对旋转不变;而在目标检测、语义分割里,中间层需要保留方向信息,最后的 head 再根据任务决定是否做群池化。不要把网络所有层都设成 invariant,否则空间位置或方向信息有可能提前丢失。
第四,训练稳定性上,等变卷积的自定义底层实现细节较多,更建议优先使用 escnn 这类成熟库。不要一开始就自己手写群卷积核约束,因为卷积核的旋转关系、边界填充、采样方式都会影响数值稳定性。先把成熟库跑通,再根据业务需要做二次开发。
第五,等变模型不是对数据增强的“替代品”。即使结构上已经旋转等变,仍然可以保留一些与旋转无关的增强,比如颜色抖动、噪声、裁剪。裁剪存在的意义是模拟不同位置和尺度变化,这与旋转没有冲突。合理的组合是:几何增强解决模型能力之外的几何先验,像素级增强解决光照和噪声变化。
第六,在项目落地前,建立一套标准的验证指标。不要只看总体准确率,至少要看“旋转一致性”指标。最简单的定义是:把测试集图片旋转多个角度,统计每张图的预测标签与原始预测标签一致的比例。这个指标能帮助你判断网络结构是否真正做到了等变。最好再准备一个包含连续角度的旋转测试集,用来评估离散群之外角度上的表现。
第七,关于参数和计算量。等变模型常通过增加群表示维度来提升表达力,因此参数数量和普通 CNN 不能直接对比。你真正应该对比的是“等变模型在某层用 16 channels 的表示”和“普通 CNN 用 64 channels 的特征”在同等精度下的参数量与延迟。在很多任务上,等变模型可以用更少的参数量达到同等旋转稳健精度,但延迟不一定更低,需要针对具体硬件做 benchmark。
12. 总结与推荐学习路线
旋转等变性是机器学习中为数不多能把几何先验直接写进网络结构的思路。它把“物体旋转后语义不变”从数据层提升到模型层,通过群卷积、Steerable 核等机制让网络对旋转有明确、可验证的响应方式。与数据增强相比,它更省参数、更能处理训练阶段未出现过的方向,也让模型行为更可控。
如果你想继续深入,建议按下面的路线学习。
第一步,吃透离散群等变卷积。用论文《Group Equivariant Convolutional Networks》作为起点,理解 lifting convolution 和 group convolution 的区别。第二步,阅读《Steerable CNNs》和《General E(2)-Equivariant Steerable CNNs》,理解 kernel constraint、irreducible representation 等概念。第三步,用 escnn 在 CIFAR-10、Rotated-MNIST 或你自己的业务数据上复现实验,重点比较旋转一致性和参数效率。第四步,如果研究方向是 3D 感知,再学习 SE(3) 等变网络在点云、分子表示和机器人操作中的应用。
建议你从本文的最小示例开始,先把代码跑通,再逐步替换成自己的数据。你会发现,等变网络真正强大的地方不是某一层卷积“很神奇”,而是它迫使你把一个领域里最本质的几何性质想清楚:这个任务的什么变换不应该改变结果,什么变换应该以可预测的方式改变结果,然后把这个答案写进网络结构。这种思考方式,比套用任何一个现成模型都更有价值。