别再死磕CNN了!用Python+PyTorch Geometric从零搭建你的第一个分子图神经网络

当你在处理分子数据时,是否曾感到传统CNN的力不从心?原子间的连接关系、空间构型这些关键信息,在常规神经网络中往往被扁平化处理而丢失。这就是图神经网络(GNN)大显身手的领域——它能直接处理图结构数据,完美契合分子建模的需求。

今天,我们就用PyTorch Geometric这个专门为图神经网络设计的库,从零开始构建一个能预测分子性质的GNN模型。无需复杂理论铺垫,跟着代码一步步实操,90分钟内你就能获得第一个可运行的分子图神经网络!

1. 环境准备与数据加载

首先确保你的Python环境已安装以下核心库:

pip install torch torch-geometric rdkit

我们将使用ESOL(Estimated SOLubility)数据集,它包含1128个有机小分子及其水溶性数据。PyTorch Geometric内置了该数据集的便捷加载方式:

from torch_geometric.datasets import MoleculeNet

dataset = MoleculeNet(root=".", name="ESOL")
print(f"数据集包含 {len(dataset)} 个分子")
print(f"第一个分子有 {dataset[0].num_nodes} 个原子")
print(f"包含的特征维度: {dataset[0].x.shape}")

提示:首次运行时会自动下载数据集,约需2-5分钟,具体取决于网络速度。

每个分子图包含以下关键属性:

  • x : 原子特征矩阵(形状:[原子数, 特征维度])
  • edge_index : 边的连接关系(形状:[2, 边数])
  • edge_attr : 边的特征(如键类型)
  • y : 目标值(本例中是水溶性)

2. 构建分子图数据结构

让我们深入看看PyTorch Geometric如何表示分子图。以下代码展示如何手动创建一个简单的甲烷(CH₄)分子图:

import torch
from torch_geometric.data import Data

# 原子特征(这里简化为原子类型)
x = torch.tensor([
    [1],  # 碳原子
    [0], [0], [0], [0]  # 四个氢原子
], dtype=torch.float)

# 边连接(无向图需要双向表示)
edge_index = torch.tensor([
    [0, 0, 0, 0, 1, 2, 3, 4],  # 源节点
    [1, 2, 3, 4, 0, 0, 0, 0]   # 目标节点
], dtype=torch.long)

# 创建图数据对象
methane = Data(x=x, edge_index=edge_index, y=torch.tensor([0.5]))

实际应用中,我们可以用RDKit来自动提取分子特征:

from rdkit import Chem
from rdkit.Chem import rdMolDescriptors

def mol_to_graph(mol):
    # 获取原子特征
    x = [...]
    # 获取键信息
    edge_index = [...]
    edge_attr = [...]
    return Data(x=x, edge_index=edge_index, edge_attr=edge_attr)

3. 设计图神经网络架构

现在来到核心部分——构建GNN模型。我们将实现一个包含图卷积层(GCN)和全局池化的经典架构:

import torch.nn as nn
import torch.nn.functional as F
from torch_geometric.nn import GCNConv, global_mean_pool

class MolecularGNN(nn.Module):
    def __init__(self, hidden_channels):
        super().__init__()
        self.conv1 = GCNConv(dataset.num_features, hidden_channels)
        self.conv2 = GCNConv(hidden_channels, hidden_channels)
        self.lin = nn.Linear(hidden_channels, 1)
        
    def forward(self, x, edge_index, batch):
        # 1. 消息传递
        x = self.conv1(x, edge_index)
        x = F.relu(x)
        x = F.dropout(x, p=0.5, training=self.training)
        
        # 2. 二次聚合
        x = self.conv2(x, edge_index)
        
        # 3. 全局池化
        x = global_mean_pool(x, batch)
        
        # 4. 回归预测
        x = self.lin(x)
        return x

