
合作方上周丢给我一个需求手上有病人和对照的单细胞转录组数据要训练一个分类模型把两组样本区分开。换作以前我第一反应是找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 Encoderhidden size 2564个注意力头最大序列长度2048这在Transformer模型里算非常轻量的。也正因为轻量单卡微调是可行的。1.3 什么情况下值得用Geneformer什么情况不必每个细胞对应一个类别标签这类任务正是Geneformer下游微调的强项尤其是训练标注样本不多比如只有几例病人和几例对照类别之间的差异不是单个基因表达量高低而是基因组合和调控网络层面的差异需要模型迁移到新的数据集上继续微调。如果只是对不同细胞类型做注释且marker基因又很明确那用传统方法可能更快。不要在简单问题上强行上Transformer这是我一贯的原则。2. 环境与依赖把Hugging Face生态拉起来2.1 基础安装我建议先建一个干净的conda环境conda create -n geneformer python3.10 -y conda activate geneformer pip install transformers4.30 datasets accelerate huggingface_hub pip install scanpy anndata scikit-learn pandas numpy单细胞下游分析还需要leidenalg、harmony之类看具体场景再装。核心就是transformers和datasets这两个是整个微调链路的地基。2.2 从Hugging Face Hub拉取Geneformer资源Geneformer官方权重目前挂在Hugging Face Hub上仓库ID是ctheodoris/Geneformer。里面有几个关键文件模型权重文件pytorch_model.bingene 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_okTrue) hf_hub_download( repo_idrepo_id, filenamepytorch_model.bin, local_dirlocal_dir, ) hf_hub_download( repo_idrepo_id, filenametoken_dictionary.json, local_dirlocal_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名对不上。解决思路有两种把官方权重key做一次映射比如把encoder.layer.0.attention.self.query.weight映射成bert.encoder.layer.0.attention.self.query.weight先创建BertForSequenceClassification再load_state_dict时用strictFalse加上key名替换逻辑。我在项目里通常会先把官方权重加载成一个普通字典然后做一层字符串替换再扔给model.bert加载。代码如下import torch from transformers import BertForSequenceClassification, BertConfig state_dict torch.load(./geneformer_model/pytorch_model.bin, map_locationcpu) # 示例原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。过程拆开看取一个细胞找到所有表达量大于0的基因按表达量从高到低排序把排序后的基因替换成词典里的token id截断到固定长度默认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_length2048, pad_token_id25428): 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]), descranking 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_size25429, hidden_size256, num_hidden_layers6, num_attention_heads4, intermediate_size512, max_position_embeddings2048, pad_token_id25428, num_labels2, ) model BertForSequenceClassification(config)这里vocab_size和pad_token_id必须和官方token字典严格一致否则加载权重时embedding维度对不上。然后加载之前映射好的权重model.bert.load_state_dict(mapped_state_dict, strictFalse)分类头是随机初始化的不需要从预训练权重加载。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全部设为Truefor param in model.parameters(): param.requires_grad True学习率也要注意。预训练模型微调的学习率一般设在1e-5到3e-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_ids、attention_mask、labels组织成datasets.Datasetfrom 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_strategyepoch, save_strategyepoch, learning_rate1e-5, per_device_train_batch_size16, per_device_eval_batch_size16, num_train_epochs5, gradient_accumulation_steps1, fp16True, load_best_model_at_endTrue, metric_for_best_modelf1, logging_steps20, report_totensorboard, )metric_for_best_modelf1需要你提供计算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, averagemacro), precision_macro: precision_score(labels, preds, averagemacro), recall_macro: recall_score(labels, preds, averagemacro), }然后实例化Trainertrainer Trainer( modelmodel, argstraining_args, train_datasettrain_dataset, eval_datasetval_dataset, compute_metricscompute_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_length2048但并不是所有细胞都适合这个长度。如果数据里大部分细胞的非零基因数只有几百那用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或8gradient_accumulation_steps调到4或8等效batch size变成16或64实在不行开gradient_checkpointingTrue用显存换一点时间。打开梯度检查点以后训练速度会慢一些但对20GB显存左右的老卡很友好。6.5 权重和字典版本对不上最后提醒一个容易被忽略的问题Hugging Face Hub上的Geneformer权重可能有更新token_dictionary.json也可能存在版本差异。如果权重变了但字典没更新模型预测结果会非常离奇而且很难排查。我一般会在加载后用一个小测试验证随机选几个细胞跑一次前向看看预测概率是否接近0.5随机初始化时应该接近然后看看训练loss能否正常下降。如果第一轮loss就nan优先检查token字典和模型vocab_size是否匹配。这整套链路跑下来最花时间的部分反而不是模型训练而是数据格式转换和权重适配。但只要把序列生成这个函数写好后续换一批数据再训练就很顺手了。