139

GraphDTA 论文精读:DTA 预测为什么不能只换药物表示

GraphDTA 用图表示替代 SMILES 字符串做亲和力预测。这篇讲清它的改进、以及 DTA 任务中更根本的评测问题。

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 划分评测,那才是部署时面对的情形;
  • 整体相关性高不代表单靶点内排序好——应分靶点计算并报告分布。

延伸资源