三维图网络 三维图网络(3D graph network)把 GNN 扩展到带有空间几何的数据上,其中旋转和平移必须被正确处理。本文件涵盖几何图、SE(3)/E(n) 等变性、SchNet、DimeNet、EGNN、张量场网络(tensor field network),以及在分子性质预测、蛋白质结构、材料科学和药物发现中的应用——这些架构学会了从三维物理世界中学习。 文件 03 和文件 04 里的 GNN 作用在抽象的图上:节点有特征,边编码连接,但没有三维空间的概念。社交网络图没有几何。但 GNN 最有影响力的许多应用涉及活在物理三维空间里的数据:分子、蛋白质、晶体、点云。对这些来说,节点的空间位置携带了关键信息,而抽象 GNN 把它忽略了。
三维图网络(3D graph network)把 GNN 扩展到带有空间几何的数据上,其中旋转和平移必须被正确处理。本文件涵盖几何图、SE(3)/E(n) 等变性、SchNet、DimeNet、EGNN、张量场网络(tensor field network),以及在分子性质预测、蛋白质结构、材料科学和药物发现中的应用——这些架构学会了从三维物理世界中学习。
文件 03 和文件 04 里的 GNN 作用在抽象的图上:节点有特征,边编码连接,但没有三维空间的概念。社交网络图没有几何。但 GNN 最有影响力的许多应用涉及活在物理三维空间里的数据:分子、蛋白质、晶体、点云。对这些来说,节点的空间位置携带了关键信息,而抽象 GNN 把它忽略了。
挑战在于三维数据具有几何对称性(文件 01):旋转一个分子不会改变它的性质,平移也不会。一个三维 GNN 必须尊重这些对称性。一个在你旋转分子时会改变的能量预测,在物理上是错的。
**几何图(geometric graph)**是嵌入在三维空间中的图。每个节点 i 除了特征向量 \mathbf{h}_i 之外,还有一个位置 \mathbf{r}_i \in \mathbb{R}^3。边可以由空间邻近性(连接距离在 r_{\text{cut}} 以内的节点)来定义,而不是由显式的化学键。
对分子而言,几何图的节点是原子(特征包括元素类型、电荷等),边是化学键。三维位置 \mathbf{r}_i 是原子坐标,由量子力学或实验测量(X 射线晶体学、冷冻电镜)确定。
对点云(来自 LiDAR 或三维扫描仪,第 8 章和第 11 章)而言,每个点是一个带有位置和可选特征(颜色、强度)的节点。边连接相近的点,形成一个 k 近邻(k-nearest-neighbour,kNN)图或一个半径图。
用于消息传递的关键几何量有:
原子间距离(interatomic distances):d_{ij} = \|\mathbf{r}_i - \mathbf{r}_j\|。距离对旋转和平移不变。两个具有相同原子间距离的分子形状相同,与朝向无关。
键角(bond angles):在节点 i 处向量 \mathbf{r}_j - \mathbf{r}_i 与 \mathbf{r}_k - \mathbf{r}_i 之间的夹角 \theta_{ijk}。键角捕捉了超越两两距离的局部几何。
二面角(dihedral / torsion angles):由 (i, j, k) 和 (j, k, l) 定义的两个平面之间的夹角 \phi_{ijkl}。二面角捕捉结构在三维中的扭转,对蛋白质骨架几何至关重要。
相对位置向量(relative position vectors):\mathbf{r}_{ij} = \mathbf{r}_j - \mathbf{r}_i。这些是平移不变的,但不是旋转不变的。使用它们需要等变(不只是不变)的架构。
三维物理数据对应的对称群是欧几里得群(Euclidean group) E(3),由所有旋转、反射和平移组成。子群 SE(3)(特殊欧几里得群,Special Euclidean)包括旋转和平移,但不含反射。
一个三维 GNN 应当:
这些约束直接对应文件 01 中的不变性/等变性框架,现在专门应用到三维的旋转和平移群上。
这里有两类设计思路:
SchNet(Schütt 等,2017)是奠基性的不变三维 GNN。它的关键创新是连续滤波器卷积(continuous filter convolution):不再使用一组固定的边类型(像分子 GNN 里的化学键类型),SchNet 直接从原子间距离生成消息滤波器。
距离 d_{ij} 先用**径向基函数(radial basis functions,RBFs)**展开成一个特征向量:
每个基函数是一个以 \mu_k 为中心、宽度为 \gamma_k 的高斯。这类似于距离的可学习位置编码:连续的距离被映射到一个高维特征空间,网络可以在其中学习依赖于距离的相互作用。中心 \mu_k 通常从 0 到截断半径均匀分布。
SchNet 中从节点 j 到节点 i 的消息是:
其中 W_{\text{filter}} 是一个把 RBF 展开映射成滤波器向量的 MLP,\odot 是逐元素乘法(Hadamard 积,第 2 章)。滤波器依赖于距离,所以近处的原子与远处的原子相互作用方式不同。逐元素乘法就像一个门控机制(第 6 章):依赖于距离的滤波器控制着每个特征维度有多少能通过。
因为 SchNet 只用距离(不变的),整个模型自动对旋转和平移不变。除了这个设计选择之外,不需要对对称性做任何特殊处理。
单靠距离无法完全确定三维结构。两种不同的分子构象可以有完全相同的两两距离,但键角不同(这是"距离几何歧义"问题)。DimeNet(Gasteiger 等,2020)把键角纳入消息传递。
DimeNet 使用有向消息传递(directional message passing):消息沿有向边流动,边 (j \to i) 上的消息受边 (k \to j) 与 (j \to i) 之间夹角的影响:
角度 \theta_{kji} 用球面 Bessel 函数和球面调和函数展开(它们是球面上角度信息的自然基,类似于距离的 RBF)。这让模型在保持不变性的同时获得了方向信息。
SphereNet(Liu 等,2022)更进一步,纳入二面角 \phi_{lkji},捕捉完整的三维扭转结构。层级是:
每提高一层都增加几何分辨率,代价是计算复杂度(距离是 O(|E|),角度是 O(|E| \cdot k),二面角是 O(|E| \cdot k^2),其中 k 是平均度)。
EGNN(Satorras 等,2021)采取等变思路:它不只用不变特征,而是在每一层同时更新节点特征和节点位置,全程保持等变性。
EGNN 对节点 i 的更新是:
关键在于位置更新:节点位置由相对位置向量 (\mathbf{r}_i - \mathbf{r}_j) 的加权和来调整。权重来自消息函数 \phi_r,而它只依赖于不变的量(特征和距离)。这种构造可证明是等变的:如果所有输入位置都被 R 旋转,所有输出位置也会被同一个 R 旋转。
EGNN 优雅的地方在于,它在不显式使用球面调和函数或不可约表示的情况下就实现了等变性。相对位置向量携带方向信息,而不变的消息函数控制这些方向信息如何被使用。
这种简洁性是有代价的:EGNN 只使用向量表示(1 阶)。如果不加扩展,它无法表示像四极矩或应力张量这样的高阶张量。
张量场网络(Tensor Field Networks)(Thomas 等,2018)及其后继者(SE(3)-Transformers、MACE、Equiformer)使用旋转群的**不可约表示(irreducible representations)**的完整工具箱来构建等变层。
在表示论中(联系到第 2 章的线性代数),三维中的旋转可以分解为由整数阶 \ell 刻画的不可约成分:
这些叫做球张量(spherical tensors),它们在旋转 R 下通过 Wigner-D 矩阵 D^\ell(R) 变换:标量不变,向量按 R 旋转,2 阶张量按一个更复杂的矩阵旋转。
使用球张量的等变消息传递用 Clebsch-Gordan 张量积来组合不同阶的特征:
Clebsch-Gordan 系数 C 是固定的数学常数,保证张量积是等变的。这是矩阵乘法的 SO(3) 等变类比。
MACE(Batatia 等,2022)使用高阶消息(多个邻居特征的乘积),以更少的消息传递层数达到高精度。通过构造多体相互作用(距离给出的 2 体、角度给出的 3 体、张量积给出的多体),MACE 高效地捕捉了复杂的原子相互作用。
Equiformer(Liao 和 Smidt,2023)把等变的球张量特征与 Transformer 注意力机制(文件 04)结合起来,创建了一个 SE(3) 等变的图 Transformer。注意力分数从不变特征算出,而 value 聚合则作用于等变的张量特征。
分子性质预测:给定一个分子的三维结构,预测能量、力、偶极矩、HOMO-LUMO 能隙、毒性、溶解度等性质。这是三维 GNN 最成熟的应用。在量子化学数据集(QM9、OC20)上训练的模型在许多性质上达到了化学精度,使得对数百万候选分子的虚拟筛选成为可能。
分子动力学加速:用量子力学(密度泛函理论,density functional theory,DFT)计算原子间的力极其昂贵(对 n 个电子是 O(n^3))。一个训练来预测力的三维 GNN 可以在分子动力学模拟中替代 DFT,在保持接近 DFT 精度的同时实现 10^3 到 10^6 倍的加速。这使得更大系统、更长时间尺度的模拟成为可能,揭示传统方法看不到的现象。
蛋白质结构:蛋白质是折叠成复杂三维结构的氨基酸链。蛋白质骨架是一个几何图,节点是残基,边连接空间上相近的残基。三维 GNN 用于蛋白质功能预测、结合位点识别和蛋白质设计(反向折叠:给定一个目标结构,预测氨基酸序列)。AlphaFold 使用几何和基于图的推理来从序列预测蛋白质结构。
材料科学与催化:晶体材料具有周期性的三维结构。GNN 建模重复的单胞并预测材料性质:带隙、形成能、机械强度。Open Catalyst Project(OC20/OC22)基准测试 GNN 在催化表面上预测吸附能的能力,加速为可再生能源寻找新催化剂。
药物发现:三维 GNN 预测药物分子将如何与目标蛋白结合。结合亲和力依赖于药物与蛋白结合口袋之间的三维形状互补性和化学相互作用。像 DiffDock 这样的模型把等变 GNN 与扩散模型(第 8 章)结合起来,预测结合姿态(药物在蛋白口袋里的三维朝向)。
上面所有的架构都是分析已有的图。**图生成(graph generation)**则创造新的图:设计一个具有所需性质的分子、生成用于测试的合成社交网络、或提出一种新的蛋白质结构。这是图级预测的生成式对应物。
挑战在于图是离散的、大小可变、且组合爆炸的。生成一张图意味着决定要创建多少节点、它们有什么特征、以及连接哪些节点对。可能的图的空间随节点数超指数增长。
**自回归生成(autoregressive generation)**一次一个节点(或一条边)地构建图。GraphRNN(You 等,2018)顺序地生成图:一个 RNN 维护一个状态,每一步生成一个新节点,并决定把它连到哪些已有节点上。生成顺序对本质上无序的图强加了一种人为的序列,但 BFS 顺序能让最近生成的节点保持相关性,从而有所帮助。
**基于 VAE 的生成(VAE-based generation)**把图编码到一个连续的潜在空间(用一个 GNN 编码器),再从采样的潜在向量解码出新的图。GraphVAE 一次性生成一个概率邻接矩阵 \hat{A} \in [0, 1]^{n \times n},但这会以 O(n^2) 增长,并且产生必须经过阈值化的稠密输出。潜在空间允许平滑插值:在两个分子嵌入之间移动会生成化学上有效的中间结构。
**基于扩散的生成(diffusion-based generation)**把扩散框架(第 8 章)应用到图上。前向过程逐步给节点特征和边结构加噪。反向过程学习去噪,从噪声中生成有效的图。DiGress(Vignac 等,2023)对节点类型和边类型都应用离散扩散,自然地处理了图数据的类别属性。
对于分子生成,关键约束是化学有效性:生成的分子必须服从化合价规则(碳形成 4 个键,氧形成 2 个,等等)。像 **Junction Tree VAE(JT-VAE)**这样的方法把分子分解成有效的子结构(环、链、官能团),再通过组装这些构件来生成,从构造上保证有效性。
**目标导向生成(goal-directed generation)**针对特定性质进行优化:生成一个对目标蛋白有高结合亲和力、低毒性且溶解度好的分子。这把图生成与性质预测(用一个三维 GNN 作为性质评估器)放在一个循环里:生成 → 评估 → 改进。强化学习(第 6 章)或贝叶斯优化引导在化学空间中的搜索。
DiffDock(Corso 等,2023)使用 SE(3) 等变扩散来预测药物分子如何对接到蛋白结合口袋中。该模型通过从随机放置开始去噪来生成三维结合姿态(药物相对于蛋白的位置和朝向),把本文件的三维等变网络与第 8 章的扩散框架结合起来。
import jax import jax.numpy as jnp # 水分子:O 在原点,两个 H 原子 positions = jnp.array([[0.0, 0.0, 0.0], # O [0.96, 0.0, 0.0], # H1 [-0.24, 0.93, 0.0]]) # H2 # 节点特征:[原子序数] features = jnp.array([[8.0], [1.0], [1.0]]) # 计算两两距离(不变) def pairwise_distances(pos): diff = pos[:, None, :] - pos[None, :, :] return jnp.sqrt(jnp.sum(diff**2, axis=-1) + 1e-8) # 简单的基于距离的消息传递 def invariant_message_pass(features, positions): dists = pairwise_distances(positions) # 用 4 个中心做 RBF 展开 centres = jnp.array([0.5, 1.0, 1.5, 2.0]) rbf = jnp.exp(-5.0 * (dists[:, :, None] - centres[None, None, :]) ** 2) # 消息:用依赖距离的滤波器给特征加权 messages = jnp.einsum("ij,jd->id", rbf.sum(axis=-1), features) return messages output1 = invariant_message_pass(features, positions) # 把分子绕 z 轴旋转 90 度 R = jnp.array([[0, -1, 0], [1, 0, 0], [0, 0, 1]], dtype=float) rotated_positions = (R @ positions.T).T output2 = invariant_message_pass(features, rotated_positions) print(f"Original output:\n{output1}") print(f"\nRotated output:\n{output2}") print(f"\nInvariant: {jnp.allclose(output1, output2, atol=1e-5)}")
import jax.numpy as jnp def bond_angle(r_i, r_j, r_k): """在节点 j 处,边 j->i 与 j->k 之间的夹角。""" v1 = r_i - r_j v2 = r_k - r_j cos_angle = jnp.dot(v1, v2) / (jnp.linalg.norm(v1) * jnp.linalg.norm(v2)) return jnp.arccos(jnp.clip(cos_angle, -1, 1)) # 三个原子 r1 = jnp.array([1.0, 0.0, 0.0]) r2 = jnp.array([0.0, 0.0, 0.0]) r3 = jnp.array([0.0, 1.0, 0.0]) angle_original = bond_angle(r1, r2, r3) print(f"Original angle: {jnp.degrees(angle_original):.1f}°") # 施加一个随机旋转 R = jnp.array([[0.36, 0.48, -0.80], [-0.80, 0.60, 0.00], [0.48, 0.64, 0.60]]) r1_rot, r2_rot, r3_rot = R @ r1, R @ r2, R @ r3 angle_rotated = bond_angle(r1_rot, r2_rot, r3_rot) print(f"Rotated angle: {jnp.degrees(angle_rotated):.1f}°") print(f"Invariant: {jnp.allclose(angle_original, angle_rotated, atol=1e-4)}")
import jax import jax.numpy as jnp def egnn_position_update(positions, features): """简单的 EGNN 风格等变位置更新。""" n = positions.shape[0] new_positions = jnp.zeros_like(positions) for i in range(n): shift = jnp.zeros(3) for j in range(n): if i != j: r_ij = positions[i] - positions[j] d_ij = jnp.linalg.norm(r_ij) # 按距离加权(简单做法:距离的倒数) weight = 1.0 / (d_ij + 1.0) # 再按特征相似度缩放 feat_sim = jnp.dot(features[i], features[j]) shift = shift + weight * feat_sim * r_ij new_positions = new_positions.at[i].set(positions[i] + 0.1 * shift) return new_positions # 3 个原子 pos = jnp.array([[0.0, 0.0, 0.0], [1.0, 0.0, 0.0], [0.0, 1.0, 0.0]]) feat = jnp.array([[1.0, 0.5], [0.5, 1.0], [0.8, 0.3]]) # 更新位置 pos_new = egnn_position_update(pos, feat) # 现在旋转输入、再更新,并检查输出是否相应地被旋转 R = jnp.array([[0.0, -1.0, 0.0], [1.0, 0.0, 0.0], [0.0, 0.0, 1.0]]) pos_rot = (R @ pos.T).T pos_new_from_rot = egnn_position_update(pos_rot, feat) # 应该与把原始输出旋转一下相同 pos_new_then_rot = (R @ pos_new.T).T print(f"Update then rotate:\n{jnp.round(pos_new_then_rot, 4)}") print(f"\nRotate then update:\n{jnp.round(pos_new_from_rot, 4)}") print(f"\nEquivariant: {jnp.allclose(pos_new_then_rot, pos_new_from_rot, atol=1e-4)}")