liuhongwei-2026/sptransformer-npu
模型介绍
文件和版本
Pull Requests
讨论
分析

SpliceTransformer(multimolecule/sptransformer)在昇腾 NPU 上的服务化推理

1. 简介

  • 模型来源: multimolecule/sptransformer(gitcode 镜像: hf_mirrors/multimolecule/sptransformer)
  • 模型任务: 组织特异性 RNA 剪接位点预测 —— 对 pre-mRNA/RNA 序列的每个核苷酸位置预测三类剪接位点分数(no_splice / acceptor(受体位点)/ donor(供体位点)),并输出 15 个组织的剪接使用度(tissue_0 ~ tissue_14)
  • 参数量: 17.074M(17,074,386)
  • 模型架构: SpTransformerModel —— 两个 SpliceAI 风格膨胀残差卷积特征提取器(hidden 128 / 64,各 16 层,两侧各 4,000 nt 上下文)+ 局部(windowed-local)与 Sinkhorn(sorted-bucket)混合稀疏注意力 Transformer(8 层 × 8 头,attention hidden 256)
  • 适配状态: SUCCESS
  • 适配时间: 2026-08-19

SpliceTransformer(SPT,ShenLab-Genomics,Nature Methods 2024)以原初 RNA 序列为输入, 复用两个 SpliceAI 风格卷积编码器提取逐位置序列特征,再用混合稀疏注意力块建模长程依赖, 对每个位置预测其作为剪接受体/供体的概率,以及组织特异性的剪接使用度分数。

2. 验证环境

组件版本
torch2.9.0
torch-npu2.9.0.post1
transformers5.9.0
multimolecule0.2.1
fastapi0.123.10
CANN8.5.1
NPUAscend 910(64GB HBM)

3. 模型结构

SpTransformer 将 one-hot 编码(ACGU,N 为全零掩码,RnaTokenizer 将 T 转 U)的序列两侧各补 context = 4000 nt 零填充,经两个膨胀残差卷积编码器提取逐位置特征后,与可训练投影特征融合, 送入混合稀疏注意力块(8 层;前 2 头为窗口局部注意力、后 6 头为 Sinkhorn 排序桶注意力), 最后经预测头输出每位置 num_splice_labels + num_tissues = 18 通道分数并中心裁剪回输入长度。 关键配置:

配置项值
模型类型SpTransformerModel
层数8(稀疏注意力)× 8 头(2 局部 + 6 Sinkhorn)
卷积编码器2 个 SpliceAI 风格(hidden 128 / 64,各 16 残差块)
attention_hidden_size256
context(两侧填充)4000(4000 + 4000)
max_seq_len(注意力块上限)8192
bucket_size64
输出通道3 剪接位点(softmax 概率)+ 15 组织使用度(原生 logits 尺度)
词汇表5(A/C/G/U/N + 填充,RnaTokenizer T→U)
权重文件model.safetensors(68,412,160 B)
权重 sha256cc7e5f725c3a941b68427eb4f87ffb790a38bb901f80d6a26d3a92f1bf3c8f1a

输入长度说明:模型两侧各补 4,000 nt 上下文,注意力块最大序列长度为 max_seq_len = 8192。因此输入长度 L 满足 L + 2*context <= max_seq_len (即 L <= 192 nt)时,每个位置均保有完整 4,000 nt 侧翼上下文;更长序列仍可 推理(输出长度 = 输入长度),但靠近两端的位点侧翼上下文会被中心裁剪、预测退化。

4. 服务化推理验证结果

4.1 CLI 推理(HBB pre-mRNA 演示序列,1401 nt)

[demo] 未提供序列,使用 HBB pre-mRNA 演示序列(1401 nt)
设备:     npu:0
模型加载完成,耗时 6.9s,参数量 17.074M

推理耗时: 275.2 ms
输出分数矩阵: (1401, 18)  (行=碱基位置, 列=['no_splice', 'acceptor', 'donor'] + 15 组织)
NaN 检查: False
候选剪接位点数(threshold=0.5): 3
最高 acceptor 概率: 0.8913 @ pos 617
最高 donor 概率:    0.8997 @ pos 486

候选剪接位点 Top-K:
  pos   486  donor     score=0.8997  上下文: ...TGGGCAGGTTGGT...
  pos   617  acceptor  score=0.8913  上下文: ...CCTTAGGCTGCTG...
  pos   377  acceptor  score=0.6104  上下文: ...CACTAGCAACCTC...
