作者:昇腾实战派

知识地图Geneformer 官方 Hugging Face 仓库Geneformer 官方文档


一、背景概述

单细胞转录组测序技术的快速发展,让人类能够以前所未有的分辨率观测基因表达状态。然而,基因调控网络的复杂性、疾病样本稀缺、以及批次效应等问题,使得传统的监督学习方法在许多生物学任务上难以取得稳定的效果。迁移学习和大规模预训练模型为这些问题提供了新的思路:先在海量无标注数据上学习通用的“基因语言”,再在小样本下游任务上进行微调。

Geneformer 正是这样一款面向单细胞转录组的基础 Transformer 模型。它由 Christina Theodoris 等人提出,先后在 Genecorpus-30M(约 3000 万人类单细胞转录组)和 Genecorpus-104M(约 1.04 亿)上进行自监督预训练,能够将每个细胞的转录组表示为基因排序序列,并在此基础上学习基因之间的上下文关系与网络层级。相关工作已发表于 Nature,并在疾病靶点发现、虚拟基因扰动、细胞状态分类等任务中得到了实验验证。

本文档基于 Geneformer 官方 README 与当前昇腾 NPU 适配仓库,系统介绍 Geneformer 的模型原理、昇腾环境准备、权重/数据集获取、快速推理以及常见微调任务(细胞分类、基因分类、细胞嵌入可视化、多任务分类、虚拟扰动)的完整操作流程,力求为生物医学 AI 在昇腾平台上的落地提供一份可复现的实操指南。


二、模型介绍

2.1 模型定位

Geneformer 是一个 Encoder-only 的 Transformer 基础模型,专门针对单细胞转录组数据设计。与 NLP 中的 BERT 类似,Geneformer 在海量无标注转录组数据上使用 掩码语言建模(Masked Language Modeling, MLM) 目标进行预训练:在每个细胞中随机掩码 15% 的基因,模型需要根据剩余基因的上下文预测被掩码位置的基因。通过这种方式,模型在完全自监督的情况下学习到了基因之间的共现关系和网络层级。

预训练完成后,Geneformer 可以:

  • 零样本学习(Zero-shot):直接用于虚拟基因扰动、批次整合、基因上下文特异性分析等任务;
  • 微调(Fine-tuning):在少量标注数据上针对细胞类型注释、疾病分类、转录因子剂量敏感性、染色质动力学等下游任务进行微调。

2.2 核心机制:Rank Value Encoding

Geneformer 的输入不是原始表达量,而是一种 Rank Value Encoding(排序值编码)

  1. 对每个基因,计算其在整个预训练语料(Genecorpus)中的非零中位表达量;
  2. 将单个细胞中每个基因的计数归一化到该中位值;
  3. 按归一化后的值对基因进行排序,得到该细胞的基因序列;
  4. 将基因名映射为 Token ID,从而把每个细胞转成一条“基因句子”。

这种编码方式有两个显著优势:

  • 弱化管家基因(housekeeping genes): ubiquitously 高表达的基因会被排序到较低位置,从而减少其对模型的干扰;
  • 突出状态标志基因:转录因子等低表达但高区分度的基因会排在前面,帮助模型捕捉细胞状态。

2.3 模型规格

当前仓库中主要包含以下预训练模型(以本地实际目录为准):

模型目录 训练数据 参数量 输入长度 词表大小 说明
gf-6L-30M-i2048 Genecorpus-30M 30M 2048 ~25K 30M 系列 6 层模型,本文主要示例
gf-12L-30M-i2048 Genecorpus-30M 30M 2048 ~25K 30M 系列 12 层模型
gf-12L-95M-i4096 Genecorpus-95M 95M 4096 ~25K 95M 系列 12 层模型
gf-20L-95M-i4096 Genecorpus-95M 95M 4096 ~25K 95M 系列 20 层模型
Geneformer-V2-104M Genecorpus-104M 104M 4096 ~20K V2 更新版本
Geneformer-V2-316M Genecorpus-104M 316M 4096 ~20K V2 默认大模型
Geneformer-V2-104M_CLcancer 癌症数据持续学习 104M 4096 ~20K 癌症领域微调模型

注意:30M 系列模型需要配合 special_token=Falsemodel_input_size=2048 使用;95M/V2 系列则使用默认的 special_token=Truemodel_input_size=4096。具体见 examples/tokenizing_scRNAseq_data.ipynb

2.4 支持的下游任务

