用Python和PyTorch Geometric从零构建分子图神经网络实战指南

在药物发现和材料设计领域,传统深度学习方法常遭遇瓶颈——分子本质上是由原子和化学键构成的图结构数据,而卷积神经网络(CNN)等架构是为网格化数据设计的。这就像试图用渔网捕捉空气,工具与目标存在根本性不匹配。本文将带您突破这一限制,使用PyTorch Geometric(PyG)构建能直接处理分子图的神经网络。

1. 为什么分子数据需要特殊处理?

每个化学家都知道,苯环的六边形结构决定了其稳定性,双键的位置影响着反应活性。这些关键信息在SMILES字符串或分子指纹中是被压缩甚至丢失的。当我们把分子表示为图结构时:

  • 节点 代表原子,可包含原子类型、电荷等特征
  • 代表化学键,可存储键类型、长度等属性
  • 全局属性 如分子量、手性等作为补充
from rdkit import Chem
from rdkit.Chem import AllChem

mol = Chem.MolFromSmiles('CCO')  # 乙醇分子
print(f"原子数: {mol.GetNumAtoms()}, 键数: {mol.GetNumBonds()}")

传统神经网络处理这类数据时面临三大挑战:

  1. 非欧几里得结构 :分子图中的节点邻居数量不固定 2.** 边信息重要性**:单键/双键传递的电子信息完全不同
  2. 旋转不变性 :分子旋转不应影响预测结果

2. 构建分子图数据集

使用RDKit将SMILES转换为图数据是关键第一步。以下是创建可训练数据集的完整流程:

import torch
from torch_geometric.data import Data
from rdkit.Chem import Descriptors

