用 SMILESGNN 把分子图和序列交叉注意力融合能让毒性预测变得可解释
这篇 arXiv:2609.28553v1 论文提出了一个叫 SMILESGNN 的多模态架构,核心是通过 cross-attention 融合 SMILES Transformer 编码器和 GATv2 图编码器。它解决了单模态模型在临床毒性预测中要么缺乏图属性解释(序列模型)、要么性能不足的问题,在 ClinTox 数据集上用 0.4M 参数跑出了 0.987 的 AUC-ROC。
为什么需要这种交叉注意力融合架构
在药物毒性预测里,最头疼的是类别极度不平衡和基于骨架的泛化问题。以前大家要么用 SMILES Transformer 这种序列模型,要么用图神经网络(GNN),但这两者有天然的矛盾:序列模型虽然捕捉全局信息强,但没法直接告诉你分子结构里哪个原子导致了毒性;GNN 能做子结构分析,但往往在预测精度上不如预训练的 Transformer。
SMILESGNN 的逻辑是把两者都跑一遍,但不用简单的拼接(Concatenation),而是用交叉注意力机制做融合。这样在预测阶段,它依然保留了一个显式的图分支,可以直接挂载 GNNExplainer 来分析具体是哪个子结构触发了毒性预测,解决了临床上对「可解释性」的刚需。
模型具体的实现路径与变体
SMILESGNN 的结构分为两个主要分支:一个是处理 SMILES 字符串的 Transformer 编码器,另一个是处理分子图的 GATv2 编码器。
- 基础版本: 采用标准的 SMILES Transformer 和 GATv2,通过 cross-attention 模块将两种模态的特征进行融合。
- 增强版本(SMILESGNN-PT): 为了提升效果,这个变体把 Backbone 换成了预训练的 ChemBERTa-2。
在 ClinTox 和 Tox21 上的实测数据
论文给出了具体的对比数字,证明了这种融合方式在极小参数量下依然能打。
- ClinTox 数据集:
- Tox21 数据集(包含 12 个任务):
关于可解释性的技术实现
如果你想在自己的项目里实现类似的解释功能,可以参考它的逻辑:
1. 特征提取: 分别通过 $\text{SMILES} → \text{Transformer}$ 和 $\text{Graph} → \text{GATv2}$ 得到两组特征向量。
2. 交叉融合: 使用 $\text{Attention}(Q_{graph}, K_{smiles}, V_{smiles})$ 让图特征去查询序列特征,增强表示能力。
3. 回溯解释: 因为 GATv2 分支在预测路径中被保留,可以直接调用 GNNExplainer。
# 伪代码逻辑参考
# 1. 输入分子数据
smiles_feat = smiles_transformer_encoder(smiles_input)
graph_feat = gatv2_encoder(graph_input)
# 2. 交叉注意力融合 (Cross-Attention)
# 用图特征作为 Query,序列特征作为 Key 和 Value
fused_feat = cross_attention(query=graph_feat, key=smiles_feat, value=smiles_feat)
# 3. 毒性预测
prediction = mlp_classifier(fused_feat)
# 4. 可解释性分析 (利用保留的 GATv2 分支)
explainer = GNNExplainer(model=gatv2_encoder, graph=graph_input)
important_substructure = explainer.explain_node(node_idx)
总的来说,这篇论文证明了 cross-attention 是一个非常实用的融合替代方案。它在不牺牲预测性能的前提下,通过保留图分支,让开发者能够通过 GNNExplainer 清楚地看到分子中哪些部分是「有毒」的,这比单纯拿一个 AUC 分数要有用得多。
0.4M 参数能跑 0.987 简直是神仙效率,我之前死磕大模型反而过拟合了,赶紧试下这个架构。