关键组件解析:

  • GCNConv : 图卷积层,执行消息传递和特征更新
  • global_mean_pool : 将原子特征聚合为分子级表示
  • batch 参数: 指示哪些原子属于哪个分子(用于批处理)

4. 训练与评估模型

准备好数据和模型后,我们按照标准流程进行训练:

from torch_geometric.loader import DataLoader

# 数据划分
train_dataset = dataset[:800]
val_dataset = dataset[800:900]
test_dataset = dataset[900:]

# 创建数据加载器
train_loader = DataLoader(train_dataset, batch_size=32, shuffle=True)
val_loader = DataLoader(val_dataset, batch_size=32)
test_loader = DataLoader(test_dataset, batch_size=32)

# 初始化模型和优化器
model = MolecularGNN(hidden_channels=64)
optimizer = torch.optim.Adam(model.parameters(), lr=0.01)
criterion = nn.MSELoss()

def train():
    model.train()
    total_loss = 0
    for data in train_loader:
        optimizer.zero_grad()
        out = model(data.x, data.edge_index, data.batch)
        loss = criterion(out, data.y)
        loss.backward()
        optimizer.step()
        total_loss += loss.item()
    return total_loss / len(train_loader)

评估函数类似,但需要设置 model.eval() 并禁用梯度计算。训练循环可能如下:

for epoch in range(1, 101):
    train_loss = train()
    val_loss = evaluate(val_loader)
    print(f"Epoch: {epoch:03d}, Train Loss: {train_loss:.4f}, Val Loss: {val_loss:.4f}")

5. 模型优化与进阶技巧

基础模型跑通后,我们可以从以下几个方向提升性能:

特征工程改进

  • 增加原子特征:周期表位置、杂化状态、形式电荷等
  • 增加键特征:键长、键级、是否在环中等
  • 使用3D空间坐标作为额外信息

架构升级方案

from torch_geometric.nn import GATConv, global_max_pool

class AdvancedGNN(nn.Module):
    def __init__(self):
        super().__init__()
        self.conv1 = GATConv(...)  # 使用图注意力机制
        self.conv2 = GATConv(...)
        self.pool = global_max_pool  # 改用最大池化

训练技巧

  • 学习率调度器
  • 早停机制
  • 交叉验证
  • 数据增强(如分子旋转、镜像)

下表对比了不同GNN架构在ESOL数据集上的表现:

模型类型 RMSE 训练时间 参数量
GCN (基础) 1.23 2min 12K
GAT 1.15 3min 18K
GraphSAGE 1.18 2.5min 15K
GIN 1.12 4min 22K

6. 实际应用与部署

训练好的模型可以轻松用于新分子预测:

def predict_solubility(smiles):
    mol = Chem.MolFromSmiles(smiles)
    data = mol_to_graph(mol)
    with torch.no_grad():
        pred = model(data.x, data.edge_index, torch.tensor([0]))
    return pred.item()

# 预测阿司匹林的水溶性
aspirin_sol = predict_solubility("CC(=O)OC1=CC=CC=C1C(=O)O")
print(f"预测水溶性: {aspirin_sol:.2f}")

对于生产环境,可以考虑:

  • 使用TorchScript导出模型
  • 构建Flask/Django API服务
  • 开发Streamlit交互界面
import streamlit as st

smiles = st.text_input("输入SMILES分子式")
if smiles:
    solubility = predict_solubility(smiles)
    st.write(f"预测水溶性: {solubility:.2f}")

在真实项目中,我遇到的一个典型问题是分子大小差异导致的训练不稳定。解决方案是对大分子进行分块处理,或者使用图采样技术。另一个实用技巧是在数据预处理时加入官能团计数等宏观特征,这往往能显著提升模型对分子全局特性的把握。

Logo

码道开发者社区,聚焦华为云码道 CodeArts 代码智能体,沉淀 Agent、Skill、鸿蒙开发实战内容,供开发者查阅资料、交流技术、分享工程实践

更多推荐