图神经网络(GNN)是分子建模最自然的架构:分子本来就是图。但「表示自然」不等于「效果更好」——在很多实际任务上,GNN 打不过认真调优的 ECFP 指纹基线。理解为什么,比追新架构更有价值。
消息传递的基本机制
# 分子 → 图
# 节点 = 原子(特征:元素、电荷、杂化、芳香性、成环...)
# 边 = 化学键(特征:键级、是否共轭、是否成环...)
#
# 消息传递(每一层):
#
# 1) 消息:每个原子从邻居收集信息
# m_v = Σ_{w∈N(v)} M(h_v, h_w, e_vw)
#
# 2) 更新:用消息更新自身状态
# h_v' = U(h_v, m_v)
#
# 经过 K 层后,每个原子的表示包含了
# K 跳邻域内的信息
#
# 3) 读出:把所有原子表示聚合成分子表示
# h_G = READOUT({h_v})
# 常用:求和、平均、最大值、注意力加权
#
# 【关键理解】:
# K 层 GNN 的感受野是 K 跳
# → 与 ECFP 半径 K 的信息范围【几乎等价】
# → 这解释了为什么两者性能常常接近
#
# 主要变体:
# GCN 谱域卷积的简化
# GAT 用注意力加权邻居
# GIN 理论表达能力最强(等价于 1-WL 测试)
# MPNN 通用框架
# D-MPNN 在有向边上传递(见 131)
# AttentiveFP 分子级注意力读出
为什么常常输给指纹基线
# 原因一:【信息范围相当】
# ECFP(r=2) 编码了每个原子 2 跳内的环境
# 3 层 GNN 也是 3 跳
# → 【信息量本质上差不多】
# ECFP 是「枚举 + 哈希」,GNN 是「学习聚合」
# 在数据少时,枚举比学习更稳健
#
# 原因二:【数据量不足】
# GNN 参数多,需要更多数据才能学好
# 而分子数据集常常只有几百到几千
# → 【过拟合风险高】
#
# 原因三:【全局性质表达不高效】
# 分子量、logP、TPSA 这类全局性质,
# GNN 需要通过多层聚合间接学习
# 而它们可以直接计算
# → 【这就是为什么「GNN + RDKit 描述符」
# 比纯 GNN 好得多】(见 131)
#
# 原因四:【过平滑(over-smoothing)】
# 层数增加时,所有节点的表示趋于相同
# → 【限制了 GNN 的深度】(通常 3~5 层)
# → 因此难以捕捉长距离的分子内相互作用
#
# 原因五:【表达能力的理论上限】
# 标准的消息传递 GNN 的判别能力
# 不超过 1-WL 图同构测试
# → 【某些不同的分子图,GNN 无法区分】
# → 典型例子:某些环系的组合
#
# 原因六:【工程成本】
# GNN 需要 GPU、训练时间长、超参数多
# 而 ECFP + LightGBM 几分钟就能训完
# → 在提升不明显时,成本不划算
GNN 真正有优势的场景
| 场景 | 为什么 GNN 更好 |
|---|---|
| 大数据集(> 1 万) | 有足够数据学习任务特定的表示 |
| 多任务学习 | 共享表示,数据少的任务受益 |
| 需要原子级输出 | 如预测每个原子的性质、反应位点 |
| 需要三维信息 | 等变 GNN 能处理坐标(见 160《等变神经网络》) |
| 端到端优化 | 表示与任务联合学习 |
| 生成任务 | 图生成需要图的表示 |
| 迁移学习 | 预训练的表示可迁移 |
「需要原子级输出」是 GNN 最不可替代的优势:指纹给出的是分子级向量,无法回答「哪个原子会被代谢」「哪个位点会发生反应」这类问题。
实现
import torch
import torch.nn as nn
from torch_geometric.nn import GINEConv, global_add_pool
from torch_geometric.data import Data
from rdkit import Chem
def mol_to_graph(smiles, y=None):
mol = Chem.MolFromSmiles(smiles)
if mol is None:
return None
# 原子特征
x = []
for a in mol.GetAtoms():
x.append([
a.GetAtomicNum(), a.GetDegree(), a.GetFormalCharge(),
int(a.GetHybridization()), int(a.GetIsAromatic()),
a.GetTotalNumHs(), int(a.IsInRing()),
])
# 边(无向图需要双向)
ei, ea = [], []
for b in mol.GetBonds():
i, j = b.GetBeginAtomIdx(), b.GetEndAtomIdx()
feat = [int(b.GetBondTypeAsDouble() * 2), int(b.GetIsConjugated()),
int(b.IsInRing())]
ei += [[i, j], [j, i]]
ea += [feat, feat]
return Data(x=torch.tensor(x, dtype=torch.float),
edge_index=torch.tensor(ei, dtype=torch.long).t().contiguous(),
edge_attr=torch.tensor(ea, dtype=torch.float),
y=torch.tensor([y], dtype=torch.float) if y is not None else None)
class MolGNN(nn.Module):
def __init__(self, node_dim=7, edge_dim=3, hidden=128, n_layers=3,
n_global=12):
super().__init__()
self.node_enc = nn.Linear(node_dim, hidden)
self.edge_enc = nn.Linear(edge_dim, hidden)
self.convs = nn.ModuleList()
for _ in range(n_layers):
mlp = nn.Sequential(nn.Linear(hidden, hidden), nn.ReLU(),
nn.Linear(hidden, hidden))
self.convs.append(GINEConv(mlp))
# 【关键】:拼接全局描述符
self.head = nn.Sequential(
nn.Linear(hidden + n_global, hidden), nn.ReLU(),
nn.Dropout(0.2), nn.Linear(hidden, 1))
def forward(self, data, global_feats):
x = self.node_enc(data.x)
e = self.edge_enc(data.edge_attr)
for conv in self.convs:
x = torch.relu(conv(x, data.edge_index, e))
g = global_add_pool(x, data.batch)
return self.head(torch.cat([g, global_feats], dim=1))
# 【务必先跑的基线】:
# ECFP(2048) + RDKit 描述符 + LightGBM
# 如果 GNN 打不过它,就用基线
提高 GNN 效果的实用技巧
- 拼接全局描述符:这是最有效也最简单的改进(见 131《Chemprop / D-MPNN 论文精读》);
- 边特征不要省:键的类型、共轭、成环信息对分子性质很重要,用 GINEConv 这类支持边特征的层;
- 读出方式:求和保留分子大小信息,平均则不保留——取决于目标性质是否与大小相关;
- 层数控制在 3~5:更深会过平滑;如需更大感受野,可加虚拟节点(连接所有原子的全局节点);
- 集成:多个种子训练后平均,能明显降低方差;
- 正则化:dropout、边 dropout、早停——小数据时尤其重要;
- 不要迷信新架构:超参数调优的收益常常大于换架构。
关键要点
- K 层 GNN 的感受野与 ECFP 半径 K 几乎等价——这解释了两者性能为何常常接近;
- 拼接 RDKit 全局描述符是最有效的简单改进——GNN 表达全局性质效率低;
- 过平滑限制了 GNN 的深度(通常 3~5 层),难以捕捉长距离相互作用;
- GNN 最不可替代的优势是原子级输出与三维等变建模。
延伸资源
- Chemprop:131《Chemprop / D-MPNN 论文精读》;等变网络:160《等变神经网络》;Transformer 在分子中:161《Transformer 在分子中的应用》;
- 分子指纹:035《分子指纹是什么》;ECFP:037《ECFP 指纹详解》;PyTorch Geometric:176《PyTorch Geometric》。