☰
SE(3)-Transformer:让Transformer理解三维空间旋转与平移的等变注意力机制
2026/10/3 5:35:36 网站建设 项目流程

SE(3)-Transformers 这个名字听起来像科幻片里某种变形金刚,但它其实是深度学习里一个相当硬核的方向:让 Transformer 真正理解三维空间里的旋转和平移。对三维点云、分子结构、蛋白质骨架这类数据,传统的神经网络往往把坐标当成普通数字硬塞进去,模型确实能学,但学到的是“在某个坐标系下有效”的表示;而 SE(3)-Transformers 的目标是让模型天然具备对刚体变换的感知能力——你把输入场景转个角度、挪个位置,模型输出的特征会按照对应规则跟着变换,而不是学一套只认固定朝向的模板。这种“等变注意力”机制在处理物理系统、几何数据时优势非常明显,这篇内容就是基于实际跑通 SE(3)-Transformers 的经验,把原理、实现细节和踩坑点拆开讲清楚,适合正在研究点云任务、分子性质预测或机器人操作的工程师和研究者参考。

1. SE(3)等变性是什么?为什么三维模型需要它?

1.1 从刚体变换谈起:SE(3)群到底在描述什么

先不要被数学符号吓住。SE(3) 是 Special Euclidean Group 的缩写,描述的是三维空间里所有刚体变换的集合。刚体变换很简单:先做一个旋转,再做一个平移,组合起来就是物体在空间中“不改变形状和大小”的移动。你在桌上把手机转个角度、往右推两厘米,这就是一次 SE(3) 变换。

之所以叫“群”,是因为这些变换之间有封闭性:先做一个刚体变换,再做另一个刚体变换,结果还是一个刚体变换;每个变换都有逆变换,可以把物体变回原来的位置。这个性质很重要,它保证了我们可以在统一的框架里讨论“输入被变换后,输出应该怎么跟着变换”。

在真实世界的数据里,SE(3) 变换无处不在。一个分子在真空里怎么旋转,它的化学性质都不变;蛋白质折叠后,你把整个结构旋转,它的功能依然一样;机器人抓取一个杯子,杯子在桌面上换个位置换个角度,抓取策略本质上应该是一个模式。自然规律本身不依赖坐标系的选择,这叫做物理对称性。如果一个神经网络想逼近这类规律,那么把对称性直接编码进网络结构,是最省力也最正确的做法。

1.2 等变性 vs 不变性:二者差异与适用边界

很多人会搞混“等变”和“不变”。这两个概念确实相关,但指的不是一回事。不变性是说:不管输入怎么变换,输出都不变。比如想知道一个分子的能量,分子无论怎么旋转,能量是同一个数值。这种任务里模型最终只需要输出标量,很多标量回归网络能做到旋转不变。

等变性则更强:输入做了某个变换,输出特征也要按照对应的方式变换。拿图像分类举例,普通 CNN 对一张猫的图片做平移,卷积特征图也会跟着平移,这就是平移等变,最后通过池化汇总得到“这是猫”的不变判断。在 SE(3)-Transformer 里,输入点云旋转后,网络中间层输出的向量特征、张量特征要按照旋转矩阵对应的不可约表示去旋转。这种性质对于需要中间特征去指导后续任务的场景非常关键:

  • 机器人抓取中,输出的抓取姿态必须跟随物体旋转而旋转,不能“原地不变”。
  • 三维点云配准中,估计出的变换参数要和输入坐标的变化保持协同。
  • 多尺度场景下,如果模型只输出标量,往往丢失了物体朝向等几何信息。

总结一下就是:不变性适合最终只关心一个数值的任务,等变性适合中间特征和输出仍然有几何意义的任务。SE(3)-Transformer 因为把等变性嵌入到注意力每一层,所以同一个模型既能做等变输出,也能通过最后只取标量分支来实现不变输出,覆盖范围更广。

1.3 普通Transformer在三维任务里的局限

Transformer 在序列领域很强,核心是自注意力机制:对每个 token 算 query、key、value,然后加权求和。这套机制本身并不理解“旋转”和“平移”。如果你只是把点云里每个点的 x、y、z 坐标拼到 token 特征里,旋转一下输入,坐标发生变化,注意力分数也会跟着变,网络输出当然不稳定。