SUCCESS

预测结果与人类 beta-globin(HBB,NCBI NG_000007.3)真实剪接位点吻合:

  • donor @ pos 486 = exon1/intron1 供体位点(GT),score 0.8997(最高)
  • acceptor @ pos 617 = intron1/exon2 受体位点(AG),score 0.8913(最高)
  • pos 377 为内含子内潜在受体位点(AG),score 0.6104

4.2 HTTP 服务化推理

$ curl -s -X POST http://127.0.0.1:8000/v1/predict -H "Content-Type: application/json" -d '{}'
{
  "length": 1401, "context": 4000, "max_seq_len": 8192,
  "top_acceptor": {"position": 617, "nucleotide": "G", "score": 0.8913},
  "top_donor":    {"position": 486, "nucleotide": "G", "score": 0.8997},
  "splice_sites": [{"position": 486, "type": "donor", "score": 0.8997}, ...],
  "inference_ms": 276.18
}

4.3 剪接改变(variant effect)推理

以完整 HBB pre-mRNA(1401 nt)为参考序列,做两类致病剪接突变验证:

(1)exon1/intron1 供体位点 G→A 点突变(pos 486,破坏规范 GT)

POST /v1/variant-effect
  {"reference": "<HBB 1401nt>", "alternative": "<HBB 1401nt, pos486 G->A>"}
  pos   486  G->A  donor     delta=-0.5438  loss     (供体位点使用显著下降)
  pos   470  G->G  donor     delta=+0.1182  gain
  pos   448  G->G  donor     delta=+0.0900  gain
  inference_ms: 70.9

(2)intron1/exon2 受体位点 AG→AC(pos 616 G→C,破坏受体 AG 二核苷酸)

POST /v1/variant-effect
  pos   617  G->G  acceptor  delta=-0.8868  loss     (受体位点使用几乎归零)
  pos   377  C->C  acceptor  delta=-0.2562  loss
  pos   486  G->G  donor     delta=-0.2494  loss
  inference_ms: 68.4

5. 快速开始(复现步骤)

5.1 获取权重

# 方式一:git clone 镜像(无 git-lfs 时走 LFS batch API,见下)
git clone https://gitcode.com/hf_mirrors/multimolecule/sptransformer.git /opt/atomgit/models/sptransformer

# 方式二:LFS batch API 直取 model.safetensors(无 git-lfs 环境)
#   oid=cc7e5f725c3a941b68427eb4f87ffb790a38bb901f80d6a26d3a92f1bf3c8f1a
#   curl -X POST https://gitcode.com/hf_mirrors/multimolecule/sptransformer.git/info/lfs/objects/batch \
#     -d '{"operation":"download","objects":[{"oid":"cc7e...8f1a","size":68412160}]}'   # 取签名 URL 后下载
# 校验:
sha256sum /opt/atomgit/models/sptransformer/model.safetensors
# cc7e5f725c3a941b68427eb4f87ffb790a38bb901f80d6a26d3a92f1bf3c8f1a

5.2 安装依赖

python3 -m venv .venv && source .venv/bin/activate
pip install -r requirements.txt
# 若需 vllm/vllm-ascend 请用独立 venv(本模型为卷积 + 稀疏注意力 Transformer,
# 无需 vLLM,与 transformers>=5 不冲突)

5.3 命令行推理

python3 inference.py \
    --model-path /opt/atomgit/models/sptransformer \
    --device npu:0

# 指定序列
python3 inference.py \
    --model-path /opt/atomgit/models/sptransformer \
    --sequence "ATGGTGCATCTGACTCCTGAGGAGAA..." --device npu:0

# 附加 15 组织使用度输出
python3 inference.py \
    --model-path /opt/atomgit/models/sptransformer \
    --include-tissue --device npu:0

# 剪接改变(variant effect):参考与替代序列必须同长。
# 示例:以 HBB 全长(1401nt,inference.py 内 DEFAULT_SEQUENCE)为参考,
# 将 pos486 供体位点 G 改为 A(破坏规范 GT)作为替代序列。
python3 inference.py \
    --model-path /opt/atomgit/models/sptransformer \
    --reference-file ./ref.fa \
    --alternative-file ./alt.fa --device npu:0

说明:--reference/--alternative 也可直接传序列字符串;二者长度必须相等 (仅支持同长变异,如 SNP/点突变)。输入长度超过 max_seq_len - 2*context (192 nt)时,窗口两端位点侧翼上下文被中心裁剪,建议使用包含完整剪接位点 上下文的序列(如 HBB 全长)或逐窗口推理。