Geneformer 已验证的下游任务包括但不限于:

  • 细胞分类:疾病状态、细胞类型、分化阶段等;
  • 基因分类:转录因子剂量敏感性、染色质状态、调控范围等;
  • 细胞嵌入提取与可视化:UMAP、热图等;
  • 多任务细胞分类:同时预测细胞类型与疾病状态;
  • 虚拟基因扰动(in silico perturbation):预测删除/过表达某基因后细胞状态的变化;
  • 虚拟治疗(in silico treatment):筛选潜在治疗靶点。

三、版本信息

本次实践使用的软硬件环境如下(请根据自身机器型号选择合适的 CANN 镜像):

软件/硬件 版本
HDK 25.3.rc1.2
CANN 8.3.RC1
Python 3.10
PyTorch 2.6.0
torch-npu 2.6.0
transformers 4.37.2(推荐)

提示:Atlas 800I A3 机器请下载带 a3 关键字的 CANN 镜像;Atlas 800I A2 机器请下载带 910b 关键字的版本。


四、环境准备

4.1 创建 Conda 环境

conda create --name Genecorpus python=3.10 -y
conda activate Genecorpus

4.2 克隆并安装 Geneformer

# 安装 git-lfs(如未安装)
git lfs install

# 克隆本仓库(昇腾 NPU 适配版)
git clone https://atomgit.com/AI4Science/Geneformer.git
cd Geneformer

# 配置华为云 PyPI 镜像加速下载
export PIP_INDEX_URL=https://repo.huaweicloud.com/repository/pypi/simple/

# 安装 PyTorch 与 torch_npu
pip install torch==2.6.0 torch_npu==2.6.0

# 安装 Geneformer 及其依赖
pip install .

4.3 验证 NPU 环境

执行以下命令检查 PyTorch 与 torch_npu 是否安装成功:

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

若输出类似以下内容,说明 NPU 可用:

tensor([[-0.6066,  6.3385,  0.0379,  3.3356],
        [ 2.9243,  3.3134, -1.5465,  0.1916],
        [-2.1807,  0.2008, -1.1431,  2.1523]], device='npu:0')

如果报错,可按以下顺序排查:

  1. 是否已 source /usr/local/Ascend/ascend-toolkit/set_env.sh
  2. decorator 等 torch_npu 运行时依赖是否安装;
  3. CANN、HDK、torch、torch_npu 版本是否匹配;
  4. 使用 npu-smi info 查看驱动与设备是否正常。

4.4 昇腾迁移通用代码

在需要调用 NPU 的 Python 脚本头部,添加以下两行以启用自动迁移适配:

import torch_npu
from torch_npu.contrib import transfer_to_npu

提示:本仓库的示例脚本(如 cell_classification.pymultitask_cell_classification.py)已默认包含上述导入。


五、权重与数据集准备

5.1 下载 Genecorpus-30M 示例数据

git clone https://gitee.com/hf-datasets/Genecorpus-30M.git

克隆完成后,示例输入文件位于:

Genecorpus-30M/example_input_files/
├── cell_classification/disease_classification/human_dcm_hcm_nf.dataset
├── gene_classification/dosage_sensitive_tfs/
│   ├── dosage_sensitivity_TFs.pickle
│   └── gc-30M_sample50k.dataset
└── token_dictionary.pkl

5.2 获取预训练权重

本仓库已内置部分模型的配置文件,权重文件可通过以下方式获取:

  • 方式一(推荐):运行仓库自带的下载脚本

    python3 download_model.py
    

    该脚本会自动从 Hugging Face 或镜像站下载 gf-6L-30M-i2048model.safetensorstraining_args.bin

  • 方式二:手动从 Hugging Face 拉取

    # 以 gf-6L-30M-i2048 为例
    huggingface-cli download ctheodoris/Geneformer \
        gf-6L-30M-i2048/model.safetensors \
        gf-6L-30M-i2048/training_args.bin \
        --local-dir /root/Geneformer
    
  • 方式三:使用 hf-mirror.com 等国内镜像

    export HF_ENDPOINT=https://hf-mirror.com
    huggingface-cli download ctheodoris/Geneformer \
        gf-6L-30M-i2048/model.safetensors \
        gf-6L-30M-i2048/training_args.bin \
        --local-dir /root/Geneformer
    

路径说明:下文示例默认模型路径为 /root/Geneformer/gf-6L-30M-i2048,数据路径为 /root/Geneformer/Genecorpus-30M/...,请根据实际目录进行修改。


六、数据预处理:将单细胞数据转为 Rank Value Encoding

Geneformer 要求输入为 tokenized 的 .dataset 格式(Hugging Face Datasets 结构)。仓库提供了 TranscriptomeTokenizer 用于将 .loom.h5ad 原始计数数据转为 rank value encoding:

from geneformer import TranscriptomeTokenizer

# 保留细胞类型与组织信息作为下游标签
 tk = TranscriptomeTokenizer(
    {"cell_type": "cell_type", "organ_major": "organ"},
    nproc=16,
)