很多人尝试的补救办法是数据增强:训练时随机旋转,让模型见过足够多的姿态。这确实有效果,但本质上是让模型去“记住”各种朝向的分布,而不是真正理解对称性。增强可以覆盖有限的旋转采样,很难覆盖连续空间里无穷多种姿态;而且一旦训练集里某个朝向出现得少,模型对这个方向的泛化就会打折扣。

另一个问题是坐标的平移敏感性。如果直接用绝对坐标作为输入特征,物体放在场景左边还是右边,网络看到的信号就完全不同。要解决这个问题,通常要人为构造相对位置信息,比如计算点与点之间的欧氏距离。但距离是标量,在旋转和平移下保持不变,它本身只提供了几何约束,却丢失了方向信息。如何在保留方向信息的同时又不破坏等价关系,这正是 SE(3)-Transformer 要解决的核心矛盾:既要用方向向量,又要让方向向量在输入旋转时按对应规则旋转,而不是被网络“碾平”成普通数值。

2. SE(3)-Transformer核心机制拆解:注意力如何做到等变?

2.1 输入组织与特征表示:图结构+不可约表示

SE(3)-Transformer 处理的数据不是规则网格,而是点云或图结构。输入是一组节点,每个节点有三维坐标,节点之间通过边连接,边上可以有距离等几何属性。这个结构和分子图、蛋白质残基接触图、点云近邻图天然匹配。实际使用中,一般用 k 近邻图或者半径图来定义节点之间的边,控制计算量。

节点特征的设计是整个方法的关键。普通 GNN 里节点特征就是一个向量,但 SE(3)-Transformer 里每个节点的特征被组织成“多重不可约表示”(irreps)。简单理解,就是特征被拆成不同“阶”的分量:0 阶是标量,旋转下不变;1 阶是三维向量,旋转下像普通空间向量一样被旋转矩阵作用;2 阶以上是更复杂的张量,对应更高频的角向变化。每个阶都有若干通道,e3nn 里用类似"16x0e + 8x1o + 4x2e"的字符串表示,意思是 16 个标量通道、8 个向量通道(奇偶性 opposite)、4 个二阶张量通道(偶数 parity)。

这一设计的直接好处是:模型输出的中间特征里,不同阶的信息有明确的几何含义。标量通道可以编码密度、能量等不变信息;向量通道可以编码方向、流场;高阶通道可以编码更复杂的局部几何模式。这比单纯把所有信息塞进一个扁平向量要干净得多,也为最后输出旋转等变特征打下了物理基础。

2.2 注意力分数:为什么内积是安全操作

Transformer 的注意力公式不复杂:query 和 key 做内积,得到标量分数。在等变模型里,query 和 key 的选择要非常小心。SE(3)-Transformer 的做法是:query 和 key 节点特征的全部阶分量都经过一个“等变线性层”,然后只取其中的标量部分(0 阶分量)做内积。

为什么要只取标量部分?因为两个向量直接做内积,会得到旋转不变的标量;对两个 1 阶向量做内积,结果在旋转下不变。这意味着注意力分数天然具有旋转不变性:输入旋转后,注意力分数保持不变。这正是我们需要的性质——一个点在旋转后的点云中,应该仍然关注它在原图中关注的邻居。

注意力公式里的偏置项同样有讲究。偏置通常由一个标量 RBF 特征网络构成,输入是节点之间的欧氏距离。距离本身在旋转和平移下不变,所以偏置也不变。模型还可以在边上额外拼接方向向量,利用一个等变网络把方向和径向距离一起编码进偏置,但最终输出的偏置依然是标量,这样不会破坏等变性。

这里有一个容易理解的类比:两个人认路,不管地图拿正了还是倒转了,他们对“前面那个路口往左转”的判断是一致的。注意力分数就是这个“跨姿态的稳定判断”。

2.3 消息传递与等变线性层:如何不破坏变换性质

有了注意力分数,接下来要做加权聚合。普通注意力把节点 j 的 value 乘以注意力权重再求和。在 SE(3)-Transformer 里,value 是节点特征经过等变线性变换后的结果,它同样可以包含多个阶的分量。因为注意力权重是标量,标量乘以一个一阶向量再求和,结果的变换性质仍然是一阶向量——旋转矩阵乘法对求和是线性的,标量系数放在前面不会改变变换规则,这一串操作保持了等变性。