5.4 服务化推理(FastAPI)

python3 inference.py --serve \
    --model-path /opt/atomgit/models/sptransformer \
    --host 0.0.0.0 --port 8000 --device npu:0

5.5 调用示例

# 健康检查
curl -s http://127.0.0.1:8000/health
# {"status":"ok","model":"multimolecule/sptransformer","device":"npu:0",
#  "npu":{"available":true,"device_count":2,"name":"Ascend910_9362"}}

# 模型信息
curl -s http://127.0.0.1:8000/v1/models

# 剪接位点预测(sequence 缺省用 HBB 演示序列)
curl -s -X POST http://127.0.0.1:8000/v1/predict \
  -H "Content-Type: application/json" \
  -d '{"sequence":"ATGGTGCATCTGACTCCTGAGGAGAA...","threshold":0.5,"top_k":20,"include_tissue":true}'

# 剪接改变预测(参考/替代必须同长;示例:HBB 全长 pos486 donor G->A,见 4.3 节)
curl -s -X POST http://127.0.0.1:8000/v1/variant-effect \
  -H "Content-Type: application/json" \
  -d '{"reference":"<HBB 1401nt 序列>","alternative":"<HBB 1401nt 序列,pos486 位 G 改为 A>"}'

6. 服务化推理 API

端点方法说明
/healthGET健康检查(返回 NPU 设备信息与模型状态)
/v1/modelsGET模型信息(模型 ID、任务、参数量、上下文、通道数)
/v1/predictPOST剪接位点预测;请求体 {"sequence": "...", "threshold": 0.5, "top_k": 20, "include_scores": true, "include_tissue": true}
/v1/variant-effectPOST剪接改变预测;请求体 {"reference": "...", "alternative": "..."}(同长)

/v1/predict 返回:逐位置 no_splice/acceptor/donor softmax 概率与 15 组织使用度、 候选剪接位点(acceptor/donor 概率 ≥ threshold)、Top acceptor/donor 位点; include_scores=true 时附带完整逐位置分数表。

7. NPU 适配说明

  • SpTransformer 为膨胀残差卷积 + 混合稀疏注意力 Transformer(非自回归 LLM), vLLM / vllm-ascend 无法服务;采用 multimolecule + torch_npu + FastAPI 方案在 昇腾 NPU 上推理。
  • 模型通过 multimolecule.SpTransformerModel.from_pretrained 加载(config architectures: SpTransformerModel),计算在昇腾 NPU(torch_npu 注册的 npu 后端)上进行。
  • one-hot 编码 F.one_hot 在 NPU 上仅触发 "internal format" 告警(输出正常, 无需修补),与 Enformer / DeltaSplice 部署一致。
  • 输入支持 DNA/RNA 大小写,RnaTokenizer 自动将 T 转换为 U、未知碱基映射为 N (N 在 one-hot 中为全零向量,等效掩码);tokenize 时 add_special_tokens=False。
  • model.postprocess(outputs) 返回 (scores, channels):前 3 通道 (no_splice/acceptor/donor)为 softmax 归一化概率,后 15 通道为组织使用度原生 logits 尺度(越大越易在该组织发生剪接)。
  • 服务化用 FastAPI(非 OpenAI 格式),inference.py 支持 CLI + --serve 双模式。

8. 目录结构

sptransformer-npu/
├── inference.py        # 推理脚本(CLI + FastAPI 服务化)
├── README.md           # 部署说明文档
├── requirements.txt    # 环境依赖清单
└── assets/             # 截图素材(agent_workflow / npu_device_call / model_result)

9. 精度与性能

场景输入长度耗时
CLI 剪接位点推理(首次加载 6.9s)1401 nt275.2 ms
HTTP /v1/predict1401 nt276.2 ms
HTTP /v1/variant-effect(首轮算子编译后)1401 nt(同长)70.9 ms
  • 模型权重 model.safetensors(68,412,160 B,sha256 cc7e5f72...bf3c8f1a)来自 gitcode 镜像 hf_mirrors/multimolecule/sptransformer,与 HF 官方权重一致。
  • 预测结果与 HBB 真实剪接位点吻合(donor@486=0.8997 / acceptor@617=0.8913); 破坏规范 GT/AG 的点突变均被正确预测为剪接使用下降。