# file_format 支持 "loom" 或 "h5ad"
tk.tokenize_data(
    "loom_data_directory",   # 原始数据目录
    "output_directory",      # 输出目录
    "output_prefix",         # 输出前缀
    file_format="loom",
)

注意:30M 系列模型需设置 special_token=Falsemodel_input_size=2048;95M/V2 系列使用默认值即可。具体参数请参见 examples/tokenizing_scRNAseq_data.ipynb


七、快速体验:加载模型并提取细胞嵌入

完成环境准备后,可以通过 EmbExtractor 快速提取细胞嵌入:

import torch_npu
from torch_npu.contrib import transfer_to_npu
from geneformer import EmbExtractor

embex = EmbExtractor(
    model_type="CellClassifier",
    num_classes=3,
    filter_data={"cell_type": ["Cardiomyocyte1", "Cardiomyocyte2", "Cardiomyocyte3"]},
    max_ncells=1000,
    emb_layer=0,
    emb_mode="cell",          # 30M 词表无 <cls> token,使用 mean-pooled 细胞嵌入
    emb_label=["disease", "cell_type"],
    labels_to_plot=["disease"],
    forward_batch_size=200,
    nproc=16,
    token_dictionary_file="/root/Geneformer/Genecorpus-30M/token_dictionary.pkl",
)

embs = embex.extract_embs(
    "/root/Geneformer/gf-6L-30M-i2048",  # 预训练模型目录
    "/root/Geneformer/Genecorpus-30M/example_input_files/cell_classification/disease_classification/human_dcm_hcm_nf.dataset",
    "/root/Geneformer/examples/output/cell_embeddings",
    "cm_emb",
)

# 绘制 UMAP
embex.plot_embs(
    embs=embs,
    plot_style="umap",
    output_directory="/root/Geneformer/examples/output/cell_embeddings",
    output_prefix="emb_plot",
)

运行示例脚本:

cd /root/Geneformer/examples
python3 extract_and_plot_cell_embeddings.py

八、微调实战

进入示例目录并配置昇腾环境:

cd /root/Geneformer/examples
source /usr/local/Ascend/ascend-toolkit/set_env.sh

8.1 细胞分类

本示例使用人心肌病数据(human_dcm_hcm_nf.dataset),将心肌细胞按疾病状态(nf/hcm/dcm)进行分类。

8.1.1 NPU 适配补丁

如果使用的是未修改的 Geneformer 源码,需要在 evaluation_utils.py 中将 cuda 改为 npu。当前仓库已默认完成以下修改:

# geneformer/evaluation_utils.py
import torch_npu
from torch_npu.contrib import transfer_to_npu

device = torch.device('npu' if torch_npu.npu.is_available() else 'cpu')

# 将 input_ids / attention_mask / labels 的 .to("cuda") 改为 .to(device)
input_ids=torch.tensor(np.array(input_data_batch, dtype=np.int64)).to(device),
attention_mask=torch.tensor(np.array(attn_msk_batch, dtype=np.int64)).to(device),
labels=torch.tensor(np.array(label_batch, dtype=np.int64)).to(device),

提示:当前 cell_classification.py 已包含 import torch_nputransfer_to_npu

8.1.2 运行训练
python3 cell_classification.py

等价地,也可使用更简洁的封装脚本:

python3 run_cell_classification.py

脚本会自动完成数据准备、交叉验证、测试评估、混淆矩阵与预测结果可视化,并输出到 examples/output/ 目录。

8.2 基因分类

本示例根据转录因子剂量敏感性对基因进行分类。

python3 gene_classification.py

该脚本包含两部分:

  1. 使用 5 折交叉验证训练并评估基因分类器;
  2. 使用全部数据训练最终模型并保存。

8.3 多任务细胞分类

多任务分类器可同时预测多个细胞属性(如细胞类型 + 疾病状态)。本仓库提供了数据准备脚本与训练脚本:

# 1. 准备训练/验证/测试数据
python3 prepare_mtl_data.py

# 2. 运行多任务训练
python3 multitask_cell_classification.py

multitask_cell_classification.py 默认使用 Optuna 进行超参搜索(n_trials=2 仅用于快速跑通,生产环境建议 ≥50),然后加载最优模型在测试集上评估。

注意:脚本中注释掉的虚拟扰动部分依赖 bitsandbytes 8-bit 量化,当前在昇腾 NPU 上不支持,请在 CUDA 环境下运行。

8.4 虚拟基因扰动(可选)

Geneformer 支持零样本虚拟扰动:删除或过表达某个基因,观察细胞嵌入向目标状态的偏移,从而发现疾病驱动基因或治疗靶点。

