034

分子图表示:原子、键与图神经网络的连接方式

分子图是 GNN 的输入形式。这篇讲清图的构建、原子与键特征的设计,以及常见的实现错误。

分子本质上就是图:原子是节点,化学键是边。这个对应关系让图神经网络成为分子建模的自然选择——但「图怎么建、特征怎么设计」对结果的影响,往往大于选哪个 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。

延伸资源