DeepFRI_pytorch在昇腾的部署实践
作者:昇腾实战派
知识地图:https://blog.csdn.net/Lumos_Lovegood/article/details/161601003
背景概述
随着蛋白质序列数据库(如UniProt,目前已包含超过1亿条序列)的爆发式增长,如何高效地预测蛋白质功能已成为计算生物学领域的核心挑战。传统的基于序列比对(BLAST)或基于特征工程的方法在面对低序列相似性蛋白质时往往表现不佳,尤其是对于新测序的蛋白质或孤儿蛋白质。
DeepFRI(Deep Functional Residue Identification)由Gligorijević等人于2021年在Nature Communications上发表,是一款结合蛋白质序列信息和三维结构信息的深度学习模型,利用图卷积神经网络(GCN)将蛋白质的三维结构表示为图,从中学习功能相关模式,同时融合蛋白质语言模型提取的序列特征,实现对蛋白质功能的高精度预测。
本文介绍 DeepFRI 模型的PyTorch + 昇腾 Ascend NPU 适配版本——从原始 TensorFlow/Keras 项目中抽取推理与权重转换的最小闭环,在昇腾 AI 平台上完成部署、迁移与精度验证,为蛋白质功能预测的工业级推理场景提供高效、可复现的技术方案。
模型介绍
DeepFRI 概述
DeepFRI 的核心目标是预测蛋白质的生物学功能注释,包括:
- 基因本体(Gene Ontology, GO)注释:分子功能(MF)、生物过程(BP)、细胞组分(CC)
- 酶分类(Enzyme Commission, EC)编号
与传统方法不同,DeepFRI 的创新之处在于将蛋白质结构编码为接触图(Contact Map)——一种图结构表示,节点代表氨基酸残基,边代表残基间的空间接近关系(Cα原子距离 ≤ 10Å),然后利用图卷积网络在该图上传播特征,捕获序列中远距离残基在三维空间上的相互作用模式。
整体架构
DeepFRI 的数据流分为三个阶段:
第一阶段:LSTM 蛋白质语言模型(序列特征提取)
预训练的 LSTM 语言模型(LSTM-LM)在 Pfam 数据库约1000万个蛋白质结构域序列上训练,用于从蛋白质氨基酸序列中提取残基级别的上下文特征。模型由两层单向 LSTM 组成(隐藏维度512),输出拼接后产生1024维的残基级特征向量。
第二阶段:图卷积网络(GCN)处理结构数据
- 接触图被转换为邻接矩阵,每个氨基酸残基对应图中的一个节点
- GCN 接收两个输入:接触图的邻接矩阵 + LSTM 提取的残基级特征矩阵
- 通过多层图卷积操作(3层 MultiGraphConv,每层512维)传播特征
- 使用 SumPooling 将节点级特征聚合为蛋白质级全局表示
第三阶段:功能预测输出
全连接层(FuncPredictor)将蛋白质级表示映射到功能标签空间,输出每个 GO term / EC number 的预测概率。
两条推理路径
| 路径 | 输入 | 特征提取 | 预测网络 |
|---|---|---|---|
| GCN 路径 | PDB 结构文件 / 接触图 | LSTM-LM → 残基特征 + 接触图邻接矩阵 | 图卷积网络 |
| CNN 路径 | 氨基酸序列 | LSTM-LM → 残基特征 | 一维卷积网络(DeepCNN) |
GCN 路径利用了结构信息,预测精度更高;CNN 路径仅需序列,适用于缺少结构数据的场景。
残基级功能解释
DeepFRI 不仅输出蛋白质的功能预测,还利用 Grad-CAM 技术生成残基级别的功能关联图谱(Class Activation Map),标识出可能参与该功能的关键氨基酸位置,为蛋白质功能提供位点级注解。
应用场景
- 蛋白质功能注释:对新测序基因的蛋白产物进行自动功能预测
- 酶工程:预测蛋白酶的EC编号,辅助代谢途径重建
- 药物靶标发现:通过预测分子功能推断蛋白在细胞通路中的角色
- 疾病机制研究:揭示致病蛋白的功能异常
PyTorch + 昇腾 NPU 适配版本
迁移动机
原始 DeepFRI 基于 TensorFlow 1.x / Keras 实现,依赖tf.keras生态进行训练和推理。为在昇腾 Ascend NPU 上高效运行,本项目将推理核心代码转换为 PyTorch 实现,并通过torch_npu适配昇腾硬件加速。
仓库结构
DeepFRI_Pytorch/ ├── deepfrier/ │ ├── torch_layers.py # 图卷积层、池化层、功能预测层的 PyTorch 实现 │ ├── torch_model.py # LSTMLanguageModel、DeepFRIGCN、DeepFRICNN 模型定义 │ ├── torch_predictor.py # 推理预测器封装 │ └── utils.py # 数据处理工具函数 ├── examples/ # 示例输入(PDB文件、接触图、FASTA序列) ├── figs/ # 模型架构图 ├── scripts/ │ └── prepare_models.sh # 权重下载与转换一键脚本 ├── trained_models/ # 转换后的 PyTorch 权重存放目录 ├── benchmark_inference.py # 推理性能基准测试 ├── convert_weights.py # HDF5 → PyTorch state_dict 权重转换 ├── predict.py # 主推理入口 ├── verify_accuracy.py # 精度验证脚本 ├── requirements.txt ├── environment.yml └── setup.py核心实现
图卷积层(MultiGraphConv):对邻接矩阵进行三种归一化处理(原始矩阵、非对称归一化、对称归一化),将节点特征与三种归一化邻接矩阵相乘后拼接,通过线性变换产生输出。
LSTM 语言模型:双层单向 LSTM,输出两层隐状态拼接,产生1024维残基级特征。
权重转换要点:
- TensorFlow
Conv1D权重维度(K, Cin, Cout)→ PyTorch(Cout, Cin, K) - TensorFlow
BatchNorm默认eps=1e-3,PyTorch 中必须保持一致 - CuDNNLSTM 的 HDF5 权重转换到
nn.LSTM时需要按 TensorFlow 官方 HDF5 兼容逻辑做 CuDNN layout 到标准 LSTM layout 的转换,再合并 bias
版本信息
| 软件 | 版本 |
|---|---|
| CANN | 8.2+ |
| Python | 3.10 |
| PyTorch | 2.5.1 |
| torch_npu | 2.5.1 |
环境配置
创建 Conda 环境
conda create-ndeepfri_npupython=3.10-yconda activate deepfri_npu克隆代码
gitclone https://gitcode.com/AI4Science/DeepFRI_Pytorch.gitcdDeepFRI_Pytorch安装依赖
exportPIP_INDEX_URL=https://repo.huaweicloud.com/repository/pypi/simple/ pipinstall-rrequirements.txt配置昇腾环境
source/usr/local/Ascend/ascend-toolkit/set_env.shexportASCEND_RT_VISIBLE_DEVICES=0可使用npu-smi info命令检查驱动是否正常。
模型权重准备
本仓库不直接提交上游预训练权重(体积较大),需要从上游下载并转换。
下载上游 GPU 版权重包
curl-Lhttps://users.flatironinstitute.org/~renfrew/DeepFRI_data/trained_models.tar.gz-otrained_models.tar.gz解压并转换
tarxzf trained_models.tar.gz-C.--no-same-owner python convert_weights.py转换输出示例:
Converting LSTM LM weights... Saved 8 tensors Converting GCN model: mf ... Saved 10 tensors Converting GCN model: bp ... Saved 10 tensors Converting GCN model: cc ... Saved 10 tensors Converting GCN model: ec ... Saved 10 tensors Converting CNN model: ec ... Saved 38 tensors Converting CNN model: mf ... Saved 38 tensors Converting CNN model: bp ... Saved 38 tensors Converting CNN model: cc ... Saved 38 tensors All models converted successfully!也可使用一键脚本:
bashscripts/prepare_models.sh trained_models.tar.gz转换完成后,目录应包含:
trained_models/pytorch/ ├── lstm_lm.pt ├── DeepCNN-MERGED_biological_process.pt ├── DeepCNN-MERGED_cellular_component.pt ├── DeepCNN-MERGED_enzyme_commission.pt ├── DeepCNN-MERGED_molecular_function.pt ├── DeepFRI-MERGED_MultiGraphConv_3x512_fcd_1024_ca_10A_cellular_component.pt ├── DeepFRI-MERGED_MultiGraphConv_3x512_fcd_1024_ca_10A_enzyme_commission.pt ├── DeepFRI-MERGED_MultiGraphConv_3x512_fcd_1024_ca_10A_molecular_function.pt └── DeepFRI-MERGED_MultiGraphConv_3x512_fcd_2048_ca_10A_biological_process.pt迁移适配要点
TensorFlow → PyTorch 关键差异
| 问题 | 解决方案 |
|---|---|
TFBatchNorm默认eps=1e-3 | PyTorch CNN 中设置eps=1e-3保持一致 |
TFConv1D权重(K, Cin, Cout) | 转置为 PyTorch(Cout, Cin, K) |
CuDNNLSTM →nn.LSTM | 按 TF 官方 CuDNN layout 转换逻辑处理,合并 bias |
| GCN 路径对 LSTM 权重更敏感 | 需严格对齐 CuDNNLSTM 到标准 LSTM 的转换 |
昇腾 NPU 适配
PyTorch 版本天然支持通过torch_npu在昇腾 NPU 上运行,无需额外迁移代码。只需在推理时指定设备:
python predict.py--seq'...'-ontmf--devicenpu:0如果运行时环境不完整,可能会在aclInit阶段失败,例如出现507000或1343225857错误码。
推理命令
1. 序列输入,CNN 路径
python predict.py\--seq'SMTDLLSAEDIKKAIGAFTAADSFDHKKFFQMVGLKKKSADDVKKVFHILDKDKDGFIDEDELGSILKGFSSDARDLSAKETKTLMAAGDKDGDGKIGVEEFSTLVAES'\-ontmf\--devicenpu:0\--verbose上游参考输出:
Protein GO-term/EC-number Score GO-term/EC-number name query_prot GO:0005509 0.99769 calcium ion bindingPyTorch NPU 复现结果:
[PASS] query_prot GO:0005509 calcium ion binding expected=0.99769 actual=0.99769 diff=0.0000032. FASTA 输入,CNN 路径
python predict.py\--fasta_fnexamples/pdb_chains.fasta\-ontmf\--devicenpu:0\--verbose3. PDB 输入,GCN 路径
python predict.py\--pdb_fnexamples/pdb_files/1S3P-A.pdb\-ontmf\--devicenpu:0\--verbose上游参考输出:
query_prot GO:0005509 0.99824 calcium ion bindingPyTorch NPU 复现结果:
[PASS] query_prot GO:0005509 calcium ion binding expected=0.99824 actual=0.99824 diff=0.000001精度验证
CPU 验证
python verify_accuracy.py--devicecpu输出示例:
[PASS] query_prot GO:0005509 calcium ion binding expected=0.99769 actual=0.99769 diff=0.000003 [PASS] 1S3P-A GO:0005509 calcium ion binding expected=0.99769 actual=0.99769 diff=0.000003 [PASS] 2J9H-A GO:0004364 glutathione transferase activity expected=0.46937 actual=0.46937 diff=0.000003 [PASS] 2J9H-A GO:0016765 transferase activity, transferring alkyl or aryl (other than methyl) groups expected=0.19910 actual=0.19910 diff=0.000001 [PASS] gcn_pdb GO:0005509 calcium ion binding expected=0.99824 actual=0.99824 diff=0.000001 [OK] MF top: GO:0005509 score=0.99769 (calcium ion binding) [1 predictions] [OK] BP top: GO:0051179 score=0.14491 (localization) [4 predictions] [OK] CC top: GO:0005829 score=0.23144 (cytosol) [7 predictions] [OK] EC no predictions above threshold (expected for some proteins)NPU 验证
python verify_accuracy.py--devicenpu:0精度对齐结果
转换后的 PyTorch 权重在 CPU 上与上游 README 参考值完全对齐:
| 测试用例 | GO term | 期望值 | 复现值 | 差异 |
|---|---|---|---|---|
| query_prot (CNN/seq) | GO:0005509 | 0.99769 | 0.99769 | 0.000003 |
| 1S3P-A (CNN/fasta) | GO:0005509 | 0.99769 | 0.99769 | 0.000003 |
| 2J9H-A (CNN/fasta) | GO:0004364 | 0.46937 | 0.46937 | 0.000003 |
| 2J9H-A (CNN/fasta) | GO:0016765 | 0.19910 | 0.19910 | 0.000001 |
| query_prot (GCN/pdb) | GO:0005509 | 0.99824 | 0.99824 | 0.000001 |
额外 ontology 验证:
- BPtop prediction:
GO:0051179score=0.14491 (localization) - CCtop prediction:
GO:0005829score=0.23144 (cytosol) - EC: 对该测试序列没有超过阈值的预测(符合预期)
性能测试
单条序列推理
python benchmark_inference.py--devicenpu:0--modeseq--ontologymf--warmup3--iters10CPU 基准结果:
| 指标 | 数值 |
|---|---|
| Mean latency | 392.988 ms |
| Median latency | 418.371 ms |
| P95 latency | 482.988 ms |
| Min latency | 315.802 ms |
| Throughput | 2.545 items/s |
FASTA 批量推理
python benchmark_inference.py--devicenpu:0--modefasta--ontologymf--warmup2--iters5CPU 基准结果:
| 指标 | 数值 |
|---|---|
| Items per iteration | 4 |
| Mean latency | 1102.067 ms |
| Median latency | 1269.700 ms |
| P95 latency | 1301.482 ms |
| Throughput | 3.630 items/s |
已知限制
- 本仓库不包含原始 TensorFlow 训练代码,仅聚焦于 PyTorch 推理
- 上游 GCN 权重比 CNN 权重更敏感,因为经过了 CuDNNLSTM →
nn.LSTM的转换路径 - 如果 Ascend 910 会话没有正确映射设备节点,即使 Python 包安装正确,
torch_npu仍会在初始化阶段失败
参考文献
- Gligorijević V, Renfrew P D, Kosciolek T, et al. Structure-based protein function prediction using graph convolutional networks[J]. Nature Communications, 2021, 12(1): 1-14.
- 上游代码仓库:https://github.com/flatironinstitute/DeepFRI
- PyTorch 昇腾适配版:https://gitcode.com/AI4Science/DeepFRI_Pytorch