from geneformer import InSilicoPerturber, InSilicoPerturberStats, EmbExtractor

# 1. 计算起始状态、目标状态和替代状态的嵌入
embex = EmbExtractor(
    model_type="CellClassifier",
    num_classes=3,
    filter_data={"cell_type": ["Cardiomyocyte1", "Cardiomyocyte2", "Cardiomyocyte3"]},
    max_ncells=1000,
    emb_layer=0,
    summary_stat="exact_mean",
    forward_batch_size=256,
    nproc=16,
)

state_embs_dict = embex.get_state_embs(
    cell_states_to_model={
        "state_key": "disease",
        "start_state": "dcm",
        "goal_state": "nf",
        "alt_states": ["hcm"],
    },
    model_directory="/path/to/fine_tuned_CellClassifier",
    input_data_file="/path/to/input_data",
    output_directory="/path/to/output_directory",
    output_prefix="isp_state_embs",
)

# 2. 执行扰动
isp = InSilicoPerturber(
    perturb_type="delete",
    genes_to_perturb="all",
    model_type="CellClassifier",
    num_classes=3,
    emb_mode="cell",
    cell_emb_style="mean_pool",
    cell_states_to_model={...},
    state_embs_dict=state_embs_dict,
    max_ncells=2000,
    emb_layer=0,
    forward_batch_size=400,
    nproc=16,
)

isp.perturb_data(
    model_directory="/path/to/fine_tuned_CellClassifier",
    input_data_file="/path/to/input_data",
    output_directory="/path/to/isp_output_directory",
    output_prefix="isp_result",
)

# 3. 统计显著性
ispstats = InSilicoPerturberStats(
    mode="goal_state_shift",
    genes_perturbed="all",
)
ispstats.get_stats(
    "/path/to/isp_output_directory",
    None,
    "/path/to/isp_stats_output_directory",
    "isp_result",
)

完整示例见 examples/in_silico_perturbation.ipynb


九、Docker 镜像(开箱即用)

为简化部署,本仓库提供了预置的 Docker 镜像,已包含 Conda 环境、模型权重与 Geneformer 源码:

# 拉取镜像
sudo docker pull swr.cn-north-4.myhuaweicloud.com/ascend_ai4s/geneformer:a3

# 启动容器(根据实际 Ascend 设备挂载对应路径)
sudo docker run --privileged -it -u root --ipc=host --network=host \
    --device=/dev/davinci_manager \
    --device=/dev/devmm_svm \
    --device=/dev/hisi_hdc \
    -v /usr/local/Ascend/driver:/usr/local/Ascend/driver \
    -v /usr/local/Ascend/ascend-toolkit:/usr/local/Ascend/ascend-toolkit \
    -v /root/Geneformer:/root/Geneformer \
    swr.cn-north-4.myhuaweicloud.com/ascend_ai4s/geneformer:a3 /bin/bash

容器内:

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

十、常见问题与注意事项

  1. 30M 系列词表问题

    • 使用 30M 系列模型时,务必将 token_dictionary_file 指向 Genecorpus-30M/token_dictionary.pkl,否则 Classifier / EmbExtractor 会使用 V2 默认词表,导致维度不匹配。
  2. evaluation_strategy / eval_strategy 报错

    • 当前仓库的 classifier.py 已默认在 eval_data is None 时设置 eval_strategy="no";若使用其他分支或旧版本,可将相关 evaluation_strategy 统一改为 "no"
  3. 训练脚本找不到数据集或模型

    • 请检查 cell_classification.pygene_classification.py 等脚本中的绝对路径是否与本机一致,建议使用相对路径或环境变量管理。
  4. bitsandbytes 量化不支持 NPU

    • MTLCellClassifier-Quantized 与部分虚拟扰动示例依赖 CUDA 8-bit 量化,昇腾 NPU 暂不支持,请在 CUDA 环境中运行或改用非量化版本。
  5. transformers 版本冲突

    • 本仓库推荐 transformers==4.37.2。若遇到 AdamW 相关报错,可参考 examples/run_cell_classification.py 中的补丁方案。
  6. NPU 内存不足

    • 可适当降低 per_device_train_batch_sizeforward_batch_sizemax_ncells 等参数;或使用更小的 6 层 30M 模型。

附录:官方资源链接

  • Hugging Face 模型仓库:https://huggingface.co/ctheodoris/Geneformer
  • 数据集仓库(Genecorpus-30M):https://huggingface.co/datasets/ctheodoris/Genecorpus-30M
  • 官方文档:https://geneformer.readthedocs.io/en/latest/
  • 原始论文:https://rdcu.be/ddrx0
Logo

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

更多推荐