但要注意,value 的生成不能使用普通矩阵乘法直接对扁平特征做线性变换,因为那会混合不同阶的特征,破坏变换性质。正确做法是使用张量积分解:把不同阶的输入特征分别映射到目标阶,同时保证“输入是 1 阶向量、输出也必须是 1 阶向量”这样的对应关系。e3nn 里的TensorProduct封装了这一复杂过程。实际使用中,通常构造一个约化的张量积,把输入阶数按照 Clebsch-Gordan 规则组合到输出阶数,并由可学习权重控制每个组合通道的强度。

消息传递结束之后,每个节点会把聚合结果和自身特征做残差连接,再进一层等变归一化。残差连接是安全的:同阶相加不会影响变换性质。归一化则需要额外注意,普通 BatchNorm 对所有样本统一计算均值方差,这在大量点云里会引入跨样本的统计信息,未必会直接破坏等变性,但会带来分布偏移和训练不稳定。SE(3)-Transformer 通常采用对每个节点独立计算范数的归一化,或者干脆用等变层归一化(对每个节点、每个阶内部做归一化),保证归一化不依赖整体数据的旋转方向。

2.4 等变非线性与归一化:一个容易翻车的环节

普通神经网络里 ReLU、GELU 对每个标量通道独立激活,但在等变模型里不能直接对向量分量用 ReLU:ReLU 是逐元素操作,它作用在向量坐标上会改变向量的长度和方向,而旋转后这个运算结果不能和原来的旋转结果对齐。SE(3)-Transformer 采取的策略是门控非线性:对每个向量或高阶通道,先用一个小网络从标量通道计算一个门控值,再用这个标量门控去缩放向量特征。因为缩放系数是标量,向量方向保持不变,只有长度和符号被调节,这样非线性操作就不会破坏等变性。

这个设计也被称为“等变 MLP”或“gated nonlinearity”。实际代码里,每个 block 的局部更新可能有两次:一次是对标量通道做常规激活再加权重;第二次是用标量门控去缩放张量通道。模型里比较成功的做法是采用一种类似 pre-LN 的结构,先对输出的各阶分量做范数归一化和门控,再走残差连接。这里有一个实操心得:训练初期,如果高阶特征通道学习不充分,门控值往往很小,导致梯度传播弱。可以先用较小的num_degrees跑通流程,逐步增加阶数,不要一上来就堆 4 阶甚至 5 阶,训练不稳定很容易劝退人。

3. 实操指南:从零配置SE(3)-Transformer

3.1 环境与代码选择

SE(3)-Transformer 有一版官方实现基于 PyTorch 和 e3nn,GitHub 上也有非官方的 PyTorch 移植版,代码风格更友好,我实际用后者比较多。安装时需要注意 e3nn 的版本差异,不同版本之间 irreps 字符串解析规则和旋转矩阵函数名有变动,建议直接用项目 requirements 里锁定的版本,不要贸然 upgrade。

基本环境是 Python 3.8+、PyTorch 1.10+、e3nn 0.4 或更新的 0.5 版本。如果跑分子性质数据集 QM9,还需要下载数据集并处理成图结构。动手前先跑通官方 README 里的 demo,验证 forward 能跑、loss 能下降,再换自己的数据,这样能把环境问题与模型问题分开排查。

写一段最简初始化代码做参考:

import torch from se3_transformer_pytorch import SE3Transformer model = SE3Transformer( dim=64, depth=4, input_degrees=1, num_degrees=3, output_degrees=1, reduce_dim_out=False ) coors = torch.randn(2, 20, 3) feats = torch.randn(2, 20, 1) # 每个点一个标量初始特征 out = model(coors, feats) print(out.shape)

这段代码创建了一个输入为标量特征、输出也是标量特征的 SE(3)-Transformer。输入特征维度可以改成和任务匹配的通道数,但注意input_degrees要和输入特征里各阶通道总数匹配,不能用普通扁平特征直接塞进去。

3.2 关键参数选择:dim、num_degrees、depth怎么定

参数选择直接影响模型容量和等变表达能力。dim指每个阶的通道数,但它不是总特征维度,而是“每个阶的基础通道数”或某种缩放因子,具体语义要看实现。经验上,dim在 32 到 128 之间比较常见,分子性质预测通常用 64;num_degrees指模型内部使用到的最大不可约表示阶数,2 到 4 都可以跑,数据几何结构越复杂(比如需要表达局部曲率、手性),越需要高阶特征,但计算量和内存也随之增长。

