GraphDTA(Nguyen 等,Bioinformatics 2021)把药物-靶点亲和力(DTA)预测中的药物表示从 SMILES 字符串换成了分子图。改进是真实的——但这篇论文更重要的价值在于它引出的问题:DTA 任务的评测设置本身有多大问题。
模型结构
# GraphDTA 的架构很直接:
#
# 药物侧:分子图 → GNN(GCN/GAT/GIN/组合)→ 药物向量
# 蛋白侧:氨基酸序列 → 1D CNN → 蛋白向量
# 融合: 拼接两个向量 → 全连接层 → 预测亲和力
#
# 相对 DeepDTA(见 140)的唯一改动:
# 药物侧从「SMILES 字符 CNN」换成「分子图 GNN」
#
# 报告的结果:
# 在 Davis 与 KIBA 数据集上,MSE 降低、CI 提高
#
# 【论文自己的结论】:
# 图表示优于字符串表示
#
# 但这个结论需要放在更大的背景下看
更根本的问题:蛋白侧太弱
| 侧 | GraphDTA 的处理 | 问题 |
|---|---|---|
| 药物 | 分子图 + GNN | 合理 |
| 蛋白 | 氨基酸序列 + 1D CNN | 丢失了几乎所有结构信息 |
这是整类 DTA 方法的共同短板。结合发生在三维的结合口袋中,而一维序列的 CNN 很难学到「哪些残基构成口袋、它们在空间上如何排列」。把药物侧从字符串升级到图,却不动蛋白侧——改进的空间从一开始就受限。
评测设置的关键问题
# Davis 与 KIBA 数据集的结构:
# Davis: 68 个药物 × 442 个激酶 = 约 3 万个测量
# KIBA: 2111 个药物 × 229 个靶点
#
# 【关键】:这是一个「矩阵」,不是独立样本
#
# 随机划分的后果:
# 同一个药物的其它靶点数据在训练集中
# 同一个靶点的其它药物数据在训练集中
# → 模型可以学「这个药物平均活性高」
# 和「这个靶点平均容易被抑制」
# → 【不需要真正理解相互作用就能得高分】
#
# 验证方法(很有说服力的对照实验):
# 基线 1:只用药物 ID 的平均值预测
# 基线 2:只用靶点 ID 的平均值预测
# 基线 3:药物均值 + 靶点均值(类似矩阵分解)
#
# → 多项研究发现,这些【平凡基线】
# 在随机划分下就能取得接近深度模型的成绩
#
# 正确的划分(TDC 提供,见 133):
# cold drug: 测试集药物在训练集中完全没出现
# cold protein: 测试集靶点在训练集中完全没出现
# cold both: 两者都没出现
#
# → 在这些设置下,所有方法的性能都【大幅下降】
# → 而这才是实际部署时面对的情形
# 【读 DTA 论文时应该问】:
# 1) 用了什么划分?
# 2) 有没有报告平凡基线的成绩?
# 3) cold split 下的结果是多少?
# 三个问题都答不上来的论文,结论要打折扣
实现要点
import torch
import torch.nn as nn
from torch_geometric.nn import GINConv, global_mean_pool
class GraphDTA(nn.Module):
def __init__(self, num_features_drug=78, num_features_prot=25,
embed_dim=128, output_dim=128, dropout=0.2):
super().__init__()
# 药物侧:GIN
nn1 = nn.Sequential(nn.Linear(num_features_drug, embed_dim),
nn.ReLU(), nn.Linear(embed_dim, embed_dim))
self.conv1 = GINConv(nn1)
nn2 = nn.Sequential(nn.Linear(embed_dim, embed_dim),
nn.ReLU(), nn.Linear(embed_dim, embed_dim))
self.conv2 = GINConv(nn2)
self.fc_drug = nn.Linear(embed_dim, output_dim)
# 蛋白侧:embedding + 1D CNN
self.embed_prot = nn.Embedding(num_features_prot + 1, embed_dim)
self.conv_prot = nn.Conv1d(embed_dim, 32, kernel_size=8)
self.fc_prot = nn.Linear(32 * 121, output_dim)
# 融合
self.fc1 = nn.Linear(2 * output_dim, 1024)
self.fc2 = nn.Linear(1024, 512)
self.out = nn.Linear(512, 1)
self.dropout = nn.Dropout(dropout)
self.relu = nn.ReLU()
def forward(self, data):
x, edge_index, batch = data.x, data.edge_index, data.batch
x = self.relu(self.conv1(x, edge_index))
x = self.relu(self.conv2(x, edge_index))
x = global_mean_pool(x, batch)
xd = self.relu(self.fc_drug(x))
xt = self.embed_prot(data.target)
xt = self.relu(self.conv_prot(xt.permute(0, 2, 1)))
xt = self.relu(self.fc_prot(xt.flatten(1)))
xc = torch.cat([xd, xt], dim=1)
xc = self.dropout(self.relu(self.fc1(xc)))
xc = self.dropout(self.relu(self.fc2(xc)))
return self.out(xc)
# 【务必同时实现的对照基线】:
def trivial_baseline(train_df, test_df):
"""只用药物均值与靶点均值预测"""
global_mean = train_df["y"].mean()
drug_mean = train_df.groupby("drug_id")["y"].mean()
prot_mean = train_df.groupby("prot_id")["y"].mean()
preds = []
for _, row in test_df.iterrows():
d = drug_mean.get(row["drug_id"], global_mean)
p = prot_mean.get(row["prot_id"], global_mean)
preds.append((d + p) / 2)
return preds
# 如果深度模型打不过这个基线,
# 说明它没有学到真正的相互作用
评价指标的选择
- MSE / RMSE:常用,但对异常值敏感;
- CI(Concordance Index):衡量排序一致性,比 MSE 更贴近实际需求(我们通常关心排序而非绝对值);
- Spearman / Pearson:相关性;
- 关键提醒:整体相关性高不代表在单个靶点内的排序好。DTA 数据跨靶点的活性范围很宽,跨靶点的方差会让整体相关性看起来很好——应该分靶点计算相关性并报告分布;
- 实际关心的问题:对某个靶点,模型能否把真正的活性分子排在前面?这需要按靶点分组评价。
这一系列工作的价值与局限
| 价值 | 局限 | |
|---|---|---|
| 方法 | 确立了 DTA 的深度学习范式 | 蛋白侧表示过于简化 |
| 表示 | 验证了图优于字符串 | 提升幅度受限于整体架构 |
| 评测 | 提供了可比较的基准 | 划分方式导致高估 |
| 实用性 | 可做大规模粗筛 | 不能替代对接或实验 |
现在的更好方向:用蛋白语言模型表示(见 143《ESM-2 论文精读》)或直接用结构(见 113《AlphaFold3 论文精读》、121《Boltz-2 技术报告精读》)替代序列 CNN——把蛋白侧的信息量补上,才是这类方法的主要改进空间。
关键要点
- 只把药物侧从字符串升级到图,蛋白侧的序列 CNN 才是真正的瓶颈;
- 随机划分下,只用药物均值 + 靶点均值的平凡基线就能接近深度模型;
- 必须用 cold drug / cold protein 划分评测,那才是部署时面对的情形;
- 整体相关性高不代表单靶点内排序好——应分靶点计算并报告分布。
延伸资源
- DeepDTA:140《DeepDTA 论文精读》;MolTrans:141《MolTrans 论文精读》;DTI 模型:157《AI DTI 预测模型》;
- TDC 划分:133《TDC 论文精读》;Benchmark 陷阱:169《Benchmark 陷阱》;ESM-2:143《ESM-2 论文精读》。