用Geneformer和Hugging Face微调实现单细胞转录组分类实践
2026/9/16 21:08:02 网站建设 项目流程

合作方上周丢给我一个需求:手上有病人和对照的单细胞转录组数据,要训练一个分类模型,把两组样本区分开。换作以前,我第一反应是找marker基因、抽PCA特征,然后上随机森林或者XGBoost就完事了。但这次我选了另一条路——用Geneformer配合Hugging Face Transformers,把每个细胞当“一句话”来做分类。

选这条路不是跟风。单细胞转录组矩阵动辄几万维,直接喂给传统分类器噪声很大,而且基因表达不是独立变量,基因与基因之间的调控关系在“分选细胞类型”这种任务里很有价值。Geneformer在约3000万个单细胞转录组上做过自监督预训练,已经学过基因调控的上下文信息,我们只需要把它当BERT一样微调,就能得到一个分类器。这篇文章把我从数据准备、模型加载、训练到踩坑的完整链路写清楚,适合已经会用Scanpy处理单细胞数据、但对Transformer还处于“能跑通但不太明白为什么”阶段的读者。

1. 单细胞分类的痛点与Geneformer的解题思路

1.1 传统方案为什么不够用

我最早做细胞类型注释或者疾病状态分类,走的是经典流程:Scanpy读数据、QC过滤、归一化、找高变基因、PCA降维、然后聚类。聚类出来以后用已知marker基因去注释每一群,最后统计不同组间的细胞比例差异。

这套流程放在简单问题上没问题,但一旦目标是“给单个细胞标注一个类别”——比如预测它来自疾病组还是对照组——传统做法会非常别扭。你可以在高变基因上训练一个逻辑回归或者SVM,也能在PCA embedding上跑XGBoost,但问题是:

  • 高变基因本身是统计筛选出来的,会丢掉大量低频但有判别力的基因;
  • PCA是线性变换,基因之间的非线性调控关系在里面表达不出来;
  • 不同批次、不同样本的count差异会直接影响这些模型的特征分布,一换数据就得重新调。

1.2 Geneformer的底子:BERT架构 + 单细胞预训练

Geneformer的思路其实很直接:把每个细胞的基因表达量按从高到低排序,形成一条“基因序列”,然后用一个类似BERT的Transformer模型在这个序列上做掩码预测。预训练数据是跨组织、跨疾病状态的约3000万个人类单细胞转录组,模型在这个语料上学到的是“哪些基因倾向于共同出现”“哪些基因的表达等级关系是怎样的”,本质上是在学基因调控的上下文。

模型的参数规模不大,不到1500万,配置大概是:

  • 6层Transformer Encoder
  • hidden size 256
  • 4个注意力头
  • 最大序列长度2048

这在Transformer模型里算非常轻量的。也正因为轻量,单卡微调是可行的。

1.3 什么情况下值得用Geneformer,什么情况不必

每个细胞对应一个类别标签,这类任务正是Geneformer下游微调的强项,尤其是:

  • 训练标注样本不多,比如只有几例病人和几例对照;
  • 类别之间的差异不是单个基因表达量高低,而是基因组合和调控网络层面的差异;
  • 需要模型迁移到新的数据集上继续微调。

如果只是对不同细胞类型做注释,且marker基因又很明确,那用传统方法可能更快。不要在简单问题上强行上Transformer,这是我一贯的原则。

2. 环境与依赖:把Hugging Face生态拉起来

2.1 基础安装

我建议先建一个干净的conda环境:

conda create -n geneformer python=3.10 -y conda activate geneformer pip install "transformers>=4.30" datasets accelerate huggingface_hub pip install scanpy anndata scikit-learn pandas numpy

单细胞下游分析还需要leidenalgharmony之类,看具体场景再装。核心就是transformers和datasets,这两个是整个微调链路的地基。

2.2 从Hugging Face Hub拉取Geneformer资源

Geneformer官方权重目前挂在Hugging Face Hub上,仓库ID是ctheodoris/Geneformer。里面有几个关键文件:

  • 模型权重文件(pytorch_model.bin)
  • gene token字典文件(token_dictionary.json),记录基因名到token id的映射
  • 模型配置说明

我用huggingface_hub直接拉文件到本地:

import os from huggingface_hub import hf_hub_download repo_id = "ctheodoris/Geneformer" local_dir = "./geneformer_model" os.makedirs(local_dir, exist_ok=True) hf_hub_download( repo_id=repo_id, filename="pytorch_model.bin", local_dir=local_dir, ) hf_hub_download( repo_id=repo_id, filename="token_dictionary.json", local_dir=local_dir, )

如果网络条件一般,可以设置HF_ENDPOINT走镜像站点,这里不展开,但你知道有这个方法即可。

