作者:昇腾实战派

知识地图https://blog.csdn.net/Lumos_Lovegood/article/details/161601003

背景概述

单细胞 RNA 测序(scRNA-seq)技术的爆发式发展,使人类细胞图谱等综合数据库的规模已膨胀至数千万个细胞。如何从海量且高度异质的单细胞数据中高效提取生物学知识——包括细胞类型识别、批次效应消除、基因扰动预测和调控网络推断——成为计算生物学的核心挑战。当前主流方法大多针对特定任务开发专用模型,导致"数据孤岛"和"任务孤岛"的局限。

scGPT(single-cell Generative Pre-trained Transformer)由多伦多大学 Bo Wang 实验室于 2024 年发表于 Nature Methods,是一种专为单细胞多组学设计的基础模型。scGPT 建立了语言与细胞生物学之间的平行联系——正如文本由单词组成,细胞由基因定义。模型将基因视为 Token,将细胞视为文本,在 CELLxGENE 中超过 3300 万个人类正常单细胞的 RNA 测序数据上完成生成式预训练,通过迁移学习支持多种下游任务,在相应基准上达到 SOTA。

本文介绍 scGPT 的昇腾 Ascend NPU 适配版本——对微调与推理流程进行 NPU 适配,使用昇腾原生 FlashAttention2 加速注意力计算,为大规模单细胞数据分析提供高效的硬件加速方案。

模型介绍

scGPT 概述

scGPT 的核心设计理念是将自然语言处理中的生成式预训练范式引入单细胞组学领域。模型在 51 个器官/组织、超过 3300 万个正常人类细胞上完成预训练,学习到基因和细胞的通用嵌入表示。

支持的下游任务:

  • 细胞类型注释:对未知细胞进行分类标注
  • 多批次整合:消除技术批次效应,保留生物学变异
  • 多组学整合:融合 RNA-seq 和 ATAC-seq 等多模态数据
  • 扰动响应预测:预测基因编辑后的细胞状态变化
  • 基因网络推断:从注意力权重提取基因调控网络

整体架构

scGPT 基于 Transformer 架构,但针对单细胞数据的非序列化特征进行了专门设计。

输入编码

单细胞数据的输入整合了三个维度的信息:

编码类型 说明 维度
Gene Token Embedding 基因名称的离散编码 ntoken × d_model
Expression Value Encoding 基因表达值的连续/离散编码 d_model
Condition Token 批次、模态、扰动条件等元信息 d_model

表达值编码支持三种模式:

  • continuous:连续值编码器(MLP 映射)
  • category:分箱离散化后的嵌入编码
  • scaling:归一化缩放

核心 Transformer

  • 堆叠的 Transformer Encoder 层(默认 12 层)
  • 多头注意力(默认 8 头),支持 Flash Attention 加速
  • 前馈网络隐藏维度 d_hid
  • 支持 Pre-Norm 和 Post-Norm 两种归一化方案

掩蔽注意力机制(Masked Attention)

由于基因表达数据不具有自然语言的严格语序,scGPT 设计了专门的注意力掩码策略:

  • 已知基因可互相关注(双向注意力)
  • 被掩蔽基因只能关注已知基因(单向约束)
  • 这使得模型在生成式预训练中能同时学习基因间的相互依赖关系

多任务输出头

输出头 功能 场景
ExprDecoder 基因表达值重建 预训练、批次整合
ClsDecoder 细胞分类 细胞类型注释
MVCDecoder 掩码值预测 自监督学习
AdversarialDiscriminator 对抗批次判别 批次整合

细胞嵌入策略

  • cls:使用 [CLS] token 的输出作为细胞级表示
  • avg-pool:对所有基因 token 的输出取平均
  • w-pool:加权池化

预训练与微调

预训练

  • 数据:CELLxGENE 3300 万人类正常单细胞(51 个器官/组织)
  • 任务:生成式预训练——随机掩蔽部分基因表达值,模型预测被掩蔽的值
  • 损失函数:掩蔽 MSE 损失 + 弹性细胞相似度(ECS)正则化