depth是层数,3 到 6 层基本够用。对点云任务层数太深反而容易过平滑,所有节点特征趋于一致。output_degrees根据下游任务决定:如果要输出点级 et al. 向量,比如力场预测,就把输出阶数设成 1;如果只预测标量属性,就设成 0。实际操作中,我一般先用num_degrees=2、depth=3跑通一个小版本,观察训练曲线,再增加容量。

邻居数量是另一个隐藏超参数。模型复杂度随图边数线性增长,如果每个节点连 20 个邻居,50 个节点的图还可以接受;到几千个点的大点云,全连接图直接内存爆炸。一般用 k 近邻,k 在 8 到 16 之间比较均衡。注意 k 太小时局部感受野受限,模型可能学不到长程依赖,可以在浅层用较小的 k,深层用稍大的 k,或者配合 radius graph。

3.3 等变性验证方法:训练前必须做的一步

很多跑挂的人忽略了一件事:用随机旋转和平移测试一下模型输出是否真的等变。这个测试 5 分钟就能做,却能筛掉大量实现和配置错误。

测试逻辑很简单:给定一组随机点云和特征,记录模型输出;把同样点云整体施加一个随机旋转 R 和平移 t,再输入模型,记录输出。如果模型是等变的,那么第二次输出的特征应该等于第一次输出特征经过对应旋转矩阵变换后的结果。对 0 阶输出,两者应当完全相等;对 1 阶输出,第二次输出应该等于第一次输出左乘旋转矩阵 R;更高阶输出则对应不可约表示的 Wigner-D 矩阵。

在 e3nn 里,可以用o3.Irreps.D_from_matrix得到对应阶的变换矩阵,写一个简易校验:

import torch from e3nn import o3 irreps_out = model.irreps_out rot = o3.rand_rotation() D = irreps_out.D_from_matrix(rot) x1 = torch.randn(1, 10, 3) f1 = torch.randn(1, 10, 3) # 按 input_degrees 配置 out1 = model(x1, f1) x2 = x1 @ rot.T # 注意作用方向约定 f2 = f1 # 特征本身不做处理 out2 = model(x2, f2) # 等变: out2 应变换为 out1 @ D.T if torch.allclose(out2, out1 @ D.T, atol=1e-4): print("等变校验通过") else: print("等变校验失败")

实际操作时要注意两点:第一,旋转矩阵作用在坐标上的方向要与 e3nn 的约定一致,写反了测试必挂;第二,如果模型最终做了池化或者只取标量,校验范围要相应调整。我习惯在模型实现里增加一个return_full开关,让中间各层全特征都能输出,方便调试。

3.4 训练技巧与超参数调节经验

SE(3)-Transformer 训练起来和普通 Transformer 有点不一样。第一,learning rate 不要开太大,我常用 1e-3 配合 warmup 加 cosine decay,对大模型降到 3e-4 左右。第二,梯度裁剪要留好,等变层内部的张量积运算容易产生较大梯度,我用 clip norm 1.0 比不用的稳定得多。

损失函数上没有特殊限制,标量预测就用 MSE,向量预测可以加坐标或力的监督。有一点值得注意:模型本身是等变的,理论上不需要旋转增强,但如果你训练数据里点云有边界截断(比如只保留了物体某一部分),旋转增强可能会让截断造成的伪影被放大。我实际跑点云分类时发现,去掉旋转增强后,模型在某些方向上泛化更好,因为它被迫真正依赖相对几何,而不是靠数据统计弥补朝向偏差。

训练时监控两类指标:一类是任务 loss,一类是我主动加的等变校验误差(每 1000 步重新跑一次随机旋转测试)。如果校验误差突然增大,多半是数值溢出或者某些通道变成 NaN,及早发现能省很多定位时间。

4. 常见问题与排查心得

4.1 问题速查表

现象常见原因解决思路
输出在旋转后对不上旋转矩阵作用方向写反;e3nn 版本不一致先运行官方 demo 校验;统一 D 矩阵约定
训练初期 loss 完全不下降高阶通道初始化尺度过大;学习率太高用较小 num_degrees;降低学习率;加 warmup
显存爆掉全连接图边数过多;num_degrees 太大;depth 太深限制邻居数;换小模型验证;梯度检查点
等变校验误差在几十步后变大数值溢出;某些通道变成 NaN;归一化实现破坏等变检查输入特征是否包含绝对位置;检查归一化是否按阶独立
输出只有 0 阶标量,向量通道总为 0门控非线性失效;向量通道被残差覆盖检查门控网络是否能学习到非零值;增加向量通道初始权重
平移测试失败,但旋转测试通过输入特征或偏置编码了绝对坐标确保边的特征只使用相对位置;移除绝对位置编码