顺便多说一句:Hugging Face Hub上很多模型仓库会同时在README里给用法示例,GitHub上作者也开源了训练和推理代码。但官方的训练代码为了兼顾科研场景,封装层次比较多,直接跑会有一堆参数要调。我在项目里习惯只借权重和字典,模型结构自己用Transformers拼,这样后面替换分类头、做交叉验证都更顺手。

2.3 Geneformer权重与标准BERT的兼容问题

Geneformer的骨干网络结构基本就是BERT,但它并不是直接用Hugging Face的BertModel存成model.safetensors的。它的state_dict key通常没有bert.前缀,或者层级命名和标准BERT略有不同。

我第一次加载时直接model.load_state_dict(state_dict),报了一堆unexpected key,后来检查才发现是key名对不上。

解决思路有两种:

  1. 把官方权重key做一次映射,比如把encoder.layer.0.attention.self.query.weight映射成bert.encoder.layer.0.attention.self.query.weight
  2. 先创建BertForSequenceClassification,再load_state_dict时用strict=False,加上key名替换逻辑。

我在项目里通常会先把官方权重加载成一个普通字典,然后做一层字符串替换,再扔给model.bert加载。代码如下:

import torch from transformers import BertForSequenceClassification, BertConfig state_dict = torch.load("./geneformer_model/pytorch_model.bin", map_location="cpu") # 示例:原key可能是 encoder.layer.0.attention.self.query.weight # 需要变成 bert.encoder.layer.0.attention.self.query.weight mapped_state_dict = {} for k, v in state_dict.items(): if not k.startswith("bert."): k = "bert." + k mapped_state_dict[k] = v mapped_state_dict = {k: v for k, v in mapped_state_dict.items() if k.startswith("bert.")}

这一步对应标准BERT的weight名,操作完以后load_state_dict基本能对上。

3. 数据处理核心:把count矩阵改造成token序列

3.1 rank value encoding到底是什么

Geneformer的输入不是PCA降维后的坐标,也不是直接的高维表达向量,而是每个细胞内部的“基因表达排名序列”,论文里叫rank value encoding。

过程拆开看:

  1. 取一个细胞,找到所有表达量大于0的基因;
  2. 按表达量从高到低排序;
  3. 把排序后的基因替换成词典里的token id;
  4. 截断到固定长度(默认2048)。

换句话说,表达量最高的基因在序列最前面,表达量低一点的排在后面,不表达的基因直接不进序列。每个细胞由此变成一条整数序列。

这就像给细胞写了一句“句子”:句子里的单词顺序不是自然语言的语法,而是基因的表达强度顺序。Transformer注意力机制会自动去学哪些位置和哪些基因组合有判别力。

3.2 基因ID一定要对准词典

这里我必须提一个重要细节:Geneformer的token_dictionary.json里的键是Ensembl基因ID,不是gene symbol,不是gene symbol。也就是说你不能直接把Scanpy里var_names那一列gene symbol丢进去查表,会查到一堆缺失。

正确做法是在上游把基因名映射到Ensembl ID。Scanpy里如果adata.var_names是symbol,可以用scanpy里自带的注释文件或者自己维护一个symbol到Ensembl ID的映射表来做。映射完成以后,再和token字典求交集,丢掉那些不在字典里的基因。

我通常会在读取数据后先做这一步:

import json import pandas as pd with open("./geneformer_model/token_dictionary.json", "r") as f: token_dict = json.load(f) # 假设 adata.var_names 是 gene symbol # 先用你自己的注释表把 symbol 换成 ensembl_id symbol_to_ensembl = {...} # 你维护的映射表 adata.var["ensembl_id"] = adata.var_names.map(symbol_to_ensembl) adata = adata[:, [x in token_dict for x in adata.var["ensembl_id"]]]

这一步会在后续查表的时候省掉大量麻烦,否则debug到深夜都查不出为什么序列全是pad。

3.3 把表达矩阵转换成token序列

核心转换函数长这样:

import numpy as np from tqdm import tqdm def rank_gene_sequences(adata, token_dict, max_length=2048, pad_token_id=25428): X = adata.X if hasattr(X, "toarray"): X = X.toarray() var_ids = adata.var["ensembl_id"].values input_ids = [] attention_masks = [] for i in tqdm(range(X.shape[0]), desc="ranking cells"): row = X[i] nonzero_idx = np.where(row > 0)[0] if len(nonzero_idx) == 0: input_ids.append([pad_token_id] * max_length) attention_masks.append([0] * max_length) continue # 按表达量降序 rows_sorted = nonzero_idx[np.argsort(-row[nonzero_idx])] # 截断或padding rows_sorted = rows_sorted[:max_length] seq = [] for gene_idx in rows_sorted: gid = var_ids[gene_idx] if gid in token_dict: seq.append(token_dict[gid]) seq_len = len(seq) if seq_len < max_length: seq = seq + [pad_token_id] * (max_length - seq_len) mask = [1] * seq_len + [0] * (max_length - seq_len) else: mask = [1] * max_length input_ids.append(seq) attention_masks.append(mask) return np.array(input_ids), np.array(attention_masks)

