文件最后提交记录最后更新时间
2 个月前
2 个月前
2 个月前
README

ESM-IF1 昇腾 NPU 适配与性能优化

将 Facebook Research 的 ESM-IF1 逆向折叠模型 (esm_if1_gvp4_t16_142M_UR50) 从 GPU 迁移到华为昇腾 Ascend 910 NPU,并提供约 2.2 倍的推理加速。

环境要求

硬件

  • 华为昇腾 Ascend 910 系列 NPU

软件依赖

依赖 版本要求 说明
Python >= 3.11
PyTorch 2.9.0 需配套 torch_npu
torch_npu 2.9.0 昇腾 PyTorch 插件
CANN 8.5.1 昇腾计算框架
fair-esm 2.0.0 pip install fair-esm
biotite >= 1.0 pip install biotite
scipy fair-esm 依赖,自动安装
numpy fair-esm 依赖,自动安装

快速开始

1. 安装基础环境

pip install fair-esm biotite

确保 torch_npu 已正确安装且 CANN 环境已加载。

2. 应用补丁

补丁包含两部分改动:

  • NPU 适配:修改 esm 库中 2 个文件,解决设备不匹配问题
  • 性能优化:新增 esm_if1_npu_optimized_v2.py,提供 ~2.2x 加速

方式 A:patch 命令(推荐)

需要知道 fair-esm 的安装路径:

ESM_PATH=$(python3 -c "import esm; print(esm.__path__[0])")

# 例如: /usr/local/python3.11.14/lib/python3.11/site-packages/esm
# 补丁路径前缀为 esm/,所以去掉末尾的 /esm
SITE_DIR=$(dirname $ESM_PATH)

cd $SITE_DIR
patch -p1 < /path/to/esm_if1_npu_v2.diff

方式 B:git apply

适用于从源码管理 esm 的场景:

cd /path/to/esm-source
git apply /path/to/esm_if1_npu_v2.diff

方式 C:git am(保留 commit 信息)

cd /path/to/esm-source
git am /path/to/esm_if1_npu_v2.patch

3. 验证安装

基础推理(使用原始 model.sample)

python3 test_inference.py

预期输出:

加载模型...
模型已加载到 npu:0
输入坐标形状: torch.Size([395, 3, 3])(残基数, 3种原子, 3维坐标)

生成 3 个序列样本(温度=0.6):
样本 1: SRLEKWWASPLTQESDPGASSDDQLFDPEDFLKLLPDTADENLPNWFKFHVGIRKNRMYE... (长度 395)
样本 2: STYKRWIASPLTQLLDPELTEEDLLFDPRLMDALLPEKEEENEPAWRRFWSGIRKYRMYD... (长度 395)
样本 3: TRLTAWRASPLTQEKDPDLSSADLLFDPADLLALIPDKVPDKLPAWLQFQVQIRRHRMYE... (长度 395)

推理完成!

性能测试(使用优化模块)

python3 test_time.py

预期输出:

平均推理时间: ~3.1 秒
序列长度: 395 残基

使用优化模块

优化模块 OptimizedESMIF1 封装了原始模型,提供相同的采样接口但速度更快:

import torch
import esm
from esm_if1_npu_optimized_v2 import OptimizedESMIF1

device = torch.device("npu:0")
model, _ = esm.pretrained.esm_if1_gvp4_t16_142M_UR50()
model = model.to(device).eval()

opt = OptimizedESMIF1(model)

# coords: numpy array, shape [L, 3, 3]
# confidence: list of floats, length L
sequence = opt.sample(coords, confidence=[1.0] * len(coords), temperature=1.0)

性能参考(Ascend 910, 3 次平均)

链长度 原始 (s) 优化 V2 (s) 加速比
100 0.749 0.341 2.19x
300 2.146 0.948 2.26x
647 5.404 3.119 1.73x
1000 10.625 4.738 2.24x
1686 18.210 8.045 2.26x

补丁内容详解

NPU 适配修改(2 个文件)

esm/inverse_folding/gvp_transformer.py:

  • 自动检测输入 coords 所在设备(NPU),沿调用链传递 device 参数
  • 在正确设备上创建 sampled_tokens,避免 CPU/NPU 设备不匹配

esm/inverse_folding/util.py:

  • torch.tensor() 改为 torch.as_tensor().detach(),避免对已有 NPU tensor 的重复创建和梯度追踪

性能优化(新增 1 个文件)

esm_if1_npu_optimized_v2.py — 9 项优化:

  1. Encoder KV 预缓存 — 消除每步 16 次 linear
  2. 合并 QKV 投影 — 3 次 linear 合并为 1 次
  3. bmm 替代 einsum — 走 NPU Cube 引擎高速路径
  4. 逐 token 嵌入 — 只嵌入当前 token
  5. 预计算位置编码 — 一次性计算
  6. jit.script 全层融合 — 消除 Python 调用开销
  7. 预乘 attention scale 到权重 — 消除运行时碎片操作
  8. 预分配 KV cache — index write 替代 torch.cat
  9. 去除 float()/type_as() — 全程 fp32

注意事项

  • ESM 库升级:pip install --upgrade fair-esm 会覆盖补丁,需重新应用
  • 设备选择:测试脚本默认使用 npu:0,多卡环境按需修改 torch.device("npu:X")
  • torch_npu 警告:Cannot create tensor with internal format 警告不影响推理结果
  • 精度:全程 fp32,与 CPU/GPU 推理结果一致