
如果你做过图像分类或者目标检测一定遇到过这种“诡异”现象训练集里都是竖着摆放的物体模型表现很好一旦把测试图旋转 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 representationGeometricTensor把普通 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_rotations8): 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, dim0).mean(dim0)它简单可靠但它仍然是“集成”而不是“等变”。模型内部没有群结构每个旋转方向都是独立做一次前向推理计算量也要乘以 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 python3.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(N8) 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_size5, padding2), enn.InnerBatchNorm(hidden1_type), enn.ReLU(hidden1_type), enn.PointwiseAvgPool(hidden1_type, kernel_size2, stride2), enn.R2Conv(hidden1_type, hidden2_type, kernel_size5, padding2), enn.InnerBatchNorm(hidden2_type), enn.ReLU(hidden2_type), enn.PointwiseAvgPool(hidden2_type, kernel_size2, stride2), enn.R2Conv(hidden2_type, out_type, kernel_size1, padding0), ) 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, trainTrue, downloadTrue, transformtransform ) train_loader DataLoader(train_set, batch_size64, shuffleTrue) model RotationEquivariantCNN() optimizer optim.Adam(model.parameters(), lr1e-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(dim1) labels).sum().item() total images.shape[0] return total_loss / total, correct / total for epoch in range(3): loss, acc train_one_epoch() print(fepoch {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(dim1) 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(dim1) correct (rotated_pred base_pred).sum().item() total images.shape[0] return correct / total test_set torchvision.datasets.MNIST( root./data, trainFalse, downloadTrue, transformtransform ) test_loader DataLoader(test_set, batch_size64, shuffleFalse) 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(N8)时模型只在旋转 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) 等变网络在点云、分子表示和机器人操作中的应用。建议你从本文的最小示例开始先把代码跑通再逐步替换成自己的数据。你会发现等变网络真正强大的地方不是某一层卷积“很神奇”而是它迫使你把一个领域里最本质的几何性质想清楚这个任务的什么变换不应该改变结果什么变换应该以可预测的方式改变结果然后把这个答案写进网络结构。这种思考方式比套用任何一个现成模型都更有价值。