处理分子的三维坐标时,模型必须尊重物理的对称性:把整个分子旋转或平移,它的性质不应改变,而预测的坐标应该跟着变。把这个约束编码进架构,而非指望模型从数据中学会,是几何深度学习的核心思想。
不变性与等变性
# 【不变性 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) 归一化坐标时用了全局参考系 → 可能破坏对称性
#
# → 【每次实现都应该跑等变性单元测试】
关键要点
- 不变性用于标量预测,等变性用于向量/坐标预测——先分清需要哪种;
- 架构保证的对称性比数据增强学到的样本效率高得多,这在数据稀缺的化学中尤其关键;
- 每次实现都应跑等变性单元测试——旋转平移后检查输出,能立刻发现实现错误;
- 先用二维基线验证「三维是否必要」;构象生成流程必须固定,否则结果不可复现。
延伸资源
- GNN:159《图神经网络 GNN》;EquiBind:102《EquiBind 论文精读》;DiffDock:101《DiffDock 论文精读》;
- AlphaFold2:112《AlphaFold2 论文精读》;Uni-Mol:142《Uni-Mol 论文精读》;构象生成:049《构象生成入门》。