分子本质上就是图:原子是节点,化学键是边。这个对应关系让图神经网络成为分子建模的自然选择——但「图怎么建、特征怎么设计」对结果的影响,往往大于选哪个 GNN 架构。
图的构建
from rdkit import Chem
import torch
from torch_geometric.data import Data
def mol_to_graph(smiles, y=None):
mol = Chem.MolFromSmiles(smiles)
if mol is None:
return None
# ---- 节点(原子)特征 ----
x = []
for atom in mol.GetAtoms():
x.append(atom_features(atom))
# ---- 边(化学键)----
edge_index, edge_attr = [], []
for bond in mol.GetBonds():
i = bond.GetBeginAtomIdx()
j = bond.GetEndAtomIdx()
f = bond_features(bond)
# 【关键:无向图必须存两个方向】
edge_index += [[i, j], [j, i]]
edge_attr += [f, f]
return Data(
x=torch.tensor(x, dtype=torch.float),
edge_index=torch.tensor(edge_index, dtype=torch.long).t().contiguous(),
edge_attr=torch.tensor(edge_attr, dtype=torch.float),
y=torch.tensor([y], dtype=torch.float) if y is not None else None,
smiles=smiles,
)
# 【处理没有键的分子】(如单原子离子):
# edge_index 会是空的 → 需要特殊处理
# edge_index=torch.empty((2, 0), dtype=torch.long)
原子特征的设计
def one_hot(value, choices):
"""独热编码,未知值归入最后一类"""
enc = [0] * (len(choices) + 1)
idx = choices.index(value) if value in choices else len(choices)
enc[idx] = 1
return enc
ATOM_TYPES = ["C", "N", "O", "S", "F", "Cl", "Br", "I", "P", "B", "Si"]
HYBRIDIZATIONS = [Chem.HybridizationType.SP,
Chem.HybridizationType.SP2,
Chem.HybridizationType.SP3,
Chem.HybridizationType.SP3D,
Chem.HybridizationType.SP3D2]
def atom_features(atom):
return (
one_hot(atom.GetSymbol(), ATOM_TYPES) # 元素类型
+ one_hot(atom.GetDegree(), [0, 1, 2, 3, 4, 5]) # 连接数
+ one_hot(atom.GetFormalCharge(), [-2, -1, 0, 1, 2])
+ one_hot(atom.GetTotalNumHs(), [0, 1, 2, 3, 4]) # 氢数
+ one_hot(atom.GetHybridization(), HYBRIDIZATIONS)
+ [int(atom.GetIsAromatic())]
+ [int(atom.IsInRing())]
+ [int(atom.IsInRingSize(3)), int(atom.IsInRingSize(4)),
int(atom.IsInRingSize(5)), int(atom.IsInRingSize(6)),
int(atom.IsInRingSize(7))]
+ [atom.GetMass() * 0.01] # 【归一化】
+ one_hot(int(atom.GetChiralTag()), [0, 1, 2, 3]) # 手性
)
# 【最常见的特征设计错误】:
#
# 1) 【用原始数字而非独热编码】
# atom.GetAtomicNum() 直接当特征
# → 模型会以为「碳(6) 与氮(7) 数值上接近 = 化学上相似」
# → 【这是错误的归纳偏置】
# → 【必须用独热编码】
#
# 2) 【忘记归一化连续特征】
# 原子质量(1~200)与独热编码(0/1)尺度差异巨大
# → 训练不稳定
#
# 3) 【遗漏环信息】
# 环的大小对分子性质影响很大
# 而消息传递【不能直接感知环】
# → 必须作为特征显式提供
#
# 4) 【手性未编码】
# 对映体的性质可能完全不同
# → 做手性相关的任务时必须加
#
# 5) 【未知值处理】
# 测试集出现训练集没有的元素
# → one_hot 要有「其它」类别
键特征
BOND_TYPES = [Chem.BondType.SINGLE, Chem.BondType.DOUBLE,
Chem.BondType.TRIPLE, Chem.BondType.AROMATIC]
STEREO_TYPES = [Chem.BondStereo.STEREONONE, Chem.BondStereo.STEREOZ,
Chem.BondStereo.STEREOE, Chem.BondStereo.STEREOCIS,
Chem.BondStereo.STEREOTRANS]
def bond_features(bond):
return (
one_hot(bond.GetBondType(), BOND_TYPES)
+ [int(bond.GetIsConjugated())]
+ [int(bond.IsInRing())]
+ one_hot(bond.GetStereo(), STEREO_TYPES)
)
# 【键特征常被忽略但很重要】:
# - 共轭:影响电子分布与刚性
# - 成环:影响构象自由度
# - 双键的顺反:影响三维形状
#
# 【使用键特征需要支持的 GNN 层】:
# GINEConv、NNConv、MPNN 支持边特征
# GCNConv、GATConv(基础版)不支持
# → 【选层时要注意】
图表示的固有局限
# 【局限一:不含三维信息】
# 拓扑图只知道「哪些原子相连」
# 不知道它们在空间中的相对位置
# → 对依赖三维的性质表达不足
# → 【应对】:
# - 加入三维坐标 + 等变网络(见 160)
# - 加入距离/角度作为边特征
#
# 【局限二:不能直接感知环】
# 消息传递是局部的
# → 【六元环与七元环,在局部看起来一样】
# → 【应对】:把环信息作为原子/键特征
#
# 【局限三:表达能力有上限】
# 标准 GNN 的判别力不超过 1-WL 测试(见 159)
# → 【某些不同的分子图无法区分】
# → 【应对】:加入子结构计数、环特征等
#
# 【局限四:全局性质表达低效】
# 分子量、logP 需要多层聚合才能间接学到
# → 【应对】:【直接拼接 RDKit 全局描述符】
# 这是最有效也最简单的改进(见 131)
#
# 【局限五:氢原子通常被隐含】
# 隐式氢丢失了氢的位置信息
# → 多数任务可接受
# → 但对氢键相关的任务可能有影响
# → 【应对】:Chem.AddHs() 后再建图(代价是图变大)
图 vs 指纹:实际的比较
| 分子图 + GNN | ECFP 指纹 | |
|---|---|---|
| 信息范围 | K 层 = K 跳 | 半径 K |
| 本质 | 学习聚合 | 枚举 + 哈希 |
| 小数据表现 | 容易过拟合 | 更稳健 |
| 大数据表现 | 更好 | 受限 |
| 原子级输出 | 支持 | 不支持 |
| 三维扩展 | 可以 | 需要专门的三维指纹 |
| 训练成本 | 高(需 GPU) | 低 |
| 可解释性 | 需要归因方法 | 可追溯到具体子结构 |
实践建议:先跑 ECFP + LightGBM 基线(见 131《Chemprop / D-MPNN 论文精读》)。只有在数据量足够(> 5000)或需要原子级输出/三维建模时,GNN 才有明显优势。
批处理与常见实现问题
from torch_geometric.loader import DataLoader
# PyG 的批处理:把多个图拼成一个大图
dataset = [mol_to_graph(s, y) for s, y in zip(smiles_list, labels)]
dataset = [d for d in dataset if d is not None] # 【过滤解析失败的】
loader = DataLoader(dataset, batch_size=32, shuffle=True)
for batch in loader:
print(batch.x.shape) # 所有图的节点拼在一起
print(batch.batch) # 【指示每个节点属于哪个图】
print(batch.num_graphs)
break
# 【常见错误】:
# 1) 池化时忘记传 batch 参数
# global_mean_pool(x, batch) ← 【必须传 batch】
# 忘了会把所有图混在一起
#
# 2) 忘记过滤 None(解析失败的分子)
#
# 3) 单原子分子的 edge_index 为空
# → 某些层会报错,需要特殊处理
#
# 4) 特征维度不一致
# → one_hot 的类别数必须固定
关键要点
- 原子类型必须用独热编码——用原子序数会引入「碳与氮数值接近所以相似」的错误偏置;
- 无向图必须存两个方向的边,否则消息只单向传递;
- 环信息必须显式作为特征——消息传递不能直接感知环;
- 拼接 RDKit 全局描述符是最有效的简单改进;池化时不要忘记传 batch。
延伸资源
- GNN:159《图神经网络 GNN》;等变网络:160《等变神经网络》;Chemprop:131《Chemprop / D-MPNN 论文精读》;
- 分子指纹:035《分子指纹是什么》;CS224W:006《Stanford CS224W 图神经网络》;PyTorch Geometric:176《PyTorch Geometric》。