4.2 坐标缩放与距离编码:一个容易被忽视的细节

SE(3)-Transformer 对坐标的绝对数值没有太多限制,但边的距离编码对尺度很敏感。模型通常用 RBF 基函数把距离展开成多个标量通道,RBF 的中心和带宽是固定的。如果坐标单位是纳米,距离范围是 0 到 1;换成埃,范围变成 0 到 10,同一组 RBF 参数覆盖的分布完全不同。

换数据集时,一定要重新统计坐标距离的分布,再调整 RBF 的边界和带宽。我在一次点云配准实验中,只换了一个数据集忘了调 RBF 边界,模型前 20 轮完全不学习,loss 卡在初始值附近。调完距离编码的范围之后,训练曲线立刻恢复正常。这也是等变模型里少有的“需要根据数据分布手动调”的地方,其他很多参数都相对通用。

4.3 数据归一化与中心化:平移等变的边界

模型理论上是平移等变的,但实际实现里存在一些隐藏的平移依赖。最典型的是输入特征如果包含不随坐标变化的全局统计量(比如整个点云的中心坐标),等价性就会被破坏。处理点云时,我一般先把所有点坐标减去点云质心,再做 k 近邻建图。这样模型看到的几何只依赖相对位置,对全局平移天然不敏感。

局部归一化同样要注意。如果你给每个节点特征拼接了它到质心的距离,这个特征在平移下是不变的(因为质心跟着平移,相对距离不变),没问题。但如果直接拼接节点的绝对坐标,哪怕后面接再强的网络,平移等变也救不回来。

还有一类边界情况:点云数量不同导致全局池化后的等变性质变化。比如做分子能量预测时,最终输出是全局标量,从各个节点特征加和得到,这个加和对旋转和平移都是安全的。但如果中间层使用了对所有节点计算均值并更新特征的操作,且均值本身依赖节点数量,等变性质不会受影响,但数值尺度要小心,尤其是分子大小差异大的数据集,建议用 sum pooling 时加一个对数尺度补偿。

4.4 版本兼容问题

e3nn 的接口在 0.3 到 0.5 之间变动不小。o3.Irreps,o3.TensorProduct,o3.rand_rotation这些核心接口虽然在,但参数名和默认值有差异。最省事的方法是固定一个版本,用虚拟环境隔离,不要跟着最新版走。

另外,reduce_dim_out=True与output_degrees的组合会导致输出通道数变化,不同项目里含义不同。我遇到过在旧项目里能跑的配置,换到新版本直接报维度不匹配。排查维度问题最快的办法是打印模型输出的 shape 和 irreps 信息,手动比对预期结果,不要盯着报错信息猜。

个人经验与一点后续建议

我最早接触 SE(3)-Transformer 时,总觉得等变模型那么精巧,应该比普通图神经网络强一大截。实际测试发现,效果确实好,但好得很“挑剔”:数据质量差、建图不合理、距离编码不匹配时,甚至不如简单的 PointNet 结构稳。这让我明白一个道理:等变性解决的是“几何对称性”问题,不是“特征表达丰富度”问题。它让你在同等数据量下泛化更好,但前提是你得先把几何预处理做对。

如果做三维分子性质或点云理解,我建议先用小模型、小数据集快速验证任务是否值得用等变模型:如果数据里物体朝向差异很大、旋转增强很难覆盖全域,那么 SE(3)-Transformer 的价值就很明显;如果数据已经严格对齐到某个模板(比如人脸对齐后的关键点),用普通模型可能更划算,因为模型不需要额外的等变能力。

最后分享一个小技巧:在跑大实验前,先写一个 50 行的等变性单元测试放进 CI 或模型代码库里。每次改动模型结构或更新依赖,都自动跑一遍。我靠这个测试抓过至少三次因 e3nn 版本升级导致的静默错误——模型能跑,loss 能降,但输出的等变性质已经悄悄失效,如果不校验,模型训完都不知道结果是有偏的。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询