别再死磕CNN了!用Python和PyTorch Geometric(PyG)从零搭建你的第一个分子图神经网络
·
用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()}")
传统神经网络处理这类数据时面临三大挑战:
- 非欧几里得结构 :分子图中的节点邻居数量不固定 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)
这个基础架构包含三个关键组件:
- 图卷积层 :执行消息传递和节点更新
- 全局池化 :将节点特征聚合为图级表示
- 回归头 :输出预测值
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}')
如果模型无法在微型数据集上过拟合,说明架构存在根本性问题。
更多推荐


所有评论(0)