简介:这份资源面向计算化学、材料科学与机器学习交叉方向的学习者,提供了一套基于图神经网络(GNN)预测分子能量的完整Python实现方案。分子被抽象为图结构,原子作节点、化学键作边,节点与边特征编码原子类型、键级等化学信息,模型通过多轮消息传递聚合局部环境,再经全连接层输出标量能量值,适合作为入门GNN与分子性质预测的实战案例。压缩包共33个文件、约7.13MB,以8个py源码、7个csv数据集、3个pt模型文件为主,另含zbak备份、mol分子结构、txt说明、png结果图与md文档,覆盖数据加载、图构建、模型定义、训练循环到结果可视化的全流程。资源已预先划分训练、验证与测试集,代码注释完整,可通过配置文件调整超参数或替换自有数据做迁移学习。目前已有128人学习,适合希望快速上手分子能量预测、理解图神经网络消息传递机制的自学者参考。
1. 分子能量预测为什么值得用 GNN 重做一遍
做计算化学或者药物筛选的朋友大概率都碰过这个场景:手头有几千到几十万个分子,想快速估一下它们的能量或者某个量子化学性质,DFT 算一遍动辄几小时到几天,成本根本扛不住。传统做法是退而求其次用分子指纹加 XGBoost、随机森林这类模型,但指纹本质是把分子拍扁成一维向量,键长、键角、环结构、官能团之间的空间关系全丢了。分子能量预测模型要真正做准,必须让模型看到原子之间的连接拓扑,这正是图神经网络(GNN)的用武之地——把原子当节点、化学键当边,分子天然就是一张图。
这篇笔记讲的就是怎么用 Python 从零搭一个基于 GNN 的分子能量预测模型,包括数据集的读取与图化、模型结构选型、训练参数怎么调、以及上线前怎么验证。适合已经会写 Python、懂一点深度学习、但还没系统做过图神经网络的从业者。读完你应该能自己跑通一条从 SMILES 到能量预测值的完整链路,并且知道哪些参数一动就会翻车。
2. 从 SMILES 到图张量:分子图构建的完整链路
2.1 为什么分子必须转成图结构
分子能量本质上是原子核位置和电子结构的函数。一个分子里,每个原子有自己的元素类型、杂化状态、形式电荷,原子之间通过化学键连接,键有单键、双键、芳香键之分。如果把这些信息塞进一个固定长度的向量,分子越大信息损失越严重。而图结构天然支持变长:节点数等于原子数,边数等于化学键数,每个节点和边都可以带多维特征。GNN 的消息传递机制就是让每个原子不断聚合邻居原子的信息,几轮之后每个原子的表示就编码了它周围的局部化学环境,最后池化得到整个分子的表示,再回归到能量值。
常见做法是用 RDKit 读 SMILES,然后手动抽原子特征和边特征。原子特征一般包括元素类型(one-hot 或 embedding)、度、形式电荷、手性、是否在环内、杂化方式。边特征包括键类型、是否共轭、是否在环内。这些特征维度不高,但缺一个都可能让模型在某些子结构上系统性偏差。
2.2 用 RDKit 构建分子图的代码实现
下面这段代码把一条 SMILES 转成节点特征矩阵、边索引和边特征矩阵。我一般会把原子特征维度控制在 30 以内,边特征维度 10 左右,太大容易过拟合小数据集。
import numpy as np import torch from rdkit import Chem # 允许的特征取值集合 ATOM_TYPES = ['C', 'N', 'O', 'S', 'F', 'Cl', 'Br', 'P', 'I', 'B', 'Si', 'Se'] HYBRID_TYPES = [ Chem.rdchem.HybridizationType.SP, Chem.rdchem.HybridizationType.SP2, Chem.rdchem.HybridizationType.SP3, Chem.rdchem.HybridizationType.SP3D, Chem.rdchem.HybridizationType.SP3D2, ] BOND_TYPES = [ Chem.rdchem.BondType.SINGLE, Chem.rdchem.BondType.DOUBLE, Chem.rdchem.BondType.TRIPLE, Chem.rdchem.BondType.AROMATIC, ] def one_hot(value, choices): vec = [0] * (len(choices) + 1) # 最后一位留给未知类别 if value in choices: vec[choices.index(value)] = 1 else: vec[-1] = 1 return vec def atom_features(atom): return np.array( one_hot(atom.GetSymbol(), ATOM_TYPES) + one_hot(atom.GetTotalDegree(), list(range(6))) + one_hot(atom.GetFormalCharge(), [-2, -1, 0, 1, 2]) + one_hot(atom.GetHybridization(), HYBRID_TYPES) + [int(atom.GetIsAromatic()), int(atom.IsInRing())], dtype=np.float32 ) def bond_features(bond): return np.array( one_hot(bond.GetBondType(), BOND_TYPES) + [int(bond.GetIsConjugated()), int(bond.IsInRing())], dtype=np.float32 ) def smiles_to_graph(smiles): mol = Chem.MolFromSmiles(smiles) if mol is None: return None # 加氢很重要,隐式氢会让能量预测偏差很大 mol = Chem.AddHs(mol) atom_feats = [atom_features(a) for a in mol.GetAtoms()] edge_index, edge_feats = [], [] for bond in mol.GetBonds(): i, j = bond.GetBeginAtomIdx(), bond.GetEndAtomIdx() bf = bond_features(bond) edge_index += [[i, j], [j, i]] # 无向图双向建边 edge_feats += [bf, bf] if not edge_index: return None x = torch.tensor(np.stack(atom_feats), dtype=torch.float) ei = torch.tensor(edge_index, dtype=torch.long) ea = torch.tensor(np.stack(edge_feats), dtype=torch.float) return x, ei, ea逻辑说明:one_hot函数给每个类别留了一个未知位,避免遇到训练集没出现过的元素时直接报错。smiles_to_graph里Chem.AddHs是关键一步,很多公开数据集(比如 QM9)的能量标签是在含氢构型下算的,不加氢会导致节点数对不上,模型学到的能量系统性偏低。边索引双向添加是因为大多数 GNN 层默认有向消息传递,不双向建边会丢一半邻居信息。
参数说明:ATOM_TYPES和HYBRID_TYPES可以根据你的数据集调整,如果做的是含金属配合物,得把金属元素加进去。GetTotalDegree的取值范围我设成 0 到 5,超过 5 的归到未知位。边特征里GetBondType对芳香键会返回AROMATIC,RDKit 在 sanitize 之后会自动识别,不用手动 kekulize。
2.3 数据集加载与批处理:别让 padding 拖慢训练
分子图大小不一,不能直接 stack 成一个 batch。常见做法是用 PyTorch Geometric 的DataLoader自动做图拼接,或者自己写一个 collate 函数把多个小图拼成一个大图,用batch向量标记每个节点属于哪个分子。我一般用 PyG,省事且经过大量验证。
from torch_geometric.data import Data, DataLoader def build_dataset(smiles_list, energy_list): data_list = [] for smi, e in zip(smiles_list, energy_list): g = smiles_to_graph(smi) if g is None: continue x, ei, ea = g data_list.append(Data(x=x, edge_index=ei, edge_attr=ea, y=torch.tensor([e], dtype=torch.float))) return data_list dataset = build_dataset(train_smiles, train_energies) loader = DataLoader(dataset, batch_size=64, shuffle=True)逻辑说明:Data对象把节点特征、边索引、边特征、标签打包在一起,DataLoader在取 batch 时会把 64 个分子拼成一张大图,edge_index自动偏移,batch向量自动生成。这样 GPU 利用率比逐个分子跑高很多。
参数说明:batch_size对分子图任务很敏感。分子平均节点数在 20 到 30 时,64 一般能跑满显存;如果分子很大(比如超过 100 个原子),得降到 16 或 32。shuffle=True在训练集上必须开,验证集和测试集要关掉,否则没法复现指标。
3. 模型结构选型:GCN、GIN 还是 SchNet
3.1 三种主流 GNN 层在分子任务上的差异
分子能量预测这个任务,学术界和工业界用得最多的是三类:GCN、GIN 和 SchNet。GCN 是最早的图卷积,聚合方式简单,对节点特征做归一化加权求和,优点是快、稳定,缺点是对不同邻居的重要性不加区分。GIN 在聚合时加了可学习的 epsilon 和多层感知机,表达能力更强,理论上比 GCN 更接近 WL 图同构测试,在分子性质预测上通常比 GCN 高几个点。SchNet 是专门为分子设计的,它把原子间距离作为连续滤波器的输入,适合有 3D 坐标的场景,但如果你只有 2D 拓扑,SchNet 的优势发挥不出来。
我一般会先跑一个 GIN 基线,如果效果不够再考虑上 3D 信息。对于纯 2D 拓扑的能量预测,GIN 加全局注意力池化通常能到不错的精度。
3.2 一个可复现的 GIN 回归模型
下面这个模型用 4 层 GIN,每层后面接 BatchNorm 和 ReLU,最后用全局平均池化加两层全连接输出能量值。
import torch.nn as nn import torch.nn.functional as F from torch_geometric.nn import GINConv, global_mean_pool class GINEnergyModel(nn.Module): def __init__(self, node_dim, edge_dim, hidden=128, num_layers=4): super().__init__() self.convs = nn.ModuleList() self.bns = nn.ModuleList() for i in range(num_layers): in_dim = node_dim if i == 0 else hidden mlp = nn.Sequential( nn.Linear(in_dim, hidden), nn.ReLU(), nn.Linear(hidden, hidden), ) self.convs.append(GINConv(mlp, train_eps=True)) self.bns.append(nn.BatchNorm1d(hidden)) self.fc1 = nn.Linear(hidden, hidden // 2) self.fc2 = nn.Linear(hidden // 2, 1) def forward(self, data): x, edge_index, batch = data.x, data.edge_index, data.batch for conv, bn in zip(self.convs, self.bns): x = conv(x, edge_index) x = bn(x) x = F.relu(x) x = global_mean_pool(x, batch) x = F.relu(self.fc1(x)) return self.fc2(x).squeeze(-1)逻辑说明:GINConv的train_eps=True让 epsilon 可学习,比固定值更灵活。每层后接 BatchNorm 是为了稳定训练,分子图任务里节点特征尺度差异大,不加 BN 很容易梯度爆炸。global_mean_pool把每个分子的节点表示平均成一个向量,这里也可以用 sum 或 attention pool,mean 对小分子更稳。
参数说明:hidden设 128 是经验值,数据集小于 1 万条时可以降到 64 防止过拟合,大于 10 万条可以升到 256。num_layers一般 3 到 5,层数太多会出现过平滑,所有节点表示趋同,反而掉点。fc2输出 1 维是因为能量是标量,如果要做多任务(比如同时预测能量和偶极矩),把输出维度改成任务数即可。
3.3 训练循环与损失函数选择
能量预测是回归任务,损失函数用 MSE 或 Huber。Huber 对异常值更鲁棒,如果数据集里有个别分子能量算错了,MSE 会被带偏。优化器用 AdamW,学习率 1e-3 起步,配合余弦退火。
from torch.optim import AdamW from torch.optim.lr_scheduler import CosineAnnealingLR device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') model = GINEnergyModel(node_dim=dataset[0].x.shape[1], edge_dim=dataset[0].edge_attr.shape[1]).to(device) optimizer = AdamW(model.parameters(), lr=1e-3, weight_decay=1e-5) scheduler = CosineAnnealingLR(optimizer, T_max=100) criterion = nn.HuberLoss(delta=1.0) for epoch in range(100): model.train() total_loss = 0 for batch in loader: batch = batch.to(device) optimizer.zero_grad() pred = model(batch) loss = criterion(pred, batch.y) loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=5.0) optimizer.step() total_loss += loss.item() * batch.num_graphs scheduler.step() print(f"Epoch {epoch}, loss {total_loss / len(dataset):.4f}")逻辑说明:clip_grad_norm_是分子图训练的后悔药,GNN 反向传播时梯度容易在深层爆炸,裁剪到 5.0 能救回不少训练。HuberLoss的delta控制异常值阈值,能量单位是 eV 时 1.0 比较合适,如果是 kcal/mol 得调到 5 到 10。
参数说明:weight_decay设 1e-5 到 1e-4,太大模型欠拟合,太小正则不够。T_max等于总 epoch 数,余弦退火让学习率从 1e-3 平滑降到接近 0。如果验证集 loss 连续 20 个 epoch 不降,可以提前停。
4. 避坑与排查:分子能量预测里最容易翻车的五件事
4.1 加了氢但标签没对齐,能量系统性偏移
现象:训练 loss 能降,但验证集 MAE 始终在 0.5 eV 以上,预测值整体偏高或偏低。原因:数据集标签是在含氢构型下算的,但建图时没加氢,或者加了氢但没重新算坐标。解决:确认数据集文档里能量对应的分子构型,QM9 和 MD17 都是含氢的,建图时必须AddHs。如果标签来自不含氢的简化模型,那就不能加氢。
4.2 边特征维度对不上,模型静默忽略边信息
现象:模型能跑通,但效果和只用节点特征差不多,边特征像没起作用。原因:GINConv默认不接收edge_attr,你传了它也不用。解决:要么换支持边特征的层(比如GINEConv),要么把边特征拼到节点特征里。我一般用GINEConv,它把边特征加到消息传递里,对键类型敏感的任务提升明显。
4.3 学习率太大导致 loss 变 NaN
现象:训练几个 batch 后 loss 变成 nan,梯度爆炸。原因:GNN 层数深、特征尺度大时,初始学习率 1e-3 可能太大。解决:先降到 1e-4 跑几个 epoch 看 loss 是否稳定,稳定后再逐步升。同时开梯度裁剪,max_norm设 5.0 或 10.0。另外检查输入特征有没有未归一化的连续值,比如原子坐标直接塞进去,尺度能到几十,必须标准化。
4.4 数据集划分随机切分导致数据泄漏
现象:测试集指标好得离谱,上线后一塌糊涂。原因:随机切分时,同一分子的不同构象可能同时出现在训练集和测试集,模型记住了分子而不是学了化学。解决:按分子骨架切分,用 Scaffold Split 而不是随机 Split。RDKit 有现成的MurckoScaffold可以拿骨架,按骨架分组后再切。如果做的是构象能量预测,必须按分子 ID 切分,同一分子的所有构象只能出现在一个集合里。
4.5 过平滑让深层 GNN 反而更差
现象:把层数从 4 加到 8,训练 loss 降了但验证 loss 升了,节点表示余弦相似度接近 1。原因:GNN 消息传递层数太多,所有节点表示趋同,丢失局部差异。解决:层数控制在 3 到 5,或者加残差连接、Jumping Knowledge。我一般用 4 层加残差,再深就得换策略。另外可以监控节点表示的方差,如果方差小于 1e-3,基本就是过平滑了。
5. 进阶技巧:用注意力池化和集成提升预测精度
5.1 把 mean pool 换成 attention pool
全局平均池化对所有节点一视同仁,但分子里有些原子对能量的贡献更大,比如官能团上的杂原子。注意力池化让模型自己学每个节点的权重。
from torch_geometric.nn import GlobalAttention from torch.nn import Linear class AttentiveGIN(nn.Module): def __init__(self, node_dim, hidden=128, num_layers=4): super().__init__() self.convs = nn.ModuleList() self.bns = nn.ModuleList() for i in range(num_layers): in_dim = node_dim if i == 0 else hidden mlp = nn.Sequential(nn.Linear(in_dim, hidden), nn.ReLU(), nn.Linear(hidden, hidden)) self.convs.append(GINConv(mlp, train_eps=True)) self.bns.append(nn.BatchNorm1d(hidden)) self.gate_nn = nn.Sequential(nn.Linear(hidden, hidden // 2), nn.ReLU(), nn.Linear(hidden // 2, 1)) self.pool = GlobalAttention(gate_nn=self.gate_nn) self.fc = nn.Sequential(nn.Linear(hidden, hidden // 2), nn.ReLU(), nn.Linear(hidden // 2, 1)) def forward(self, data): x, edge_index, batch = data.x, data.edge_index, data.batch for conv, bn in zip(self.convs, self.bns): x = F.relu(bn(conv(x, edge_index))) x = self.pool(x, batch) return self.fc(x).squeeze(-1)逻辑说明:GlobalAttention用一个小的 gate 网络给每个节点打分,再 softmax 归一化加权求和。gate 网络的输出维度是 1,表示每个节点的重要性。相比 mean pool,attention pool 在含杂原子多的分子上通常能降 5% 到 10% 的 MAE。
参数说明:gate_nn的隐藏层设hidden // 2就够,太大容易过拟合。如果数据集很小(小于 5000),attention pool 的优势可能被过拟合抵消,这时候还是用 mean pool 稳。
5.2 多模型集成与不确定性估计
单模型预测能量总会有波动,工业场景里往往需要给出置信区间。做法是训练 5 到 10 个不同初始化的 GIN 模型,预测时取均值和标准差。均值作为最终预测,标准差作为不确定性。如果某个分子的预测标准差特别大,说明模型对它没把握,可以挑出来用 DFT 复算。
def ensemble_predict(models, data, device): preds = [] for m in models: m.eval() with torch.no_grad(): preds.append(m(data.to(device)).cpu().numpy()) preds = np.stack(preds) # shape: (num_models, num_samples) return preds.mean(axis=0), preds.std(axis=0)逻辑说明:集成时每个模型用不同的随机种子初始化,数据划分保持一致。标准差反映的是模型间认知不确定性,不是数据噪声。如果要做更严格的不确定性估计,可以上 MC Dropout 或 Deep Ensemble,但计算成本更高。
参数说明:模型数量 5 个起步,10 个一般够用。再多边际收益递减。集成推理时间线性增长,如果线上延迟敏感,可以蒸馏成单模型。
5.3 验证方法:别只看 MAE
能量预测的评估不能只看整体 MAE。我一般会分三块看:一是按分子大小分组看 MAE,小分子和大分子的误差可能差一个量级;二是按官能团分组,看模型在含氮、含硫、含卤素子集上的表现;三是画预测值 vs 真实值的散点图,看有没有系统性偏差。如果散点图在高能量区域发散,说明模型外推能力差,训练集里高能量样本太少,得补数据。
另外,化学领域有个习惯是看化学精度(chemical accuracy),也就是 MAE 是否小于 1 kcal/mol(约 0.043 eV)。如果你的模型 MAE 在 0.05 eV 左右,已经接近这个门槛,可以考虑替代一部分低精度 DFT 做预筛选。但要注意,这个精度只在训练集覆盖的化学空间内成立,超出分布外的分子不能信。
我自己踩过的最大坑是早期只看整体 MAE,模型在含氟分子上误差是平均值的 3 倍,但被大量碳氢分子拉平了。后来按子结构分组评估才发现问题,补了含氟数据才修好。做分子能量预测,数据覆盖度比模型结构重要得多,先把化学空间铺够,再调模型。希望帮到你。
本文还有配套的精品资源,点击获取