基于Geneformer的虚拟扰动分析:从Transformer原理到SHAP可解释实践 在单细胞生物学研究里我们经常面临一个很现实的困境想验证某个基因对细胞状态的影响传统做法是做 CRISPR 干扰或基因敲除实验成本高、周期长而且很难在大量基因上平行筛选。于是“虚拟扰动”这个概念越来越被关注——用训练好的模型在计算机里模拟基因改变从而预测细胞会往哪个方向转变。本文要讲的 Geneformer就是这类方法里非常有代表性的一类基础模型。它把单细胞转录组数据当成一种“语言”用 Transformer 架构学习细胞状态的内在规律然后通过虚拟基因敲除、虚拟过表达等方式做扰动预测。这篇文章会从概念入手给出可落地的环境准备、数据处理、模型调用、扰动分析流程再加上 SHAP 可解释性分析帮你把这个技术栈串成一条完整的实操链路。同时这也是一份适合生物信息学入门者、从传统差异表达分析转向深度学习的从业者以及做药物靶点筛选的同学们的技术笔记。即使你没有深度学习背景按照文章顺序一步步操作也能跑通基础的虚拟扰动流程。1. 背景与核心概念1.1 为什么需要虚拟扰动分析真实扰动实验的原理是通过物理或化学手段改变基因的表达再测量由此带来的表型变化。比如 CRISPR-Cas9 可以敲除目标基因siRNA 可以敲低基因然后将处理后的细胞拿去测序观察差异表达基因和通路变化。但真实扰动有几个痛点通量低一次只能测几个或几十个基因成本高试剂、细胞培养、测序费用不便宜周期长从构建载体到拿到结果往往需要数周甚至数月批效应重不同批次、不同培养条件会引入噪声。虚拟扰动则把这些问题转化成了计算问题如果我们有一个足够强大的模型学习到了大规模真实转录组数据中的基因调控规律那么理论上就可以在模型中改变某个或某几个基因的表达状态然后观察细胞状态表征如何变化。这种方式本质上是在做“条件推理”而不是真正的实验干预。1.2 Geneformer 是什么Geneformer 是一个基于 Transformer 架构的预训练模型专门用于单细胞转录组数据。它由 Theodore 等人提出核心思想是把每个细胞的基因表达谱转换为类似于 NLP 中的句子每个基因对应一个“词”基因的表达量排序对应“词序”。模型通过掩码语言建模的方式在大规模人类单细胞数据上预训练学习基因之间的共表达关系、调控关系和细胞状态特征。预训练完成后Geneformer 可以通过迁移学习用于多种下游任务例如细胞类型注释基因调控网络推断扰动状态预测疾病状态分类。其中虚拟扰动分析是它最有特色的应用之一。它的特别之处在于不是简单地对基因表达量做加减法而是从模型学到的基因关系网络中推断扰动后的细胞状态。1.3 AI 扰动模型与虚拟基因敲除的关系“AI 扰动模型”是一个广义说法泛指用机器学习或深度学习模型模拟生物扰动响应的建模方法。Geneformer 属于其中的生成式/条件推断模型。虚拟基因敲除则是具体操作把某个基因的表达量设为极低值或者从基因序列中移除该基因 token再让模型重新预测细胞表征。虚拟基因敲除可以帮助我们回答这类问题这个基因对维持当前细胞状态是否关键抑制该基因后细胞是否会转向另一种类型哪些下游基因的表达会随之改变多个基因同时扰动是否会带来协同效应相比传统差异表达分析虚拟扰动不依赖“两组实验数据之间做统计检验”而是直接在模型学到的生成分布上进行推理因此能够提供一种更机制性的视角。1.4 SHAP 在这里有什么用SHAPSHapley Additive exPlanations是机器学习可解释性分析中常用的一种方法。它源于博弈论中的 Shapley 值用来衡量每个特征对模型预测结果的贡献。在虚拟扰动分析场景下SHAP 可以回答模型判断某个细胞类型时哪些基因的贡献最大虚拟敲除某个基因后为什么细胞状态向量向另一个方向偏移某个细胞样本的预测置信度主要由哪些基因决定因此Geneformer 负责“预测”SHAP 负责“解释”两者结合使用能让你的扰动分析从“看到结果”升级到“理解机制”。2. 环境准备与版本说明在进入代码之前先把环境准备好。本节给出的版本信息以通用稳定版本为例你实际使用时请根据自己的服务器环境调整。2.1 推荐运行环境基因测序数据和 Transformer 模型都需要较大的内存和 GPU 显存建议使用 Linux 服务器或者带有 NVIDIA GPU 的工作站。项目建议配置操作系统Ubuntu 20.04 或 CentOS 7GPUNVIDIA V100 / A100至少 16GB 显存内存64GB 以上硬盘500GB 以上 SSDPython3.8 或 3.9CUDA11.3 以上PyTorch1.12 以上如果只有 CPU 环境可以运行小规模示例但完整预训练模型推理会比较慢。2.2 安装 Python 依赖建议使用 conda 创建独立环境避免依赖冲突。conda create -n geneformer python3.9 conda activate geneformer然后安装 PyTorch。这里以 CUDA 11.3 为例pip install torch1.12.1cu113 torchvision0.13.1cu113 --extra-index-url https://download.pytorch.org/whl/cu113接着安装 Geneformer 所需的核心依赖pip install scanpy pip install anndata pip install loompy pip install datasets pip install transformers pip install shap pip install matplotlib pip install pandas numpyGeneformer 本身可以从 GitHub 克隆到本地通过源码方式导入git clone https://github.com/ctargon/Geneformer.git cd Geneformer pip install -e .注意如果 GitHub 最新代码有更新以官方 README 为准。本文的代码示例基于常见版本细节可能随版本变化但整体思路是稳定的。2.3 数据准备说明Geneformer 官方常用的输入格式是.loom文件。你可以用自己的单细胞数据做转换也可以从公开数据集下载。为了快速跑通流程可以先下载一个小的示例 loom 文件或者使用scanpy自带的数据集做格式转换。下面是一个从 AnnData 转换为 loom 文件的示例import anndata as ad import scanpy as sc # 读取自己的数据 adata sc.read_h5ad(your_data.h5ad) adata.var_names_make_unique() sc.pp.filter_genes(adata, min_cells3) sc.pp.filter_cells(adata, min_genes200) # 构建 loom 文件 adata.write_loom(your_data.loom)这个步骤的目的是让数据满足 Geneformer 对输入格式的要求行为细胞列为基因表达值可以用原始计数或归一化计数。3. 核心原理与关键参数拆解3.1 Geneformer 将转录组转换为“基因语言”Transformer 最初用于自然语言处理它的输入是 token 序列。Geneformer 沿用了这个思路只不过这里的 token 是基因。具体做法如下对每个细胞统计所有基因的表达量按照表达量从高到低排序保留表达量最高的前 N 个基因例如 2048 或 4096 个将排序后的基因 ID 列表作为一行输入在预训练时随机掩盖一部分基因 token让模型根据上下文预测被掩盖的基因。这种设计有一个很自然的优势基因表达量本身是连续值但经过排序后转化为了序数信息。模型不再需要直接拟合表达量的绝对大小而是学习哪些基因倾向于在同一细胞中高表达。这更接近基因调控网络的实际逻辑。3.2 预训练与微调的关系Geneformer 在大规模数据上预训练后可以针对具体任务做微调。预训练阶段模型学习到的是“基因语法”也就是基因之间共表达和调控的一般规律。微调阶段模型把这些规律应用到特定的预测任务上例如预测细胞类型。在你做虚拟扰动之前通常建议先在一个相关任务上微调模型或者直接使用官方提供的已微调模型。如果你用随机初始化的模型做扰动分析得到的输出可能没有生物学意义。3.3 虚拟扰动的数学模型视角从数学上看虚拟扰动可以这样理解原始细胞状态为 X包含一组基因表达值。编码器 E 将 X 映射为隐空间向量 hh E(X)虚拟扰动操作定义为一个函数 P它改变 X 中某个基因的状态得到 XX P(X, g)其中 g 是目标基因。然后计算扰动后的隐向量h E(X)最后我们比较 h 与 h 的距离并解码出差异表达模式。如果 h 与 h 相差很大说明这个基因对细胞状态影响显著如果相差很小说明该基因可能只是一个冗余调节因子。在实际实现中扰动操作有两种常见形式虚拟敲除将目标基因的表达值设置为 0或者从基因序列中删除该 token虚拟过表达将目标基因的表达值设置为较高水平或者在序列中把该基因 token 复制到靠前位置。3.4 虚拟敲除的常见误区很多人会以为虚拟敲除就是简单地把基因 count 改为 0 再重新跑模型但 Geneformer 的输入是排序后的 token 序列所以改动一个基因的表达值会改变基因的排序位置进而影响整个 token 序列。比如你要敲除基因 g直接删掉 g 之后其他基因的相对顺序可能不会变化太大但如果 g 本来排在第 10 位删掉后原来排第 11 的基因会顶上来。这种级联效应是模型输入设计的一部分它反映的是“如果该基因不表达细胞内基因表达程序会如何重新排列”。因此在实现虚拟敲除时不要只简单地把表达矩阵对应位置置零而应该从排序序列层面进行操作。4. 完整实战案例基于 Geneformer 的虚拟扰动分析下面我们用一个完整示例演示整个流程。示例数据为假数据或公开数据代码结构可供实际项目参考。4.1 整体流程概览1. 加载预训练模型 2. 准备 loom 数据 3. 对细胞进行编码得到隐向量 4. 对目标基因执行虚拟敲除/过表达 5. 比较扰动前后的隐向量 6. 解码或层面对比差异基因 7. 使用 SHAP 解释模型决策下面分步骤实现。4.2 创建项目结构建议工程目录如下geneformer_perturbation/ ├── data/ │ └── your_data.loom ├── models/ │ ├── geneformer_pretrained_model │ └── geneformer_config ├── src/ │ ├── load_data.py │ ├── perturb.py │ ├── interpret.py │ └── utils.py └── output/ └── results/4.3 加载预训练模型以源码方式安装 Geneformer 后可以从 Hugging Face 下载预训练模型权重。这里用AutoModel的方式加载需要注意的是Geneformer 要求的输入格式比较特殊因此需要配合官方 tokenizer 一起使用。# 文件路径src/load_data.py from transformers import AutoModel, AutoTokenizer model_path ctargon/geneformer model AutoModel.from_pretrained(model_path, trust_remote_codeTrue) tokenizer AutoTokenizer.from_pretrained(model_path, trust_remote_codeTrue) model.eval() print(Model loaded.)如果你下载的是本地权重把model_path改为本地目录即可。4.4 处理 loom 文件并生成 token 序列Geneformer 官方有一个数据加载工具通常封装在geneformer包中。为了便于理解这里给出核心逻辑# 文件路径src/load_data.py import loompy import numpy as np import pandas as pd def loom_to_tokenized(loom_path, max_len2048): 读取 loom 文件对每个细胞生成按表达量排序的基因 token 序列。 all_tokens [] cell_ids [] with loompy.connect(loom_path) as ds: gene_names ds.ra[gene][:] cell_id ds.ca[CellID][:] for idx in range(ds.shape[1]): expr ds[:, idx] # 只保留表达值大于0的基因 expressed_idx np.where(expr 0)[0] # 按表达量从高到低排序 sorted_idx expressed_idx[np.argsort(expr[expressed_idx])[::-1]] # 截断或填充到 max_len token_seq gene_names[sorted_idx][:max_len] all_tokens.append(token_seq) cell_ids.append(cell_id[idx]) return all_tokens, cell_ids tokens, cell_ids loom_to_tokenized(data/your_data.loom)实际运行时官方实现的 tokenizer 还会做更多细节处理比如去掉低质量基因、处理批次效应等。上面的代码只展示排序 token 的基本思想。4.5 对细胞进行编码将 token 序列转换为模型输入并得到每个细胞的 embedding。# 文件路径src/perturb.py import torch def encode_cells(token_sequences, model, tokenizer, batch_size32): 输入 token 序列列表输出每个细胞的 embedding。 embeddings [] for i in range(0, len(token_sequences), batch_size): batch token_sequences[i:ibatch_size] encodings tokenizer(batch, paddingTrue, truncationTrue, return_tensorspt, is_split_into_wordsTrue) with torch.no_grad(): output model(**encodings) # 取 CLS token 或平均池化作为细胞 embedding # 这里使用平均池化示例 pooled output.last_hidden_state.mean(dim1) embeddings.append(pooled.cpu().numpy()) return np.vstack(embeddings) original_embeddings encode_cells(tokens, model, tokenizer)4.6 实现虚拟基因敲除虚拟基因敲除的核心是修改 token 序列。下面我们实现一个简单版本把目标基因从每个细胞的 token 序列中移除。# 文件路径src/perturb.py def virtual_knockout(token_sequences, target_gene): 对每个细胞执行虚拟敲除从 token 序列中去掉目标基因。 perturbed_tokens [] for seq in token_sequences: seq_list list(seq) new_seq [g for g in seq_list if g ! target_gene] perturbed_tokens.append(new_seq) return perturbed_tokens target_gene TP53 ko_tokens virtual_knockout(tokens, target_gene) ko_embeddings encode_cells(ko_tokens, model, tokenizer)类似地虚拟过表达可以简单理解为把目标基因放到序列最前面def virtual_overexpression(token_sequences, target_gene, top_n1): perturbed_tokens [] for seq in token_sequences: seq_list list(seq) if target_gene in seq_list: seq_list.remove(target_gene) # 放到最前面模拟高表达 new_seq [target_gene] seq_list perturbed_tokens.append(new_seq) return perturbed_tokens oe_tokens virtual_overexpression(tokens, target_gene) oe_embeddings encode_cells(oe_tokens, model, tokenizer)这里的操作是为了演示排序序列层面的变化。更精细的实现会考虑基因之间的调控关系比如把一组基因同时过表达或敲除。4.7 比较扰动前后的细胞状态有了原始 embedding 和扰动后 embedding我们就能量化扰动影响。# 文件路径src/perturb.py from sklearn.metrics.pairwise import cosine_similarity import pandas as pd import numpy as np def compute_perturbation_effect(original, perturbed, cell_ids): 计算每个细胞扰动前后 embedding 的欧氏距离和余弦相似度。 diff original - perturbed euclidean_dist np.linalg.norm(diff, axis1) cosine_sim np.array([cosine_similarity(original[i:i1], perturbed[i:i1])[0][0] for i in range(len(original))]) result pd.DataFrame({ cell_id: cell_ids, euclidean_dist: euclidean_dist, cosine_sim: cosine_sim }) return result effect_df compute_perturbation_effect(original_embeddings, ko_embeddings, cell_ids) print(effect_df.head())你可以对每个目标基因都执行一次上述流程然后按平均欧氏距离排序得到“该细胞类型中影响最大的基因列表”。4.8 观察扰动对细胞类型特征的影响如果你已经有细胞类型注释可以分析不同细胞类型对同一基因扰动的敏感度。# 假设置细胞类型存储在 cell_types 列表中 # effect_df[cell_type] cell_types # 按细胞类型分组统计 summary effect_df.groupby(cell_type)[euclidean_dist].agg([mean, std]).sort_values(mean, ascendingFalse) print(summary)这种分析能提示我们某个基因可能对一个细胞亚群特别关键但对另一个亚群影响不大。4.9 可视化扰动距离分布为了直观展示结果可以用箱线图画出不同扰动基因下细胞 embedding 距离的分布。import matplotlib.pyplot as plt import seaborn as sns # 假设 multi_gene_result 是一个包含 gene_name, dist 的长表 DataFrame # sns.boxplot(datamulti_gene_result, xgene_name, ydist) # plt.xticks(rotation45) # plt.tight_layout() # plt.savefig(output/perturbation_boxplot.png, dpi150)实际项目中你可以将多个基因的扰动结果合并在一起然后绘制横向比较图。5. 用 SHAP 增强可解释性5.1 为什么需要 SHAP虚拟扰动分析解决了“这个基因改变后细胞状态怎么变”的问题但没有直接回答“模型为什么这么判断”。SHAP 可以为我们提供另一视角模型在区分细胞类型或预测某个状态时每个基因贡献了多少。当你拥有一个用于细胞类型分类的微调模型时可以对单个细胞的模型预测做 SHAP 解释找到推动预测结果的关键基因。5.2 构建一个简单的分类模型作为解释对象这里我们以基因表达量特征为输入做一个简单的二分类模型例如区分“扰动敏感型”和“非敏感型”。# 文件路径src/interpret.py import shap import xgboost as xgb from sklearn.model_selection import train_test_split # 假设 X 是细胞的特征矩阵行为细胞列为基因y 是二分类标签 # 这里用随机数据演示 import numpy as np X np.random.rand(1000, 50) # 1000个细胞50个基因特征 y np.random.randint(0, 2, 1000) X_train, X_test, y_train, y_test train_test_split(X, y, test_size0.2, random_state42) model xgb.XGBClassifier(n_estimators100, max_depth4, random_state42) model.fit(X_train, y_train) # 创建 SHAP 解释器 explainer shap.TreeExplainer(model) shap_values explainer.shap_values(X_test)5.3 SHAP 图怎么画新手最常问的问题就是“SHAP 图怎么画”。下面给出几种最常见的 SHAP 图绘制方法。5.3.1 Summary Plot总体特征重要性# 文件路径src/interpret.py import matplotlib.pyplot as plt shap.summary_plot(shap_values, X_test, feature_names[fgene_{i} for i in range(X.shape[1])]) plt.savefig(output/shap_summary.png, dpi150, bbox_inchestight)summary plot 中每行代表一个基因横轴是 SHAP 值颜色表示特征值高低。如果某个基因在高表达时 SHAP 值为正说明它推动模型预测为正类。5.3.2 Bar Plot平均绝对 SHAP 值shap.summary_plot(shap_values, X_test, plot_typebar) plt.savefig(output/shap_bar.png, dpi150, bbox_inchestight)bar plot 只展示平均绝对 SHAP 值适合快速查看哪些基因整体影响最大。5.3.3 Waterfall Plot单个细胞归因如果你想看某一个细胞样本的预测是如何被各个基因推高的可以使用 waterfall plot# 查看第0个测试样本 shap.plots.waterfall(shap.Explanation(valuesshap_values[0], base_valuesexplainer.expected_value, dataX_test[0], feature_names[fgene_{i} for i in range(X.shape[1])])) plt.savefig(output/shap_waterfall_sample0.png, dpi150, bbox_inchestight)waterfall plot 的起点是基础值红色柱子表示正向贡献蓝色柱子表示负向贡献最终到达模型预测值。5.4 将 SHAP 与虚拟扰动结合我们可以把 SHAP 的结果与虚拟扰动的结果交叉对比虚拟扰动距离大的基因是否也是 SHAP 值排名靠前的基因如果某个基因在 SHAP 中贡献大但虚拟敲除后状态变化不大可能说明模型虽然依赖该基因做分类但细胞状态具有较强的鲁棒性。这种对比能帮助我们识别更可信的基因靶点。6. 常见问题与排查思路在实际运行 Geneformer 虚拟扰动分析时你可能会遇到下面这些高频问题。问题现象常见原因解决思路loom 文件读取失败loom 文件损坏或格式不正确检查用 scanpy 重写 loom 文件确认基因和细胞注释字段齐全加载模型报错transformers 版本与 Geneformer 不匹配升级/降级 transformers建议参考官方 requirementstoken 序列长度过长细胞中高表达基因超过模型最大长度设置 max_len例如 2048 或 4096GPU 显存不足单个 batch 细胞数太多减小 batch_size或使用梯度累积方式虚拟敲除后 embedding 没有变化目标基因原本就不在这个细胞的表达基因列表中检查目标基因在数据中的存在情况过滤掉不表达的细胞SHAP 计算速度慢特征数量多、样本多先对基因做特征筛选或者使用 background data 更小的 explainer6.1 模型输入与基因名不一致如果你用自己的数据基因名可能是ENSEMBL ID而 Geneformer 预训练模型使用的是Gene Symbol这时需要做 ID 转换。建议用mygene或biomart做批量转换然后过滤掉无法映射的基因。6.2 虚拟敲除操作不生效请检查你的 token 序列中是否真的包含目标基因。如果目标基因在数据中的表达量本来就非常低排序后可能被截断掉了。可以考虑把 max_len 调大或者先确认该基因在目标细胞类型中是否表达。6.3 模型预测结果不稳定Transformer 模型对 token 顺序有一定敏感性。虚拟扰动造成的顺序变化可能在不同细胞中产生不同的结果。建议使用多个随机种子进行编码对同一扰动重复多次并取平均增加扰动基因的组合分析而不是只分析单个基因。7. 最佳实践与工程建议7.1 数据层面数据质量优先虚拟扰动的结果高度依赖输入数据的质量。建议先做严格的质量控制过滤掉 doublets 和低质量细胞。统一基因注释所有细胞使用同一套基因名体系避免混用 Symbol 与 Ensembl ID。考虑批次效应如果是多个数据集合并建议先做批次校正或者使用 Geneformer 自带的 batch 处理逻辑。7.2 模型层面尽量使用微调后的模型做扰动分析因为预训练模型学习的是通用规律对特定细胞类型的敏感度可能不够。不要直接把虚拟扰动等同于真实实验。模型预测的是“转录组状态偏移趋势”不是真实的表型实验数据需要后续湿实验验证。记录模型版本和权重来源保证实验可复现。7.3 扰动分析层面先做单个基因的虚拟敲除再升级到双基因组合扰动。组合扰动可以揭示基因之间的上位效应或协同作用。在比较多个基因时除了比较 embedding 距离还可以关注下游差异基因的重叠情况。如果两个基因扰动后影响的通路高度重叠它们可能处于同一调控模块中。最好把虚拟扰动结果与公开的扰动图谱数据如 Perturb-seq但注意数据版权进行对比验证模型的合理性。7.4 SHAP 分析层面SHAP 值反映的是模型内部的决策逻辑不代表真实因果效应解释时要避免过度因果化描述。对于高维基因特征可以先对特征做粗筛否则 SHAP 计算时间和内存开销会很大。也可以用shap.Explainer接口统一处理不同类型模型减少切换成本。8. 后续进阶方向如果你已经顺利跑通了基础的虚拟扰动流程下一步可以从这几个方向继续深入。8.1 多基因组合扰动实际生物学过程往往是多个基因协同作用。你可以实现一个组合扰动函数对一组基因同时进行虚拟敲除并计算组合扰动带来的状态向量变化。这需要合理的实验设计比如先按单基因扰动距离排序再取 top 20 的基因做两两组合。8.2 引入时间动态建模单细胞数据往往是静态快照但虚拟扰动可以模拟状态转移的方向。如果结合 RNA velocity 或轨迹推断方法可以让扰动分析从“状态偏移”升级为“轨迹变化”更贴近动态过程。8.3 面向药物靶点筛选识别出关键基因后可以查询这些基因是否属于已知药物靶点或者使用深度学习模型预测小分子干预后的转录组响应。这类应用需要加入药物结构信息或基因-药物关系数据库整体复杂度会明显上升但研究价值也更高。8.4 与大规模真实扰动数据对齐如果有条件可以利用公开的 Perturb-seq 等真实扰动数据进行验证。对比模型虚拟扰动结果与真实扰动实验的吻合度是整个方法可信度的关键。没有真实数据时至少应该在多种细胞类型中做交叉验证避免结论过拟合到单一数据。这篇笔记把 Geneformer 虚拟扰动分析的完整链路梳理了一遍从原理、环境、代码到 SHAP 解释都有覆盖。建议你从一个小型公开数据集开始先跑通“单个基因虚拟敲除 - embedding 距离比较 - SHAP 图绘制”的主流程再逐步扩展到多基因组合和真实数据验证。深度学习的模型是工具真正值钱的是在下游实验中把“预测”转化为“机制假设”的能力。如果实际操作中遇到其他报错欢迎在评论区交流你的具体错误信息和环境信息我会尽量帮忙定位。