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 项优化:
- Encoder KV 预缓存 — 消除每步 16 次 linear
- 合并 QKV 投影 — 3 次 linear 合并为 1 次
- bmm 替代 einsum — 走 NPU Cube 引擎高速路径
- 逐 token 嵌入 — 只嵌入当前 token
- 预计算位置编码 — 一次性计算
- jit.script 全层融合 — 消除 Python 调用开销
- 预乘 attention scale 到权重 — 消除运行时碎片操作
- 预分配 KV cache — index write 替代 torch.cat
- 去除 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 推理结果一致