☰
基于预训练语言模型与标签扩散的蛋白质功能预测实战
2026/10/3 21:37:40 网站建设 项目流程

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对低复杂度序列鲁棒长序列处理差
ProtT53BBFD+UniRef生成式,可做zero-shot推理成本高

我的建议是:如果算力有限,用ESM-2的650M版本,性价比最高。如果追求极致效果且有A100级别的卡,可以上ESM-2 3B。ProtBert在序列长度超过512时需要截断,对长蛋白不友好,慎选。

选ESM的另一个原因是它的embedding维度是1280(650M版本),信息密度足够,而且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, mask=None): # 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), dim=1) # [batch, dim] return out

注意:mask一定要处理,否则padding位置的attention会干扰结果。我一开始忘了加mask,预测结果里出现了大量假阳性。

3.2 标签扩散的数学形式和实现细节

假设我们有N个蛋白,其中前M个有标签(训练集),后N-M个无标签(待预测)。构建一个N×N的相似性矩阵W,W_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_{t+1} = 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, k=10, alpha=0.85, T=3): # embeddings: [N, D] # labels: [N, C], 前M行有值 N = embeddings.shape[0] W = kneighbors_graph(embeddings, k, mode='connectivity', include_self=True) 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(axis=1)) 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 16000

4.2 PLM embedding提取

用HuggingFace的transformers加载ESM-2:

from 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_size=8): inputs = tokenizer(sequence, return_tensors="pt", truncation=True, max_length=1024) 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 标签扩散调参实战

调参的目标是最大化Fmax(CAFA的标准指标)。我的做法是网格搜索:

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上的最优参数是k=15, alpha=0.85, T=3,Fmax达到0.62左右。对比baseline(纯BLAST)的0.48,提升明显。

4.4 后处理与输出格式

预测完得到的是每个蛋白对每个GO术语的分数。需要做几件事:

  • 阈值过滤:分数低于0.1的直接丢掉
  • 层次一致性修正:如果子节点分数>0.5,父节点分数至少设为子节点分数的0.8
  • 输出格式:标准CAFA格式,每行是protein_id + GO_term + score
def post_process(scores, go_parents, threshold=0.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 scores

5. 常见问题与排查技巧实录

5.1 预测结果全是高频标签怎么办

这是最常遇到的问题。GO术语里“binding”“catalytic activity”这些高频标签几乎出现在所有蛋白上,模型倾向于全预测这些,导致precision很低。

解决方案有三个:一是训练时对高频标签降采样,二是损失函数用focal loss,三是后处理时对高频标签设更高的阈值。我一般组合使用,高频标签阈值设0.5,低频标签设0.1。

5.2 embedding提取时显存不够

ESM-2 650M在1024长度下,batch_size=1大概占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_size=32 + fp16,单张A100每小时能处理约5万条序列。

6.2 扩展到其他功能预测任务

这套框架不只能做GO预测,稍微改改就能用于:

  • EC号预测:把GO标签换成EC号,层次结构换成EC的树状结构
  • 亚细胞定位:标签变成定位类别,扩散图不变
  • 蛋白-蛋白相互作用:把标签扩散改成边预测

核心不变的是PLM embedding + 图传播这个范式。我最近在做一个抗菌肽识别项目,也是用ESM embedding + kNN分类,效果比传统特征工程好很多。

6.3 和结构信息的结合

纯序列方法的天花板在于,有些功能只有看结构才能确定。如果有AlphaFold2预测的结构,可以把结构embedding和序列embedding拼接,再走标签扩散。我试过在CAFA3上拼接GVP(Geometric Vector Perceptron)的结构embedding,Fmax从0.62提到0.67。代价是推理时间增加3倍,因为要跑AF2。

如果算力允许,这个方向值得投入。尤其是对那些序列同源性低但结构相似的蛋白,结构信息能救命。

最后分享一个我在实际项目中的体会:这套方案的效果高度依赖embedding质量,而embedding质量又依赖PLM的预训练数据覆盖度。如果你的目标蛋白是某种极端环境微生物的,而PLM训练集里这类序列很少,效果会打折扣。这时候可以考虑用目标物种的序列对PLM做继续预训练(continue pretraining),哪怕只用几万条序列,也能提升3-5个点。这个trick在文献里提得不多,但实测有效。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询