160

等变神经网络:3D 分子模型必须尊重旋转和平移对称性

三维分子模型必须尊重旋转和平移对称性。这篇讲清不变性与等变性的区别、实现方式与选型判断。

处理分子的三维坐标时,模型必须尊重物理的对称性:把整个分子旋转或平移,它的性质不应改变,而预测的坐标应该跟着变。把这个约束编码进架构,而非指望模型从数据中学会,是几何深度学习的核心思想。

不变性与等变性

# 【不变性 invariance】
#   f(R·x + t) = f(x)
#   输入变换后,输出不变
#
#   适用:预测标量性质
#     能量、结合亲和力、溶解度、偶极矩大小
#
#   实现:只用【不变量】作为特征
#     - 原子间距离 |r_i - r_j|
#     - 键角、二面角
#     → 这些量本身就与朝向无关
#
# 【等变性 equivariance】
#   f(R·x + t) = R·f(x) + t
#   输入变换后,输出做相同的变换
#
#   适用:预测向量/张量
#     原子坐标、力、偶极矩(向量)、极化率(张量)
#
#   实现:
#     更新坐标时,只沿【相对位置向量】的方向移动
#     Δx_i = Σ_j (x_i - x_j) · φ(距离, 特征)
#     → 因为 (x_i - x_j) 会跟着旋转,
#       所以更新也跟着旋转 —— 天然等变
#
# 【为什么不用数据增强代替】
#   数据增强(随机旋转训练样本)
#   → 让模型「近似地」学会对称性
#   → 需要更多数据;且只是近似
#
#   架构保证的对称性:
#   → 精确成立,不需要学
#   → 【样本效率高得多】
#   → 【这在数据稀缺的化学中尤其重要】

主要的实现方式

方法 思路 特点
不变量网络(SchNet、DimeNet) 只用距离/角度 简单高效;无法输出向量
EGNN 沿相对位置向量更新坐标 简单、快、易实现
Tensor Field Network / e3nn 球谐函数 + 张量积 理论完备;实现复杂、慢
SE(3)-Transformer 等变注意力 表达力强;成本高
PaiNN / NequIP 标量-向量混合表示 性能与效率平衡好
不变点注意力(IPA) 局部坐标系中操作 AlphaFold2 使用(见 112《AlphaFold2 论文精读》

EGNN:最容易理解和实现的等变层

import torch
import torch.nn as nn

class EGNNLayer(nn.Module):
    """E(n) 等变图神经网络层"""
    def __init__(self, hidden=128):
        super().__init__()
        # 边消息:输入是两端特征 + 【距离平方】(不变量)
        self.edge_mlp = nn.Sequential(
            nn.Linear(hidden * 2 + 1, hidden), nn.SiLU(),
            nn.Linear(hidden, hidden), nn.SiLU())
        # 节点更新
        self.node_mlp = nn.Sequential(
            nn.Linear(hidden * 2, hidden), nn.SiLU(),
            nn.Linear(hidden, hidden))
        # 坐标更新的标量系数
        self.coord_mlp = nn.Sequential(
            nn.Linear(hidden, hidden), nn.SiLU(),
            nn.Linear(hidden, 1))

    def forward(self, h, x, edge_index):
        row, col = edge_index
        # 相对位置向量与距离(不变量)
        rel = x[row] - x[col]
        dist2 = (rel ** 2).sum(dim=-1, keepdim=True)

        # 边消息只依赖不变量 → 消息是不变的
        m = self.edge_mlp(torch.cat([h[row], h[col], dist2], dim=-1))

        # 【坐标更新:沿相对位置向量方向】
        #   系数是标量(不变),方向是相对向量(等变)
        #   → 更新量是等变的
        coord_w = self.coord_mlp(m)
        agg_x = torch.zeros_like(x).index_add_(0, row, rel * coord_w)
        x_new = x + agg_x / (torch.bincount(row, minlength=x.size(0))
                             .clamp(min=1).unsqueeze(-1))

        # 节点特征更新(不变)
        agg_h = torch.zeros_like(h).index_add_(0, row, m)
        h_new = h + self.node_mlp(torch.cat([h, agg_h], dim=-1))
        return h_new, x_new

# 【验证等变性 —— 必做的单元测试】
def test_equivariance(layer, h, x, edge_index, tol=1e-4):
    torch.manual_seed(0)
    # 随机旋转矩阵
    A = torch.randn(3, 3)
    Q, _ = torch.linalg.qr(A)
    if torch.det(Q) < 0:
        Q[:, 0] = -Q[:, 0]
    t = torch.randn(3)

    h1, x1 = layer(h, x, edge_index)
    h2, x2 = layer(h, x @ Q.T + t, edge_index)

    print("特征不变性误差:", (h1 - h2).abs().max().item())
    print("坐标等变性误差:", (x1 @ Q.T + t - x2).abs().max().item())
    # 【两者都应该接近 0】
    # 这个测试能立刻发现实现错误

在药物发现中的应用

应用 需要的对称性
分子性质预测(从三维构象) 不变性
结合姿势预测 等变性(见 102《EquiBind 论文精读》101《DiffDock 论文精读》
蛋白结构预测 等变性(见 112《AlphaFold2 论文精读》
基于结构的分子生成 等变性
力场/势能面学习 能量不变,力等变
构象生成 等变性(见 049《构象生成入门》
蛋白设计 等变性(见 151《RFdiffusion 论文精读》

选型的实际判断

# 【第一个问题】:真的需要三维吗?
#
#   很多分子性质主要由二维结构决定
#   → 用三维模型只会增加成本与过拟合风险
#   → 【先用二维基线验证三维是否必要】(见 142)
#
# 【第二个问题】:需要不变还是等变?
#
#   只预测标量 → 不变量网络就够(SchNet 类)
#     更简单、更快、更容易训练
#   需要输出坐标/向量 → 必须等变
#
# 【第三个问题】:需要多高的角度分辨率?
#
#   只用距离(SchNet):
#     无法区分某些不同的三维排布
#     → 对角度敏感的性质不够
#   加入角度(DimeNet):
#     表达力更强,成本上升
#   完整的球谐展开(e3nn):
#     理论完备,但慢很多
#
#   → 【实践中 EGNN 或 PaiNN 的性价比通常最好】
#
# 【第四个问题】:构象从哪来?
#
#   三维模型的输入需要构象
#   → 构象生成的质量直接影响结果(见 049)
#   → 【必须固定构象生成流程,否则结果不可复现】
#   → 单构象 vs 多构象平均,是重要的设计选择

# 【常见的实现错误】:
#   1) 用绝对坐标而非相对坐标 → 破坏平移不变性
#   2) 在等变层后接普通 MLP 处理坐标 → 破坏等变性
#   3) 用了坐标的某个分量(如 x 坐标)作为特征 → 破坏旋转不变性
#   4) 归一化坐标时用了全局参考系 → 可能破坏对称性
#
#   → 【每次实现都应该跑等变性单元测试】

关键要点

  • 不变性用于标量预测,等变性用于向量/坐标预测——先分清需要哪种;
  • 架构保证的对称性比数据增强学到的样本效率高得多,这在数据稀缺的化学中尤其关键;
  • 每次实现都应跑等变性单元测试——旋转平移后检查输出,能立刻发现实现错误;
  • 先用二维基线验证「三维是否必要」;构象生成流程必须固定,否则结果不可复现。

延伸资源