这个函数在几万细胞的数据量上跑会有些慢,主要慢在逐细胞循环和逐个基因查dict。如果数据量到几十万细胞,建议把X转成CSR矩阵后用稀疏排序逻辑,或者直接用numba加速排序段落。但作为第一版跑通模型,这个函数够了。

对于不同max_length的选择:Geneformer预训练时最大位置编码是2048,所以默认用2048。如果非零基因数很少,可以降到1024甚至512,速度会明显提升。5000个以上非零基因的细胞属于少数,直接截断对分类影响不大。

4. 模型组装:预训练主干加分类头

4.1 用BertConfig定义模型骨架

Geneformer预训练模型用的是BERT结构,所以分类模型我们直接用Transformers里的BertForSequenceClassification最省事。先把config定义好:

from transformers import BertConfig, BertForSequenceClassification config = BertConfig( vocab_size=25429, hidden_size=256, num_hidden_layers=6, num_attention_heads=4, intermediate_size=512, max_position_embeddings=2048, pad_token_id=25428, num_labels=2, ) model = BertForSequenceClassification(config)

这里vocab_sizepad_token_id必须和官方token字典严格一致,否则加载权重时embedding维度对不上。

然后加载之前映射好的权重:

model.bert.load_state_dict(mapped_state_dict, strict=False)

分类头是随机初始化的,不需要从预训练权重加载。

4.2 冻结策略:先学分类头,再微调全部参数

单细胞分类场景里,经常遇到标注细胞数量不多的情况。这个时候直接全参数微调,很容易让预训练学到的基因调控知识被少量样本冲掉。我的经验是先冻结主干,只训练分类头,等分类头收敛了再解冻整个模型,用小学习率做精调。

具体操作:

for param in model.bert.parameters(): param.requires_grad = False # 第一阶段只训练分类头 for param in model.classifier.parameters(): param.requires_grad = True

跑几个epoch以后,再把requires_grad全部设为True:

for param in model.parameters(): param.requires_grad = True

学习率也要注意。预训练模型微调的学习率一般设在1e-53e-5之间。我试过上来直接1e-4,训练loss虽然降得快,但验证集F1反而更差,这大概率是灾难性遗忘导致的。后来改成第一段3e-4只训分类头,第二段1e-5全参数微调,效果稳定很多。

4.3 padding和attention_mask的配合

Geneformer的token id 25428是默认pad token。在Transformers里,padding token对应的attention_mask会被自动处理为0,模型就不会对这些位置做注意力计算。

要注意的是,Geneformer原repo有一些代码会直接处理padding,但当我们用标准BertForSequenceClassification时,必须自己把attention_mask传进去。如果你的collator构造了字典,但忘了放attention_mask键,模型会把padding位置也当成真实token参与计算,效果会神秘变差。

5. 用Trainer跑训练,以及在评估中看什么

5.1 组装Hugging Face Dataset

我习惯先把上一节得到的input_idsattention_masklabels组织成datasets.Dataset

from datasets import Dataset train_data = { "input_ids": input_ids_train, "attention_mask": attention_mask_train, "labels": labels_train, } val_data = { "input_ids": input_ids_val, "attention_mask": attention_mask_val, "labels": labels_val, } train_dataset = Dataset.from_dict(train_data) val_dataset = Dataset.from_dict(val_data)

Dataset的features会自动推断,不需要额外声明。如果显存不够,可以为Dataset设置torch_format,让它在训练时才转torch tensor,省一点内存。

5.2 TrainingArguments参数怎么设

用Trainer之前,先准备好TrainingArguments。我通常这么设置:

from transformers import TrainingArguments, Trainer training_args = TrainingArguments( output_dir="./geneformer_classifier/", evaluation_strategy="epoch", save_strategy="epoch", learning_rate=1e-5, per_device_train_batch_size=16, per_device_eval_batch_size=16, num_train_epochs=5, gradient_accumulation_steps=1, fp16=True, load_best_model_at_end=True, metric_for_best_model="f1", logging_steps=20, report_to="tensorboard", )

metric_for_best_model="f1"需要你提供计算F1的函数,不然Trainer会报错。对于类别不平衡的数据,我不用accuracy做早停指标,而是用macro F1,这样得到的checkpoint更可靠。

5.3 自定义评估指标

简单写一个计算F1、precision、recall的指标函数:

from sklearn.metrics import f1_score, precision_score, recall_score import numpy as np def compute_metrics(eval_pred): logits, labels = eval_pred preds = np.argmax(logits, axis=-1) return { "accuracy": (preds == labels).mean(), "f1_macro": f1_score(labels, preds, average="macro"), "precision_macro": precision_score(labels, preds, average="macro"), "recall_macro": recall_score(labels, preds, average="macro"), }

然后实例化Trainer:

trainer = Trainer( model=model, args=training_args, train_dataset=train_dataset, eval_dataset=val_dataset, compute_metrics=compute_metrics, )

这样就能跑了。

5.4 训练过程中的监控重点

我在训练时会盯着两个东西看:第一个是训练loss有没有突然上升,一旦上升先检查学习率;第二个是验证集的macro F1和recall,尤其是少数类别的recall,如果一直偏低,后面要多想办法。

num_train_epochs不需要设太大,我常用3到5。预训练模型参数已经知道基因关系了,分类任务通常很快收敛。如果训练到第2轮验证集F1就不再涨,直接手动停,不要死等5轮结束,省时间也省显存。

6. 实战踩坑记录:从数据泄漏到显存爆炸

6.1 数据预处理阶段比模型更容易泄漏

Geneformer的rank encoding是对单个细胞内部做的排序,本身不会跨细胞泄漏信息。真正的坑在预处理:如果你对所有细胞合并后的矩阵做了一次全局的normalize_total,或者用全部数据fit的PCA再喂给模型,那训练集和验证集之间的信息就不再独立了。

比如你先对整个anndata做了sc.pp.scale,其中mean和std是全量数据算的,验证集中每个样本的标准化数值已经携带了训练集细胞的信息。模型在验证集上的性能会虚高,但一到新数据上就露馅。正确的做法是:先按样本或按批次做normalize,再做rank encoding;如果确实需要PCA或降维,也必须拆分成训练集合、验证集分别fit,验证集只transform。

更稳妥的切分方式是按患者/样本分组,而不是按细胞随机切。同一个病人的细胞本来就高度相似,随机切会让训练集和验证集出现重复病人信息,评估出来的指标不可信。我在这上面翻过车:按细胞随机划分时AUC有0.98,换成按病人划分后直接掉到0.86。后者才是真实水平。

6.2 输入长度背后的信息取舍

Geneformer默认max_length=2048,但并不是所有细胞都适合这个长度。如果数据里大部分细胞的非零基因数只有几百,那用512或1024就够,速度能快两倍。如果细胞普遍是高深度测序,非零基因数超过4000,强行截断到2048会丢掉不少表达等级信息。

但也要注意,模型的位置编码是预训练时在2048长度上学出来的,如果你把max_length设成4096,位置编码就无法直接加载,整个模型效果反而可能变差。稳妥的做法是先用2048跑一版实验结果,如果验证集F1差距明显,再考虑从头训练位置编码这种更重的方案。

6.3 类别不平衡:用WeightedRandomSampler还是改loss

单细胞分类里正负样本比经常失衡,尤其是罕见细胞类型注释或少数疾病组预测。简单用标准交叉熵会倾向把所有样本都预测成多数类。

我试过几种办法:

  • 设置class weights传进loss,简单有效;
  • WeightedRandomSampler对少数类过采样,配合早停比较好用;
  • 对极度不平衡的场景,可以换focal loss。

Transformers的Trainer默认用模型的loss,想改loss比较方便的方式是继承Trainer重写compute_loss。我用的比较多的还是给分类头前面加一个weight参数,代价最小。

6.4 显存优化:梯度累积与梯度检查点

单细胞数据量一大,batch size经常受限。2048长度的序列输入16条,A100都够呛,更别说常规的3090或者V100。我的处理常规是:

  • per_device_train_batch_size降到4或8;
  • gradient_accumulation_steps调到4或8,等效batch size变成16或64;
  • 实在不行开gradient_checkpointing=True,用显存换一点时间。

打开梯度检查点以后,训练速度会慢一些,但对20GB显存左右的老卡很友好。

6.5 权重和字典版本对不上

最后提醒一个容易被忽略的问题:Hugging Face Hub上的Geneformer权重可能有更新,token_dictionary.json也可能存在版本差异。如果权重变了但字典没更新,模型预测结果会非常离奇,而且很难排查。

我一般会在加载后用一个小测试验证:随机选几个细胞,跑一次前向,看看预测概率是否接近0.5(随机初始化时应该接近),然后看看训练loss能否正常下降。如果第一轮loss就nan,优先检查token字典和模型vocab_size是否匹配。

这整套链路跑下来,最花时间的部分反而不是模型训练,而是数据格式转换和权重适配。但只要把序列生成这个函数写好,后续换一批数据再训练就很顺手了。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询