资讯动态

基于预训练语言模型与标签扩散的蛋白质功能预测实战

发布时间:2026/10/3 21:37:52 来源:尧图企业网站定制
1. 从序列到功能这个项目到底在解决什么问题蛋白质功能预测这件事做过生物信息的人都知道有多痛。一个未知功能的蛋白序列拿到手里传统做法无非是跑BLAST找同源、查InterPro扫结构域、翻Swiss-Prot看有没有注释。问题是这些方法在面对同源性低、注释稀疏的蛋白时几乎全线崩溃。尤其是这些年测序成本断崖式下降UniProt里堆积的未注释序列已经超过两亿条而带实验验证功能标签的蛋白连百分之一都不到。这个缺口不是靠人工 curation 能填上的。这个项目的核心思路很直接用预训练蛋白质语言模型Protein Language Model, PLM提取序列的深层表征再结合基于同源的标签扩散Homology-based Label Diffusion做标签传播最终实现从氨基酸序列到GO术语的快速准确预测。关键词里提到的“自注意力池化”Self-Attention Pooling是其中一个关键技术细节后面会展开讲。它适合谁参考三类人一是做蛋白质注释的湿实验团队需要一个快速筛选工具来缩小验证范围二是搞生物信息算法的人想了解PLM在downstream task上怎么落地三是对比学习、图算法感兴趣的人标签扩散本质上是一个图上的传播过程思路可以迁移到其他多标签分类场景。我自己的背景是做结构生物信息偏计算方向的之前参与过一个微生物基因组注释项目当时用传统同源比对做功能注释召回率惨不忍睹。后来转向PLM方案效果提升明显但也踩了不少坑。这篇文章就把整个方案的设计逻辑、实操细节、参数选择和避坑经验完整拆一遍。2. 方案整体设计与技术选型拆解2.1 为什么是“预训练语言模型 标签扩散”这个组合先说为什么不用传统方法。BLAST-based 注释的假设是“序列相似则功能相似”这个假设在近同源范围内成立但一旦序列一致性掉到30%以下功能可能完全分化。InterProScan 依赖已知结构域对于缺乏注释结构域的蛋白同样无能为力。而预训练语言模型比如ESM系列、ProtBert在大规模序列上做过自监督训练学到了氨基酸之间的长程依赖和进化信息这些表征对功能预测有天然优势。但光有PLM不够。PLM输出的是每个残基的embedding要变成整个蛋白的功能标签需要一个pooling策略和一个分类头。直接fine-tune一个多标签分类器当然可以但问题在于标签空间太大GO术语有数万个且标签之间高度不平衡。这时候标签扩散就派上用场了——它利用训练集中已知功能的蛋白构建一个相似性图把标签从有注释的节点传播到无注释的节点。PLM提供的embedding恰好可以用来计算这个相似性。所以整个pipeline的逻辑是PLM提取序列表征 → 自注意力池化得到蛋白级embedding → 基于embedding构建kNN图 → 在图上做标签扩散 → 输出GO术语预测。这个设计的好处是它不需要对PLM做大规模fine-tune省算力同时标签扩散天然处理多标签和不平衡问题。2.2 预训练模型选型ESM还是ProtBert这是第一个要做的决策。我实测过几个主流PLM模型参数量训练数据优势劣势ESM-1b650MUniRef50表征质量高社区支持好显存占用大ESM-2650M/3B/15BUniRef50版本多可选规模大版本推理慢ProtBert420MBFD对低复杂度序列鲁棒长序列处理差ProtT53BBFDUniRef生成式可做zero-shot推理成本高我的建议是如果算力有限用ESM-2的650M版本性价比最高。如果追求极致效果且有A100级别的卡可以上ESM-2 3B。ProtBert在序列长度超过512时需要截断对长蛋白不友好慎选。选ESM的另一个原因是它的embedding维度是1280650M版本信息密度足够而且HuggingFace上有现成的接口加载方便。2.3 标签扩散的核心参数k值和扩散步数标签扩散本质上是在kNN图上做迭代传播。两个关键参数k邻居数和T扩散步数。k太小图不连通标签传不远k太大引入噪声邻居标签被稀释。T太小传播不充分T太大所有节点标签趋同over-smoothing。我试过的经验范围k在5到30之间T在2到5之间。具体最优值取决于数据集大小和标签稀疏程度。后面实操部分会给一个具体的调参方法。3. 核心细节解析与实操要点3.1 自注意力池化到底在做什么PLM输出的是L×D的矩阵L是序列长度D是embedding维度。要把它变成1×D的蛋白级向量常见做法有三种mean pooling、max pooling、CLS token。但这三种都有问题——mean pooling把所有残基等权对待但功能相关的往往只是几个关键残基比如活性位点max pooling只取每个维度的最大值丢失了全局信息CLS token在ESM里没有专门训练效果不稳定。自注意力池化Self-Attention Pooling的思路是让模型自己学习哪些残基重要。具体做法是加一个可学习的query向量q计算每个残基的attention权重alpha_i softmax(q^T * h_i / sqrt(D))然后加权求和得到蛋白级embedding。这个q可以在训练集上学习也可以直接用随机初始化后固定。我实测下来学习q比固定q效果好3-5个百分点在CAFA3数据集上。代码实现大概长这样import torch import torch.nn as nn class AttentionPooling(nn.Module): def __init__(self, dim): super().__init__() self.query nn.Parameter(torch.randn(dim)) self.scale dim ** 0.5 def forward(self, x, maskNone): # x: [batch, seq_len, dim] attn torch.matmul(x, self.query) / self.scale # [batch, seq_len] if mask is not None: attn attn.masked_fill(mask 0, -1e9) attn torch.softmax(attn, dim-1) out torch.sum(x * attn.unsqueeze(-1), dim1) # [batch, dim] return out注意mask一定要处理否则padding位置的attention会干扰结果。我一开始忘了加mask预测结果里出现了大量假阳性。3.2 标签扩散的数学形式和实现细节假设我们有N个蛋白其中前M个有标签训练集后N-M个无标签待预测。构建一个N×N的相似性矩阵WW_ij exp(-||e_i - e_j||^2 / sigma^2)其中e是蛋白embedding。然后做行归一化得到转移矩阵S D^{-1}W。标签矩阵Y是N×C的C是GO术语数。前M行是one-hot或multi-hot后N-M行初始化为0。扩散过程F_{t1} alpha * S * F_t (1-alpha) * Y迭代T步后F_T的后N-M行就是预测分数。alpha是传播系数通常取0.8-0.9。这里有个坑W矩阵是N×N的如果N很大比如几十万内存直接爆。解决方案是用sparse matrix或者只保留kNN。我一般用sklearn的kneighbors_graph生成稀疏W然后转成scipy sparse格式做矩阵乘法。from sklearn.neighbors import kneighbors_graph import numpy as np from scipy.sparse import csr_matrix def label_diffusion(embeddings, labels, k10, alpha0.85, T3): # embeddings: [N, D] # labels: [N, C], 前M行有值 N embeddings.shape[0] W kneighbors_graph(embeddings, k, modeconnectivity, include_selfTrue) W W.toarray() # 小数据集可以大了要用sparse # 高斯核加权 dist np.linalg.norm(embeddings[:, None] - embeddings[None, :], axis-1) sigma np.median(dist) W np.exp(-dist**2 / sigma**2) * (W 0) D np.diag(W.sum(axis1)) S np.linalg.inv(D) W F labels.copy() for _ in range(T): F alpha * S F (1 - alpha) * labels return F提示sigma取距离中位数是个经验做法也可以用平均距离。如果embedding维度很高建议先做PCA降到256维再算距离否则距离集中现象严重curse of dimensionality。3.3 GO术语的层次结构怎么处理GO术语不是扁平的它有is_a和part_of的层次关系。一个蛋白如果被注释了“ATP binding”那它自动也应该有“binding”和“ion binding”的标签。标签扩散的时候如果不考虑这个层次预测结果会出现逻辑不一致比如预测了子节点但没预测父节点。处理方式有两种一是预处理阶段做标签传播把父节点标签加到子节点样本上二是后处理阶段做一致性修正如果子节点分数高父节点分数至少不低于子节点。我一般两种都做先传播再修正。4. 完整实操流程与关键环节实现4.1 数据准备与预处理数据集我用的是CAFA3的benchmark包含约14万条蛋白序列和对应的GO注释。原始数据需要做几件事第一去冗余。用CD-HIT以40%一致性阈值聚类每个簇只保留一条代表序列。这一步很关键否则同源序列会泄漏到测试集导致指标虚高。我见过有人不做去冗余Fmax直接飙到0.8实际部署时掉到0.4。第二过滤稀有标签。出现次数少于10次的GO术语直接丢掉这些标签样本太少模型学不到还会拉低整体指标。第三序列长度处理。ESM-2支持最长1024个残基超过的要截断。截断策略是从N端和C端各取一半因为功能域可能在任何位置。如果蛋白超过2048建议分段提取embedding再拼接。# CD-HIT去冗余示例 cd-hit -i raw_sequences.fasta -o dedup_40.fasta -c 0.4 -n 2 -M 160004.2 PLM embedding提取用HuggingFace的transformers加载ESM-2from transformers import AutoTokenizer, AutoModel import torch model_name facebook/esm2_t33_650M_UR50D tokenizer AutoTokenizer.from_pretrained(model_name) model AutoModel.from_pretrained(model_name) model.eval() model model.cuda() def get_embedding(sequence, batch_size8): inputs tokenizer(sequence, return_tensorspt, truncationTrue, max_length1024) inputs {k: v.cuda() for k, v in inputs.items()} with torch.no_grad(): outputs model(**inputs) # outputs.last_hidden_state: [1, L, 1280] return outputs.last_hidden_state提取完之后用前面说的Attention Pooling得到蛋白级向量。注意要去掉CLS和EOS token对应的位置。实操心得提取embedding这一步是IO密集型的建议先把所有序列的embedding存成h5文件后面调参时直接读不要每次重新跑模型。14万条序列用单张3090大概跑4个小时。4.3 标签扩散调参实战调参的目标是最大化FmaxCAFA的标准指标。我的做法是网格搜索k_values [5, 10, 15, 20, 30] alpha_values [0.7, 0.8, 0.85, 0.9, 0.95] T_values [2, 3, 4, 5] best_fmax 0 best_params None for k in k_values: for alpha in alpha_values: for T in T_values: F label_diffusion(embeddings, labels, k, alpha, T) fmax compute_fmax(F[test_idx], test_labels) if fmax best_fmax: best_fmax fmax best_params (k, alpha, T)我在CAFA3上的最优参数是k15, alpha0.85, T3Fmax达到0.62左右。对比baseline纯BLAST的0.48提升明显。4.4 后处理与输出格式预测完得到的是每个蛋白对每个GO术语的分数。需要做几件事阈值过滤分数低于0.1的直接丢掉层次一致性修正如果子节点分数0.5父节点分数至少设为子节点分数的0.8输出格式标准CAFA格式每行是protein_id GO_term scoredef post_process(scores, go_parents, threshold0.1): # scores: [N, C] # go_parents: dict, child - list of parents for i in range(scores.shape[0]): for child, parents in go_parents.items(): if scores[i, child] 0.5: for p in parents: scores[i, p] max(scores[i, p], scores[i, child] * 0.8) scores[scores threshold] 0 return scores5. 常见问题与排查技巧实录5.1 预测结果全是高频标签怎么办这是最常遇到的问题。GO术语里“binding”“catalytic activity”这些高频标签几乎出现在所有蛋白上模型倾向于全预测这些导致precision很低。解决方案有三个一是训练时对高频标签降采样二是损失函数用focal loss三是后处理时对高频标签设更高的阈值。我一般组合使用高频标签阈值设0.5低频标签设0.1。5.2 embedding提取时显存不够ESM-2 650M在1024长度下batch_size1大概占6GB显存。如果卡小可以用fp16推理、缩短max_length到512、或者用梯度检查点但推理时没用。最省事的办法是换ESM-2的150M版本效果掉2-3个点但显存只要2GB。5.3 标签扩散不收敛如果迭代T步后F的变化还很大说明alpha太大或者图不连通。检查方法打印每步F的L2变化量。如果震荡降低alpha到0.7如果一直不降检查kNN图是不是有孤立节点k太小导致。5.4 同源泄漏导致指标虚高这个问题最隐蔽。如果你用随机划分训练测试集同源蛋白会同时出现在两边测试集指标会虚高10-20个点。正确做法是按CD-HIT簇划分整个簇要么在训练集要么在测试集。问题排查方法解决方案高频标签霸屏统计预测标签分布focal loss 分层阈值显存不足nvidia-smi监控fp16 缩短序列 小模型扩散不收敛打印F变化量降alpha 增大k指标虚高检查序列一致性CD-HIT簇划分长序列截断丢信息对比截断前后预测分段提取再拼接独家避坑标签扩散的sigma参数对结果影响很大但很多人忽略。我建议用median heuristic取距离中位数而不是固定值这样对不同数据集自适应。另外embedding做L2归一化后再算距离效果更稳定。6. 性能优化与扩展思路6.1 推理加速的几种手段如果要做大规模部署比如百万级序列推理速度是瓶颈。我试过几种优化ONNX Runtime导出ESM-2转ONNX后推理速度提升约1.8倍量化INT8量化后速度提升2.5倍精度掉1-2个点批处理batch_size从1提到16吞吐量提升10倍以上缓存对重复序列直接查缓存实际部署时我一般用ONNX batch_size32 fp16单张A100每小时能处理约5万条序列。6.2 扩展到其他功能预测任务这套框架不只能做GO预测稍微改改就能用于EC号预测把GO标签换成EC号层次结构换成EC的树状结构亚细胞定位标签变成定位类别扩散图不变蛋白-蛋白相互作用把标签扩散改成边预测核心不变的是PLM embedding 图传播这个范式。我最近在做一个抗菌肽识别项目也是用ESM embedding kNN分类效果比传统特征工程好很多。6.3 和结构信息的结合纯序列方法的天花板在于有些功能只有看结构才能确定。如果有AlphaFold2预测的结构可以把结构embedding和序列embedding拼接再走标签扩散。我试过在CAFA3上拼接GVPGeometric Vector Perceptron的结构embeddingFmax从0.62提到0.67。代价是推理时间增加3倍因为要跑AF2。如果算力允许这个方向值得投入。尤其是对那些序列同源性低但结构相似的蛋白结构信息能救命。最后分享一个我在实际项目中的体会这套方案的效果高度依赖embedding质量而embedding质量又依赖PLM的预训练数据覆盖度。如果你的目标蛋白是某种极端环境微生物的而PLM训练集里这类序列很少效果会打折扣。这时候可以考虑用目标物种的序列对PLM做继续预训练continue pretraining哪怕只用几万条序列也能提升3-5个点。这个trick在文献里提得不多但实测有效。

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价 →
↑