ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

SE(3)-Transformer:让Transformer理解三维空间旋转与平移的等变注意力机制

SE(3)-Transformer:让Transformer理解三维空间旋转与平移的等变注意力机制 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 对每个标量通道独立激活但在等变模型里不能直接对向量分量用 ReLUReLU 是逐元素操作它作用在向量坐标上会改变向量的长度和方向而旋转后这个运算结果不能和原来的旋转结果对齐。SE(3)-Transformer 采取的策略是门控非线性对每个向量或高阶通道先用一个小网络从标量通道计算一个门控值再用这个标量门控去缩放向量特征。因为缩放系数是标量向量方向保持不变只有长度和符号被调节这样非线性操作就不会破坏等变性。这个设计也被称为“等变 MLP”或“gated nonlinearity”。实际代码里每个 block 的局部更新可能有两次一次是对标量通道做常规激活再加权重第二次是用标量门控去缩放张量通道。模型里比较成功的做法是采用一种类似 pre-LN 的结构先对输出的各阶分量做范数归一化和门控再走残差连接。这里有一个实操心得训练初期如果高阶特征通道学习不充分门控值往往很小导致梯度传播弱。可以先用较小的num_degrees跑通流程逐步增加阶数不要一上来就堆 4 阶甚至 5 阶训练不稳定很容易劝退人。3. 实操指南从零配置SE(3)-Transformer3.1 环境与代码选择SE(3)-Transformer 有一版官方实现基于 PyTorch 和 e3nnGitHub 上也有非官方的 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( dim64, depth4, input_degrees1, num_degrees3, output_degrees1, reduce_dim_outFalse ) 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 之间比较常见分子性质预测通常用 64num_degrees指模型内部使用到的最大不可约表示阶数2 到 4 都可以跑数据几何结构越复杂比如需要表达局部曲率、手性越需要高阶特征但计算量和内存也随之增长。depth是层数3 到 6 层基本够用。对点云任务层数太深反而容易过平滑所有节点特征趋于一致。output_degrees根据下游任务决定如果要输出点级 et al. 向量比如力场预测就把输出阶数设成 1如果只预测标量属性就设成 0。实际操作中我一般先用num_degrees2、depth3跑通一个小版本观察训练曲线再增加容量。邻居数量是另一个隐藏超参数。模型复杂度随图边数线性增长如果每个节点连 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, atol1e-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_outTrue与output_degrees的组合会导致输出通道数变化不同项目里含义不同。我遇到过在旧项目里能跑的配置换到新版本直接报维度不匹配。排查维度问题最快的办法是打印模型输出的 shape 和 irreps 信息手动比对预期结果不要盯着报错信息猜。个人经验与一点后续建议我最早接触 SE(3)-Transformer 时总觉得等变模型那么精巧应该比普通图神经网络强一大截。实际测试发现效果确实好但好得很“挑剔”数据质量差、建图不合理、距离编码不匹配时甚至不如简单的 PointNet 结构稳。这让我明白一个道理等变性解决的是“几何对称性”问题不是“特征表达丰富度”问题。它让你在同等数据量下泛化更好但前提是你得先把几何预处理做对。如果做三维分子性质或点云理解我建议先用小模型、小数据集快速验证任务是否值得用等变模型如果数据里物体朝向差异很大、旋转增强很难覆盖全域那么 SE(3)-Transformer 的价值就很明显如果数据已经严格对齐到某个模板比如人脸对齐后的关键点用普通模型可能更划算因为模型不需要额外的等变能力。最后分享一个小技巧在跑大实验前先写一个 50 行的等变性单元测试放进 CI 或模型代码库里。每次改动模型结构或更新依赖都自动跑一遍。我靠这个测试抓过至少三次因 e3nn 版本升级导致的静默错误——模型能跑loss 能降但输出的等变性质已经悄悄失效如果不校验模型训完都不知道结果是有偏的。
返回列表