def smiles_to_graph(smiles):
    mol = Chem.MolFromSmiles(smiles)
    if not mol:
        return None
    
    # 原子特征(使用one-hot编码)
    atom_features = []
    for atom in mol.GetAtoms():
        feature = [
            atom.GetAtomicNum(),
            atom.GetDegree(),
            atom.GetFormalCharge()
        ]
        atom_features.append(feature)
    
    # 边索引和边属性
    edge_index = []
    edge_attr = []
    for bond in mol.GetBonds():
        i = bond.GetBeginAtomIdx()
        j = bond.GetEndAtomIdx()
        edge_type = bond.GetBondTypeAsDouble()
        
        edge_index.append((i, j))
        edge_index.append((j, i))  # 无向图
        edge_attr.append([edge_type])
        edge_attr.append([edge_type])
    
    # 全局属性(示例使用分子量)
    mol_weight = Descriptors.MolWt(mol)
    
    return Data(
        x=torch.tensor(atom_features, 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([[mol_weight]], dtype=torch.float)
    )

注意:实际应用中应将分子量替换为您要预测的真实目标值(如溶解度、活性等)

3. 设计图神经网络架构

PyG提供了多种现成的图卷积层,我们从最简单的GCN开始构建:

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

class MolecularGNN(nn.Module):
    def __init__(self, hidden_channels):
        super().__init__()
        self.conv1 = GCNConv(3, hidden_channels)  # 输入特征维度3
        self.conv2 = GCNConv(hidden_channels, hidden_channels)
        self.lin = nn.Linear(hidden_channels, 1)
    
    def forward(self, data):
        x, edge_index = data.x, data.edge_index
        
        # 消息传递
        x = self.conv1(x, edge_index)
        x = F.relu(x)
        x = F.dropout(x, training=self.training)
        x = self.conv2(x, edge_index)
        
        # 全局池化(取所有节点均值)
        x = torch.mean(x, dim=0)
        
        # 回归预测
        return self.lin(x)

这个基础架构包含三个关键组件:

  1. 图卷积层 :执行消息传递和节点更新
  2. 全局池化 :将节点特征聚合为图级表示
  3. 回归头 :输出预测值

4. 训练技巧与性能优化

实际训练时,以下几个技巧能显著提升模型表现:

数据预处理最佳实践

  • 对原子特征进行标准化处理
  • 平衡数据集中的分子大小分布
  • 使用数据增强(如随机旋转分子)

模型训练配置

from torch_geometric.loader import DataLoader

# 示例训练循环
def train(model, loader, optimizer):
    model.train()
    total_loss = 0
    
    for data in loader:
        optimizer.zero_grad()
        out = model(data)
        loss = F.mse_loss(out, data.y)
        loss.backward()
        optimizer.step()
        total_loss += loss.item()
    
    return total_loss / len(loader)

# 初始化
dataset = [smiles_to_graph(s) for s in smiles_list]
loader = DataLoader(dataset, batch_size=32, shuffle=True)
model = MolecularGNN(hidden_channels=64)
optimizer = torch.optim.Adam(model.parameters(), lr=0.01)

# 训练循环
for epoch in range(100):
    loss = train(model, loader, optimizer)
    print(f'Epoch: {epoch:03d}, Loss: {loss:.4f}')

高级架构改进

  • 使用EdgeConv考虑边特征
  • 添加注意力机制(GAT)
  • 采用残差连接防止梯度消失

5. 实战案例:溶解度预测

让我们用ESOL(水溶性数据集)演示完整流程。首先准备数据:

import pandas as pd
from sklearn.model_selection import train_test_split

# 加载ESOL数据集
data = pd.read_csv('esol.csv')
smiles = data['smiles'].tolist()
targets = data['measured log solubility'].values

# 转换为图数据集
graphs = []
for s, y in zip(smiles, targets):
    g = smiles_to_graph(s)
    if g:
        g.y = torch.tensor([[y]], dtype=torch.float)
        graphs.append(g)

# 划分训练测试集
train_data, test_data = train_test_split(graphs, test_size=0.2)

然后构建更强大的网络:

from torch_geometric.nn import global_mean_pool, GATConv

class AdvancedGNN(nn.Module):
    def __init__(self):
        super().__init__()
        self.gat1 = GATConv(3, 64, heads=4)
        self.gat2 = GATConv(64*4, 64)
        self.lin = nn.Sequential(
            nn.Linear(64, 32),
            nn.ReLU(),
            nn.Linear(32, 1)
        )
    
    def forward(self, data):
        x, edge_index, batch = data.x, data.edge_index, data.batch
        
        x = self.gat1(x, edge_index)
        x = F.elu(x)
        x = self.gat2(x, edge_index)
        x = global_mean_pool(x, batch)
        return self.lin(x)

训练后评估模型:

def test(model, loader):
    model.eval()
    total_error = 0
    
    for data in loader:
        with torch.no_grad():
            pred = model(data)
            error = (pred - data.y).abs().mean()
            total_error += error.item()
    
    return total_error / len(loader)

test_loader = DataLoader(test_data, batch_size=32)
print(f'Test MAE: {test(model, test_loader):.4f}')

6. 调试与常见问题解决

当模型表现不佳时,可以检查以下方面:

数据层面

  • 使用 torch_geometric.utils.degree 检查节点度数分布
  • 验证边属性是否被正确利用
  • 检查目标值的归一化是否合理

模型层面

  • 可视化消息传递路径
  • 监控各层梯度变化
  • 尝试不同的聚合函数(sum/mean/max)

训练过程

  • 调整学习率与batch size
  • 添加Layer Normalization
  • 尝试不同的损失函数组合

一个实用的调试技巧是创建小型验证集:

# 创建仅含5个分子的微型数据集
debug_data = graphs[:5]
debug_loader = DataLoader(debug_data, batch_size=2)

# 应能快速过拟合
for epoch in range(50):
    loss = train(model, debug_loader, optimizer)
    if epoch % 10 == 0:
        print(f'Debug Loss: {loss:.4f}')

如果模型无法在微型数据集上过拟合,说明架构存在根本性问题。

Logo

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

更多推荐