微调

  • 冻结预训练层,仅训练任务特定的输出头
  • 逐步解冻全部层进行端到端微调
  • 支持 Domain-Specific BatchNorm(DSBN)处理多批次数据

缩放定律

研究发现 scGPT 的性能与预训练数据规模呈正相关——当预训练数据从 3 万扩展到 3300 万个细胞时,下游任务表现持续提升,与 NLP 领域的缩放定律高度一致。

昇腾 NPU 适配版本

迁移动机

原始 scGPT 依赖 CUDA Flash Attention 进行高效注意力计算。为在昇腾 Ascend NPU 上运行,本项目使用 torch_npu 原生的 npu_fusion_attention 替代 CUDA Flash Attention,实现等效的加速效果。

仓库结构

scGPT_npu/
├── README.md
├── requirements.txt
├── pyproject.toml
├── examples/
│   ├── finetune_integration.py     # 批次整合微调脚本(NPU 适配)
│   ├── inference.py                # 推理脚本(NPU 适配)
│   ├── save/
│   │   └── scGPT_human/           # 预训练权重
│   │       ├── best_model.pt
│   │       ├── vocab.json
│   │       └── args.json
│   └── data/                       # 示例数据
├── scgpt/
│   ├── model/
│   │   ├── model.py                # TransformerModel 主模型
│   │   ├── dsbn.py                 # Domain-Specific BatchNorm
│   │   ├── generation_model.py     # 生成式模型
│   │   └── multiomic_model.py      # 多组学模型
│   ├── utils/
│   │   └── flash_attention.py      # ⭐ NPU Flash Attention 实现
│   ├── tasks/
│   │   ├── cell_emb.py             # 细胞嵌入提取
│   │   └── grn.py                  # 基因调控网络推断
│   ├── tokenizer/
│   │   └── gene_tokenizer.py       # 基因词表编码
│   ├── preprocess.py               # 数据预处理
│   ├── loss.py                     # 损失函数
│   ├── trainer.py                  # 训练器
│   └── data_collator.py            # 数据整理器
├── data/
│   ├── pbmc3k.h5ad                 # 示例数据
│   └── cellxgene/                  # 大规模数据构建脚本
└── tests/

核心迁移改动

1. Flash Attention NPU 实现

scgpt/utils/flash_attention.py 使用昇腾原生的 npu_fusion_attention 替代 CUDA flash-attn:

from torch_npu import npu_fusion_attention

支持特性:

  • FlashAttention2:下右对齐因果掩码(默认)或左上对齐因果掩码
  • 环境变量控制NPU_FA2_SPARSE_MODE=2(左上对齐)或 3(下右对齐)
  • 注意力掩码缓存:复用预计算的掩码减少内存分配
2. 推理与微调入口适配

所有入口脚本添加 NPU 自动迁移:

import torch_npu
from torch_npu.contrib import transfer_to_npu
3. NPU 性能优化配置
export CPU_AFFINITY_CONF=1        # CPU 亲和性绑定
export TASK_QUEUE_ENABLE=2        # 任务队列优化
export ASCEND_RT_VISIBLE_DEVICES=0  # 指定 NPU 卡

版本信息

软件 版本
HDK 25.5.0
CANN 8.3.RC1
Python 3.11
PyTorch 2.1.0
torch_npu 2.1.0

环境配置

创建 Conda 环境

conda create -n scgpt python=3.11 -y
conda activate scgpt

克隆代码

git lfs install
git clone https://atomgit.com/AI4Science/scGPT.git
cd scGPT

安装依赖

export PIP_INDEX_URL=https://repo.huaweicloud.com/repository/pypi/simple

cd scGPT && pip install -e .
pip install -r requirements.txt
pip install torchtext==0.15.2 torchdata==0.7.1 --no-deps

主要依赖包括:

  • torch==2.1.0torch_npu==2.1.0
  • scanpy==1.10.3:单细胞数据分析核心库
  • scvi-tools==1.2.1:单细胞变分推断工具
  • scib==1.1.5:单细胞整合基准评估
  • anndata==0.10.8:单细胞数据结构
  • einops==0.8.1:张量操作工具

