176

PyTorch Geometric:GNN 分子建模的基础框架

PyTorch Geometric 是图深度学习的主流框架,也是自己实现分子 GNN 的基础设施。这篇讲清它的数据结构、批处理机制、消息传递写法,以及分子建模里的具体用法。

PyTorch Geometric(PyG) 是图深度学习最主流的框架。它本身不是化学工具,但几乎所有需要自己动手改结构的分子 GNN 工作都建在它上面——想理解或修改 D-MPNN、SchNet、AttentiveFP 这类模型,PyG 是绕不开的一层。

安装

pip install torch_geometric
# 可选加速算子(按 torch/CUDA 版本选 wheel 源)
pip install torch_scatter torch_sparse \
  -f https://data.pyg.org/whl/torch-2.4.0+cu121.html

新版 PyG 的多数功能不再强依赖 torch_scatter,但装上通常更快。版本不匹配是最常见的安装故障。

核心数据结构

from torch_geometric.data import Data
import torch

data = Data(
    x=torch.randn(9, 32),                       # 9 个原子,每个 32 维特征
    edge_index=torch.tensor([[0, 1], [1, 0]]).t().contiguous(),  # [2, E] 稀疏边
    edge_attr=torch.randn(2, 8),                # 每条边 8 维(键类型等)
    y=torch.tensor([1.2]),                      # 图级标签
)

关键在 edge_index[2, E] 稀疏格式:不存邻接矩阵,只存边的端点对。分子平均只有二三十个原子、连接稀疏,这种表示比稠密矩阵省得多。无向键要存成两条有向边,这是新手最容易漏的一点。

批处理:为什么分子图能高效并行

PyG 把一批分子拼成一张不连通的大图,节点特征直接拼接,edge_index 加上偏移量,另用 batch 向量记录每个原子属于哪个分子。这样无需 padding 就能并行,再用 global_mean_pool(x, batch) 按分子聚合回图级表示。理解这个机制,读任何分子 GNN 代码都会顺畅很多。

from torch_geometric.loader import DataLoader
from torch_geometric.nn import GINEConv, global_mean_pool
import torch.nn as nn

loader = DataLoader(dataset, batch_size=64, shuffle=True)

class MolGNN(nn.Module):
    def __init__(self, in_dim, edge_dim, hid=256):
        super().__init__()
        mlp = nn.Sequential(nn.Linear(in_dim, hid), nn.ReLU(), nn.Linear(hid, hid))
        self.conv1 = GINEConv(mlp, edge_dim=edge_dim)
        self.head = nn.Linear(hid, 1)
    def forward(self, d):
        h = self.conv1(d.x, d.edge_index, d.edge_attr).relu()
        h = global_mean_pool(h, d.batch)     # 原子级 → 分子级
        return self.head(h)

分子建模里怎么用

  • 从 RDKit 转 PyG:用 torch_geometric.utils.from_smiles 可以一行拿到带基础原子/键特征的 Data;正经项目通常自己写特征化,控制原子类型、形式电荷、手性、环信息等。
  • 选卷积层:分子任务上 GINEConv(能用边特征)和 NNConv 是常用起点;带 3D 坐标时用 SchNetDimeNet++ 这类等变/距离感知模型。
  • 池化方式影响不小global_add_pool 对分子大小敏感(隐含了尺寸信息),global_mean_pool 尺寸无关。预测溶解度这类与分子大小强相关的性质时,两者差别明显,值得都试。
  • 预置数据集torch_geometric.datasets 里有 MoleculeNet、QM9、ZINC 等,适合快速验证结构改动。

什么时候不该用它

如果任务只是「给一批 SMILES 预测一个数值」,直接用 Chemprop(见 174《Chemprop》)比自己搭 PyG 模型快得多,效果通常还更好。PyG 的价值在于要改架构——加新的消息传递规则、融合 3D 信息、做多模态图。为了标准任务从零搭 GNN,多半是在重复造一个更差的轮子。

上手提示

  • 先吃透 edge_index 稀疏格式和 batch 向量拼图机制,这是理解一切分子 GNN 代码的钥匙;
  • 无向键必须存双向边;
  • 池化方式(sum vs mean)对与分子大小相关的性质影响很大,务必对比;
  • 标准任务先用 Chemprop,需要改架构时才动 PyG。

延伸资源