简介:这份资源是面向医药信息学、生物信息学方向的学习者与研究者整理的深度学习实战项目包,聚焦药物相互作用预测这一交叉课题,适合具备Python基础、希望了解图神经网络与分子表示建模的中高级读者参考。压缩包共19个文件,以13个py脚本和3个ipynb笔记本为主,辅以2张模型架构与图结构示意图及1个依赖说明文件,整体约621KB,结构紧凑便于快速上手。内容围绕decagon模型展开,涵盖药物分子SMILES编码与指纹特征预处理、图神经网络对分子拓扑结构的建模、二元分类标签构建、模型训练与交叉验证、AUC-ROC与AUPRC等指标评估,以及SHAP等可解释性分析思路,并配有探索性分析与测试数据集笔记本。目前已有192人学习下载,可作为复现药物相互作用预测流程、理解GNN在药理数据上应用的参考范例。
1. 药物相互作用预测为什么值得用深度学习重做一遍
两个药一起吃,会不会出事?这个问题在临床上每天都要回答无数次。传统做法靠药代动力学实验、靠文献检索、靠药师经验,覆盖的药物对数量极其有限。已知上市药物两两组合是万级到十万级的天文数字,靠人力穷举根本不现实。基于深度学习的药物相互作用预测,本质上是把「这对药能不能一起吃」变成一个二分类或多分类问题,用药物分子结构、靶点、酶、通路等特征训练模型,让模型去推断那些还没被实验验证过的组合。它解决的是「覆盖率」和「成本」两个死结:实验做不完的,模型先筛一遍,把高风险组合挑出来优先验证。适合谁?做药物警戒、临床决策支持、药企早期筛选、以及想拿深度学习做真实生物医药项目的工程师。这个方向数据公开、任务定义清晰、模型可解释性有抓手,是少有的「学术能发论文、工业能落地」的交叉领域。
2. 药物相互作用预测的数据从哪来、特征怎么建
2.1 三类主流数据源和它们的取舍
做这个任务,第一步不是搭模型,是搞清楚你手里有什么数据。常见的数据源分三类,各有各的坑。
第一类是药物分子结构数据,最典型的是 SMILES 字符串和分子指纹。SMILES 是药物的文本表示,比如阿司匹林是CC(=O)Oc1ccccc1C(=O)O。它的好处是几乎任何药物都能拿到,坏处是字符串本身没有显式编码化学性质,需要模型自己去学。分子指纹(如 Morgan fingerprint、MACCS keys)是把结构压成固定长度的 0/1 向量,计算快、可复现,但会丢失部分拓扑信息。
第二类是生物实体关联数据,包括靶点蛋白、代谢酶(尤其是 CYP450 家族)、转运体、通路。药物相互作用很大一部分机制是「A 药抑制了代谢 B 药的酶」,所以酶和靶点的共享关系是强特征。这类数据通常从 DrugBank、PubChem、ChEMBL 这类公开库整理,但要注意版本差异——不同年份的库,同一个药物的靶点标注可能不一样。
第三类是已知相互作用标签,也就是监督学习的 y。DrugBank 的 DDI 表是最常用的,但它有个致命问题:正样本多、负样本少且不可靠。没有标注相互作用的药物对,不代表真的没有相互作用,可能只是没人研究过。这是整个任务最大的数据陷阱,后面避坑章节会细说。
| 数据源 | 典型来源 | 特征形式 | 主要问题 |
|---|---|---|---|
| 分子结构 | PubChem、DrugBank | SMILES / 指纹 | 需自行编码,指纹丢信息 |
| 生物实体 | DrugBank、ChEMBL | 靶点/酶多热向量 | 版本不一致,缺失多 |
| 相互作用标签 | DrugBank DDI | 二分类/多分类 | 负样本不可靠 |
2.2 把 SMILES 变成模型能吃的张量
分子结构编码是绕不开的一步。我一般会同时准备两套特征:一套指纹向量给传统模型和快速基线,一套字符级序列给深度模型。下面这段代码把 SMILES 转成 Morgan 指纹,并做基本的合法性校验。
from rdkit import Chem from rdkit.Chem import AllChem import numpy as np def smiles_to_fingerprint(smiles, radius=2, n_bits=2048): """把 SMILES 转成 Morgan 指纹向量。 radius: 考虑周围几层原子,2 是常用值 n_bits: 指纹长度,2048 在药物任务里够用 """ mol = Chem.MolFromSmiles(smiles) if mol is None: return None # 非法 SMILES,必须过滤,否则后面全崩 fp = AllChem.GetMorganFingerprintAsBitVect(mol, radius, nBits=n_bits) return np.array(fp, dtype=np.float32) # 批量处理,记录失败的样本 def build_fingerprint_matrix(smiles_list): feats, valid_idx = [], [] for i, smi in enumerate(smiles_list): fp = smiles_to_fingerprint(smi) if fp is not None: feats.append(fp) valid_idx.append(i) return np.stack(feats), valid_idx逻辑说明:MolFromSmiles返回 None 说明这个 SMILES 写错了或者 RDKit 解析不了,必须丢弃并记录索引,否则后续矩阵对齐会错位。radius=2对应 ECFP4,是药物化学里被验证过最稳的默认值;n_bits=2048是精度和内存的折中,药物分子原子数一般不超过 100,2048 位足够稀疏表达。参数怎么改?如果你的药物分子特别大(比如多肽类),radius 可以提到 3,n_bits 提到 4096,但要注意特征维度上升后小数据集容易过拟合。
2.3 负样本构造:这一步决定模型上限
正样本从 DrugBank 拿,负样本怎么办?直接随机采样是最常见的做法,但会引入大量「假阴性」——你随机抽的一对药,可能其实有相互作用只是没被记录。我的经验是分层构造:先从「已知无相互作用」的明确记录里取一部分,再用「靶点/酶完全不重叠」的药物对补充,最后才用随机采样兜底。比例上正负 1:1 到 1:3 之间比较稳,负样本太多会让模型偏向多数类。
import random def build_negative_samples(drug_ids, positive_pairs, ratio=2): """基于随机采样构造负样本,ratio 为负正比。 注意:这只是基线做法,生产环境要叠加规则过滤。 """ pos_set = set(positive_pairs) negatives = [] target = len(positive_pairs) * ratio while len(negatives) < target: a, b = random.sample(drug_ids, 2) if (a, b) not in pos_set and (b, a) not in pos_set: negatives.append((a, b)) return negatives这段代码能跑,但我要提醒:它没有做任何机制层面的过滤。真正上线前,至少要把「共享靶点」「共享代谢酶」的药物对从负样本里剔除,否则模型学到的只是「共享靶点=有相互作用」这个捷径,换个数据集就翻车。
3. 模型选型:从指纹 MLP 到图神经网络怎么选
3.1 基线模型先跑通,别一上来就上 GNN
很多人的第一反应是直接上图神经网络,觉得分子天然是图结构。但我的血泪经验是:先用指纹 + MLP 跑一个基线,把数据管道、评估指标、划分方式全部验证一遍,再考虑复杂模型。原因很简单,如果基线只有 0.6 的 AUC,你换成 GNN 大概率也就 0.65,问题出在数据不在模型。
基线模型的结构很朴素:两个药物的指纹各过一个编码器,拼接后接分类头。
import torch import torch.nn as nn class DDI_MLP(nn.Module): def __init__(self, fp_dim=2048, hidden=512, n_class=2): super().__init__() self.encoder = nn.Sequential( nn.Linear(fp_dim, hidden), nn.ReLU(), nn.Dropout(0.3), # 药物特征维度高,dropout 必加 nn.Linear(hidden, hidden // 2), nn.ReLU(), ) self.classifier = nn.Sequential( nn.Linear(hidden, 128), nn.ReLU(), nn.Linear(128, n_class), ) def forward(self, fp_a, fp_b): h_a = self.encoder(fp_a) h_b = self.encoder(fp_b) # 拼接 + 逐元素乘积,乘积能捕捉两药特征的交互 combined = torch.cat([h_a, h_b, h_a * h_b], dim=-1) return self.classifier(combined)逻辑说明:两个药物共享同一个 encoder,这是合理的,因为药物特征空间是同一套。h_a * h_b这个逐元素乘积是关键,它显式建模了两个药物特征的交互,比单纯拼接效果好。参数上hidden=512对 2048 维指纹是合适的压缩比,dropout 0.3 是防止高维特征过拟合的常用值。如果你的数据集小于 5000 对,dropout 可以提到 0.5。
3.2 图神经网络什么时候真正带来增益
GNN 的价值在于它直接从原子和键的图结构学表示,不需要人工设计指纹。常见做法是用 GCN 或 GIN 编码每个药物分子图,得到图级表示后再做交互。什么时候 GNN 会明显超过指纹基线?我的观察是:当你的数据量足够大(万级以上药物对),且任务依赖精细的局部结构时。如果数据只有几千对,GNN 参数多、容易过拟合,反而不如指纹稳。
from torch_geometric.nn import GINConv, global_add_pool class DrugGNN(nn.Module): def __init__(self, node_dim=78, hidden=128, n_layer=3): super().__init__() self.convs = nn.ModuleList() for _ in range(n_layer): mlp = nn.Sequential(nn.Linear(hidden, hidden), nn.ReLU(), nn.Linear(hidden, hidden)) self.convs.append(GINConv(mlp)) self.input_proj = nn.Linear(node_dim, hidden) def forward(self, x, edge_index, batch): h = self.input_proj(x) for conv in self.convs: h = conv(h, edge_index).relu() return global_add_pool(h, batch) # 图级表示逻辑说明:node_dim=78是原子特征的标准维度(原子类型、度、电荷、手性等 one-hot 拼接)。n_layer=3意味着每个原子能感知到 3 跳以内的邻居,对大多数药物分子够用,层数再深会出现过平滑。global_add_pool把原子表示聚合成分子表示,也可以用 mean 或 attention pool,add 对分子大小不敏感,是我更常用的选择。
3.3 交互建模:拼接、双线性还是注意力
两个药物的表示拿到之后,怎么建模它们的交互,直接决定模型上限。最简单的拼接在基线里够用,但如果你想再往上推,双线性池化和注意力是两条路。双线性用一个小矩阵 W 建模 h_a 和 h_b 的二次交互,参数量可控;注意力则让模型自己决定关注哪些特征维度。
| 交互方式 | 参数量 | 适用场景 | 注意点 |
|---|---|---|---|
| 拼接 | 低 | 基线、小数据 | 交互靠后续全连接隐式学 |
| 逐元素乘积 | 低 | 通用增强 | 与拼接叠加最稳 |
| 双线性 | 中 | 中等数据 | 需低秩分解防过拟合 |
| 注意力 | 高 | 大数据 | 要足够数据才训得动 |
我一般的组合是「拼接 + 逐元素乘积」作为默认,数据量过万再考虑加双线性。注意力机制在 DDI 任务上不是必须的,除非你要做可解释性分析,想看模型关注了哪些子结构。
4. 训练与评估:别被虚高的 AUC 骗了
4.1 数据划分方式决定你的指标可不可信
这是整个任务里最容易翻车的地方。如果你随机划分训练集和测试集,模型会记住某些药物,测试集里出现的药物在训练集里也出现过,指标会虚高得离谱。正确的做法是按药物划分:测试集里的药物,在训练集里完全不出现。这才模拟了「预测新药相互作用」的真实场景。
from sklearn.model_selection import GroupShuffleSplit def split_by_drug(pairs, groups, test_size=0.2): """按药物分组划分,保证测试集药物不出现在训练集。 groups 是每个样本对应的药物 id 列表(取其中一个即可)。 """ gss = GroupShuffleSplit(n_splits=1, test_size=test_size, random_state=42) train_idx, test_idx = next(gss.split(pairs, groups=groups)) return train_idx, test_idx逻辑说明:GroupShuffleSplit保证同一组的样本不会同时出现在训练和测试。groups 传每个药物对里的第一个药物 id 就行。参数test_size=0.2是常规比例,但如果你的药物总数少,测试集药物太少会导致指标方差大,这时候要做交叉验证而不是单次划分。
4.2 类别不平衡和阈值选择
DDI 数据里,严重的相互作用(比如禁忌)往往只占很小比例。如果你做多分类(无/弱/中/强),类别不平衡会非常明显。处理方式有三种:重采样、类别加权损失、以及调整预测阈值。我一般用加权交叉熵,简单有效。
# 按类别频率的倒数设置权重 class_counts = torch.bincount(train_labels) weights = 1.0 / class_counts.float() weights = weights / weights.sum() criterion = nn.CrossEntropyLoss(weight=weights)逻辑说明:bincount统计每个类别的样本数,取倒数让稀有类获得更大权重。归一化是为了让权重尺度稳定,不影响学习率。注意权重不要设得太极端,否则模型会把所有样本都预测成稀有类,召回上去了但精确率崩了。阈值选择上,不要用默认的 0.5,要在验证集上画 PR 曲线,根据你的业务需求选——药物警戒场景宁可误报也不能漏报,阈值要往低里调。
4.3 评估指标:AUC 之外必须看的东西
AUC 是标配,但它对类别不平衡不敏感,容易给人虚假的安全感。我必看的还有三个:AUPRC(不平衡数据更真实)、召回率@高精度(业务关心的区间)、以及按相互作用类型分层的指标。如果模型在「禁忌」类上召回很低,那这个模型上线就是灾难。
| 指标 | 含义 | 什么时候重点看 |
|---|---|---|
| AUC | 整体排序能力 | 快速对比模型 |
| AUPRC | 正类识别能力 | 类别不平衡时 |
| Recall@Precision | 指定精度下召回 | 业务阈值选择 |
| 分层指标 | 各类型表现 | 上线前必查 |
5. 避坑指南:五个让模型翻车的真实问题
5.1 负样本假阴性导致指标虚高
现象:模型在测试集上 AUC 0.95,一换数据集掉到 0.7。原因:负样本是随机采样的,里面混了大量实际有相互作用但没被标注的药物对。模型学到的是「没标注=无相互作用」这个错误信号。解决:负样本构造时叠加机制过滤,剔除共享靶点、共享代谢酶的药物对;同时用「已知无相互作用」的明确记录优先填充。
5.2 按药物划分后指标暴跌
现象:随机划分 AUC 0.92,按药物划分只有 0.68。原因:随机划分下模型记住了药物身份,测试集药物在训练集见过。按药物划分才是真实泛化能力。解决:接受这个更低的数字,它才是真的。如果按药物划分指标太低,说明模型学的是药物记忆而非相互作用机制,要回去改特征和交互建模。
5.3 SMILES 解析失败静默丢样本
现象:训练正常,但预测时某些药物报错或结果异常。原因:RDKit 解析非法 SMILES 返回 None,如果代码里没检查,None 会传播到后面导致维度错乱或静默跳过。解决:所有 SMILES 入口都做MolFromSmiles校验,记录失败列表,人工核对。生产环境要有兜底:解析失败的药物走指纹缓存或直接拒绝预测。
5.4 特征泄漏:靶点信息混进了标签
现象:模型指标好得不真实,但换一批药就废。原因:构造特征时用了「该药物对是否有相互作用」相关的信息,比如从 DDI 记录里反推的靶点关联,等于把答案喂给了模型。解决:严格区分特征构建阶段和标签使用阶段。靶点、酶特征只能来自药物本身的独立注释,不能来自 DDI 记录。做特征时问自己一句:这个特征在预测时拿得到吗?
5.5 过平滑让深层 GNN 失效
现象:GNN 层数加到 5 层以上,效果不升反降。原因:图卷积反复聚合邻居,所有原子表示趋同,丢失区分度,这就是过平滑。解决:层数控制在 3 到 4 层;加残差连接或 Jumping Knowledge 结构;或者干脆回到指纹基线。不是越深越好,分子图通常很小,3 层足够覆盖大部分拓扑。
6. 把模型推到可用的几个进阶技巧
模型能跑出指标只是起点,要真正可用,还有几件事值得做。第一个是集成:把指纹 MLP 和 GNN 的预测做加权平均,两者犯错模式不同,集成后通常能涨 2 到 3 个点。权重不用调太细,0.5 比 0.5 起步,在验证集上微调即可。第二个是不确定性估计:用 MC Dropout 或者深度集成,对每个预测给出置信度,低置信度的样本交给人工复核,这在药物警戒场景比单纯提高准确率更有价值。
def predict_with_uncertainty(model, fp_a, fp_b, n_forward=20): """MC Dropout 估计预测不确定性。 推理时保持 dropout 开启,多次前向取均值和方差。 """ model.train() # 关键:保持 dropout 激活 preds = [] with torch.no_grad(): for _ in range(n_forward): logits = model(fp_a, fp_b) preds.append(torch.softmax(logits, dim=-1)) preds = torch.stack(preds) mean = preds.mean(dim=0) std = preds.std(dim=0) return mean, std逻辑说明:model.train()让 dropout 在推理时也生效,每次前向得到略有不同的结果,多次采样后方差就是不确定性的代理。n_forward=20是精度和耗时的折中,一般 10 到 30 之间。拿到 std 后,可以设一个阈值,std 高于阈值的预测标记为「需人工复核」,这样模型不是替代人,而是帮人排优先级。
第三个技巧是可解释性回溯。用 GNN 的时候,可以用 GNNExplainer 或者注意力权重,看模型对哪几个子结构最敏感。如果模型判断某对药有相互作用,是因为两个药共享了某个反应性官能团,这个信息对药师来说比一个概率值有用得多。我一般会把 top-5 重要子结构可视化出来,附在预测结果旁边。
最后一个习惯:永远保留一个「傻瓜基线」。我用「两药是否共享靶点」这个规则做基线,如果深度学习模型打不过这个规则,那说明模型没学到东西,别急着上线。这个习惯帮我省过好几次面子——有次模型 AUC 0.85 看着不错,结果规则基线 0.83,等于深度学习白做。后来回去查,发现是特征工程里靶点信息编码方式有问题。做这个方向,敬畏数据比迷信模型重要得多。希望帮到你。
本文还有配套的精品资源,点击获取