验证安装

python3 -c "import torch;import torch_npu; a = torch.randn(3, 4).npu(); print(a + a);"

NPU 运行环境配置

export ASCEND_RT_VISIBLE_DEVICES=0
export CPU_AFFINITY_CONF=1
export TASK_QUEUE_ENABLE=2
export WANDB_MODE=offline
export WANDB_OFFLINE=true

模型权重

预训练权重已包含在仓库中(通过 Git LFS):

examples/save/scGPT_human/
├── best_model.pt      # 预训练模型权重
├── vocab.json         # 基因词表(约 60,000 个基因)
└── args.json          # 模型超参数配置

模型参数配置(来自 args.json):

  • d_model: 512(隐藏维度)
  • nhead: 8(注意力头数)
  • nlayers: 12(Transformer 层数)
  • d_hid: 512(前馈网络维度)

微调

批次整合微调

cd examples/
python finetune_integration.py

该脚本完成:

  1. 加载 PBMC 10K 数据集(含多批次)
  2. 使用 scGPT_human 预训练权重初始化
  3. 以 40% mask ratio 进行掩蔽表达值预测
  4. 结合对抗批次判别器(DAB)消除批次效应
  5. 使用 scib 评估整合质量(AvgBIO 等指标)

主要超参数:

  • mask_ratio: 0.4
  • epochs: 30
  • n_bins: 51(表达值分箱数)
  • learning_rate: 1e-4
  • batch_size: 64

推理

细胞嵌入提取

cd examples/
python inference.py

推理流程:

  1. 加载 Kim2020 肺部单细胞数据(h5ad 格式)
  2. 筛选高变异基因(Top 3000 HVG)
  3. 使用预训练模型提取细胞嵌入
  4. 输出推理耗时统计
embed_adata = scg.tasks.embed_data(
    adata,
    model_dir,
    gene_col=gene_col,
    batch_size=64,
)

Docker 镜像

已预置 Conda 运行环境、模型权重、scGPT 源码,开箱即用:

sudo docker pull swr.cn-north-4.myhuaweicloud.com/ascend_ai4s/scgpt:v1

镜像内容:

  • 项目源码路径:/root/scGPT/
  • 已配置好 Conda 虚拟环境
  • 已内置模型权重文件,无需额外下载

迁移适配要点

Flash Attention 替换

原实现 NPU 实现
CUDA flash_attn torch_npu.npu_fusion_attention
flash_attn_func(q, k, v) npu_fusion_attention(q, k, v, head_num, ...)
CUDA causal mask NPU sparse_mode 控制(2=左上,3=下右)

自动设备迁移

import torch_npu
from torch_npu.contrib import transfer_to_npu

transfer_to_npu 自动将所有 CUDA 调用重定向到 NPU,无需手动修改设备代码。

性能优化

  • CPU_AFFINITY_CONF=1:启用 CPU 亲和性绑定,减少跨 NUMA 访问
  • TASK_QUEUE_ENABLE=2:启用任务队列优化,提升 NPU 利用率
  • Flash Attention 注意力掩码缓存:避免重复创建大尺寸掩码张量

已知限制

  • 当前版本聚焦于微调与推理,不包含完整的 3300 万细胞预训练流程
  • Flash Attention 的 npu_fusion_attention 要求输入为 float16/bfloat16
  • 大规模数据集(>100 万细胞)的微调可能需要多卡并行
  • scvi-tools 的部分功能仍在 CPU 上运行

参考文献

  • Cui, H., Wang, C., Maan, H., Pang, K., Luo, F., Duan, N., Wang, B. scGPT: toward building a foundation model for single-cell multi-omics using generative AI. Nature Methods, 21(8), 1470-1480 (2024). https://www.nature.com/articles/s41592-024-02201-0
  • 上游代码仓库:https://github.com/bowang-lab/scGPT
  • 昇腾 NPU 适配版:https://atomgit.com/AI4Science/scGPT
Logo

作为“人工智能6S店”的官方数字引擎,为AI开发者与企业提供一个覆盖软硬件全栈、一站式门户。

更多推荐