159

图神经网络 GNN:分子图模型强在哪里,为什么仍会输给指纹基线

GNN 在分子建模中强在哪里,又为什么常常输给指纹基线。这篇讲清消息传递的机制、固有局限与实践判断。

图神经网络(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 最不可替代的优势是原子级输出三维等变建模

延伸资源