ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

Geneformer虚拟扰动分析:从单细胞转录组到基因敲除可解释实践

Geneformer虚拟扰动分析:从单细胞转录组到基因敲除可解释实践 Geneformer 虚拟扰动分析是近几年单细胞转录组与深度学习方法结合后很有代表性的分析思路。它把基因表达矩阵转换成类似自然语言的 token 序列再用 Transformer 结构学习细胞状态之后通过“人为修改输入序列中的基因”来模拟扰动并用机器学习可解释性方法定位关键基因。这种分析最吸引人的地方在于它不需要真的做一次 CRISPR 实验就能在模型里回答“如果这个基因不表达细胞状态会怎么变”。本文会从原理、环境搭建、数据准备、虚拟基因敲除实现、SHAP 解释、结果验证和排错几个角度完整走一遍这个分析流程。适合已经会跑单细胞基础分析、但对大模型和可解释性还不太熟悉的读者。1. 先理解 Geneformer、虚拟扰动和虚拟基因敲除到底在做什么1.1 用一句话解释这套分析链路Geneformer 本质上是一个在单细胞转录组数据上预训练的 Transformer 模型。它把每个细胞的基因表达情况做成一段“基因 token 序列”让模型学习基因之间的共表达和调控关系。虚拟扰动分析则是利用这种学到的关系在输入层面删除、增强或替换某些基因 token再观察模型输出的隐藏状态或下游预测结果改变了什么。虚拟基因敲除是虚拟扰动里最常用的一种形式。真实实验中的基因敲除需要构建载体、转染细胞、筛选克隆周期长成本高。虚拟敲除则在模型层面把某个基因 token 从细胞状态序列中移除或者将其表达量置为不表达的等级然后看模型还能不能正确识别细胞类型、状态或疾病特征。SHAP 在这条链路里的角色不是用来做敲除本身而是用来回答“模型做出这个预测时哪些基因贡献最大”。SHAP 值可以评价每个输入基因 token 对预测结果的边际贡献也可以配合虚拟扰动结果筛选出更值得做真实实验验证的候选基因。1.2 技术定义和核心机制从技术角度说Geneformer 的输入不是普通数值表达矩阵而是按基因表达量排序后的 token 序列。每个基因对应词表里的一个固定 token表达量高低决定 token 在序列中的位置。高表达基因排在前面低表达基因排在后面。这种 rank value encoding 的设计让模型不再关心具体表达量的绝对值而是更关注基因之间的相对排序和上下文关系。虚拟扰动发生在模型推理阶段。常见做法有几种删除某个基因 token让模型在缺失该基因的条件下重新预测细胞状态。把某个基因 token 移动到序列末尾模拟低表达或沉默状态。把某个基因 token 重复或前置模拟高表达或过表达状态。将扰动后的输入与原始输入一起送入模型比较隐藏层输出差异。这些操作都只修改输入不修改模型参数。所以虚拟扰动可以批量做理论上一个细胞一条序列遍历候选基因集合就能得到一张“基因扰动影响矩阵”。1.3 和真实实验相比虚拟扰动的边界在哪里这里要非常清楚地说一句虚拟扰动是假设生成工具不是实验替代品。它只能告诉你在当前模型学到的数据分布里某个基因缺失会怎样影响细胞状态不能证明真实细胞里也会发生同样变化。原因也简单模型学的是训练集里的相关性不是因果关系。转录组数据通常只有数千到数万个细胞不代表所有细胞状态。表达量排序丢失了部分定量信息更接近“高低关系”而非精确倍数变化。SHAP 值和扰动输出都受训练分布偏差影响对稀有细胞类型可能不稳定。因此正常用法是把虚拟扰动和 SHAP 作为候选基因排序工具把排名靠前的基因再放到真实实验里验证。这也是这类方法目前最稳妥的定位。2. 环境准备Python、CUDA、Geneformer 依赖和数据集规划2.1 学习环境的最小配置建议Geneformer 这类模型对显存有一定要求但并不是所有环节都需要大显存。常见的分工方式如下环节最小配置推荐配置说明数据预处理和 token 化16 GB 内存32 GB 内存单细胞矩阵解压后可能很大模型推理8 GB 显存16 GB 以上显存批量扰动时显存压力明显模型微调12 GB 显存24 GB 以上全量微调显存要求高SHAP 解释8 GB 显存16 GB 以上KernelExplainer 类方法开销大如果本地没有 GPU也可以用 CPU 先跑一个小数据集验证流程。显存不够时优先减少批量大小其次缩短序列长度最后才是降低模型精度。注意不要把“能加载模型”当作“环境正常”。虚拟扰动分析要反复比较不同扰动输入之间的差异至少要在 2 到 3 个细胞类型上看到符合预期的结果才说明环境可信。2.2 安装 Python 和基础依赖建议使用 Python 3.9 或 3.10。Geneformer 背后依赖 PyTorch、Hugging Face Transformers、datasets 等生态版本之间需要对齐。安装 PyTorch 时要根据自己的 CUDA 版本选择命令。下面是一个常见流程实际版本号要以官方安装页为准conda create -n geneformer python3.10 -y conda activate geneformer pip install torch --index-url https://download.pytorch.org/whl/cu118 pip install transformers pip install datasets pip install scanpy anndata pip install shap pip install jupyterlab安装完成后运行一个快速检查脚本import torch import transformers import shap print(torch, torch.__version__) print(transformers, transformers.__version__) print(cuda available, torch.cuda.is_available())输出里cuda available为 True说明 GPU 可用为 False 也能跑只是速度会慢很多。2.3 数据准备从单细胞矩阵到模型可用格式Geneformer 的输入依赖基因表达矩阵。常见来源是10x Genomics 输出的matrix.mtx、barcodes.tsv、features.tsv。已经整理好的h5ad文件。GEO 或 CellxGene 上公开的单细胞数据集。使用 Scanpy 读取 h5ad 文件并检查基本结构import scanpy as sc adata sc.read_h5ad(path/to/your_data.h5ad) print(adata.shape) print(adata.obs.columns) print(adata.var_names[:10])这里要特别检查两个信息基因名格式是否统一以及表达矩阵是否已经做过对数归一化。Geneformer 的 rank value encoding 对原始数值的具体量纲不敏感但基因名必须是模型词表能识别的形式常见是GENE这种大写标准名。数据量方面学习环境不需要完整复现论文里的千万细胞量级。一个包含 2000 到 10000 个细胞的子集足够跑通虚拟扰动和 SHAP 全流程。关键不是数据多而是包含你关心的细胞类型。3. 核心流程拆分表达矩阵、Token 化、模型推理和特征提取3.1 表达矩阵如何变成基因 token 序列Geneformer 的核心输入单位是细胞每个细胞是一条序列。具体的转换过程可以拆成四步对每个细胞取出所有基因的表达量。过滤掉表达量为 0 或表达量过低的基因。把基因按表达量从高到低排序。根据基因名映射到模型词表中的 token ID得到序列。这种情况下表达量高低不再通过数值体现而是通过序列位置体现。高表达基因排在前低表达基因排在后。如果项目源码里已经有现成的 tokenizer 类优先使用官方实现。下面这段代码只是说明数据格式实际调用需要替换成你的模型名称# 假设 gene_list 是当前细胞中按表达量从高到低排列的基因列表 gene_rank_list gene_list[:max_genes] # 截断到模型支持的最大长度 # 使用模型对应的 tokenizer 把基因名映射为 token id token_ids tokenizer(gene_rank_list, max_length2048, paddingmax_length, truncationTrue) # 得到最终输入 input_ids token_ids[input_ids] attention_mask token_ids[attention_mask]实际项目中你还需要处理没有出现在词表里的基因。常见做法是丢弃不会影响主流程但要在日志里记录丢弃比例。如果丢弃比例超过 10%说明物种或基因命名方式不匹配要先修正数据不要硬跑。3.2 模型加载和隐藏状态提取Geneformer 预训练模型的核心输出不是一个简单的细胞标签而是每个 token 对应的隐藏状态。我们可以用这些隐藏状态作为“细胞嵌入”后续用于聚类、差异分析或扰动比较。加载模型的通用思路如下from transformers import AutoModel, AutoTokenizer model_name your-local-geneformer-model-path model AutoModel.from_pretrained(model_name) tokenizer AutoTokenizer.from_pretrained(model_name) model.eval()因为你可能需要批量处理数千个细胞建议使用 DataLoader 或 datasets 库管理输入。一个关键点是模型推理阶段要关闭梯度计算import torch inputs { input_ids: torch.tensor(input_ids).unsqueeze(0), attention_mask: torch.tensor(attention_mask).unsqueeze(0), } with torch.no_grad(): outputs model(**inputs) # 常见的两种使用方式 # 1. 最后一层隐藏状态 last_hidden outputs.last_hidden_state # 2. 对非 padding 位置做均值池化得到细胞级向量 valid_mask inputs[attention_mask].unsqueeze(-1).bool() pooled (last_hidden * valid_mask).sum(dim1) / valid_mask.sum(dim1)细胞级向量的质量会直接影响下游扰动比较的稳定性。建议在正式实验前先拿一批已知细胞类型标签的细胞做一次聚类或相似度检查确认不同细胞类型能被区分开再去跑扰动。3.3 学习环境如何验证模型是否真的正常工作一个常见错误是模型加载成功后就直接跑虚拟敲除结果发现不同基因敲除完全没差异。问题可能不出在敲除逻辑而是数据输入格式已经错了。建议先做两个简单检查检查输入序列中基因排序是否符合预期。打印前 20 个基因确认高表达基因排在最前。检查同一细胞重复输入两次两次得到的池化向量是否几乎一致。如果不一致说明随机 dropout 或数据加载有问题。完成这两步后再开始正式的扰动分析后续排错会省很多时间。4. 实现虚拟基因敲除和扰动分析的关键代码4.1 扰动对象是敲除一个基因还是扰动一组基因虚拟敲除不要一上来就遍历全部两万多个基因。首先要确定候选基因集合。可以考虑以下范围你关注的细胞类型中已知的关键转录因子。差异表达分析得到的 top 显著基因。某个信号通路里的核心基因。SHAP 值排序靠前的基因。对于每个候选基因分别构建原始序列和扰动序列。扰动序列则分以下几种情况扰动方式操作适用场景缺失模拟敲除把该基因 token 删除模拟完全敲除移到末尾把该基因 token 放到序列尾部模拟表达量显著下调复制并前置把该基因 token 复制后放到序列前面模拟过表达替换为另一个基因把 token A 换成 token B模拟异位表达实际项目中最常见的是“缺失模拟敲除”和“移到底部模拟沉默”。下面以缺失模拟为例说明代码结构。4.2 虚拟敲除的最小实现def remove_gene_token(input_ids, attention_mask, gene_token_id): 从 token 序列中移除指定基因 token。 返回新的 input_ids 和 attention_mask。 keep_mask input_ids ! gene_token_id new_input_ids input_ids[keep_mask] new_attention_mask attention_mask[keep_mask] # 保持长度一致尾部填充 padding token seq_len input_ids.size(0) pad_token_id tokenizer.pad_token_id if new_input_ids.size(0) seq_len: pad_len seq_len - new_input_ids.size(0) pad_ids torch.full((pad_len,), pad_token_id, dtypenew_input_ids.dtype) new_input_ids torch.cat([new_input_ids, pad_ids]) pad_mask torch.zeros(pad_len, dtypeattention_mask.dtype) new_attention_mask torch.cat([new_attention_mask, pad_mask]) return new_input_ids, new_attention_mask这里有几个细节需要注意如果同一个基因在一次输入中出现了多次而你想模拟“完全敲除”应该把所有该 token 都删除。padding token 不能放进 attention 计算里否则会把无效位置当成真实基因参与比较。删除后会改变序列长度所以需要重新 padding这也是容易出错的地方。4.3 批量扰动比较原始状态和扰动状态的差异分数拿到扰动后的输入后可以计算扰动前后细胞向量的差异。def compute_perturbation_delta(model, original_input, perturbed_input): with torch.no_grad(): orig_output model(**original_input) pert_output model(**perturbed_input) orig_vec mean_pooling(orig_output.last_hidden_state, original_input[attention_mask]) pert_vec mean_pooling(pert_output.last_hidden_state, perturbed_input[attention_mask]) delta torch.linalg.vector_norm(pert_vec - orig_vec, dim-1) return delta.item()这个delta的含义是敲除某个基因后细胞状态向量变化有多大。变化越大说明该基因对当前细胞的代表性状态影响越强。但只看向量范数还不够。更好的做法是记录扰动前后状态向量在哪些维度变化最大。扰动后细胞更接近哪一类已知细胞。多个候选基因敲除后差异分数排序是否稳定。不同 batch 之间结果是否一致。注意如果所有基因的 delta 都约等于 0先不要怀疑模型先检查你构造的扰动输入是否真的和原始输入不同。4.4 从单个细胞到细胞群体结果要怎么聚合每个细胞的基因排序和表达分布都不一样。同一个基因在一个细胞里可能排第 50 位在另一个细胞里可能排第 3000 位。所以不能把单个细胞的结果直接当成基因的全局重要性。常见聚合方式包括按细胞类型分组计算每个基因扰动分数的中位数或均值。按候选基因分组统计“扰动后改变细胞类型分类结果”的比例。对每个基因计算扰动影响最大的细胞比例做成基因排序表。输出结果可以整理成一张 CSV 表列包括基因名、细胞类型、原始表达 rank、扰动方式、扰动前状态、扰动后状态、扰动分数、是否改变预测标签。这样后续用 SHAP 解释和真实实验验证时信息才不会丢。5. 用 SHAP 解释基因贡献从扰动结果到可解释特征5.1 SHAP 在基因扰动分析里的定位SHAP 来自博弈论中的 Shapley 值核心思想是衡量每个特征对预测结果的边际贡献。在基因表达和虚拟扰动分析中SHAP 能告诉你模型最终判断某个细胞属于某个类型时哪些基因的贡献更关键。它和虚拟敲除的区别在于虚拟敲除是直接修改输入观察输出变化属于反事实推理。SHAP 是在完整输入上计算每个特征的贡献属于事后解释。两种方法可以交叉使用。SHAP 排序靠前的基因再做一遍虚拟敲除看是否真的能改变细胞状态预测虚拟敲除变化大的基因再看 SHAP 值是否也同样突出。5.2 常见实现方式Geneformer 是 Transformer 模型直接套用 TreeExplainer 不合适。常见的解释路径有以下几种对细胞级池化向量训练一个简单分类器然后对这个分类器做 SHAP 解释。使用 SHAP 的 KernelExplainer 对基因 token 输入做近似解释但计算成本很高。基于 attention 权重做简化归因虽然不是严格 SHAP但可以作为补充。下面是一段用shap.KernelExplainer解释简单模型的示例代码。它主要用于理解思路实际项目中要根据模型输出类型和特征维度调整。import shap import numpy as np # 假设 pooled 是 (样本数, 特征维度) 的细胞向量 # 假设 label 是细胞类型编码 background pooled[:100] test_data pooled[100:110] # 这里用随机森林作为说明实际中可以替换成你自己训练的细胞分类器 from sklearn.ensemble import RandomForestClassifier clf RandomForestClassifier(n_estimators50, random_state0) clf.fit(pooled, label) explainer shap.KernelExplainer(clf.predict_proba, background) shap_values explainer.shap_values(test_data, nsamples200)shap_values的维度通常是(样本数, 特征维度, 类别数)。你可以按类别取出 SHAP 值再做基因名称映射。如果你的特征列名本身就是基因名画 SHAP 图时可以直接用shap.summary_plot(shap_values[1], test_data, feature_namesgene_names)这里的gene_names要和池化向量的维度严格一一对应否则图形里的基因名会错位这是实践中最容易出的低级错误。5.3 用哪些图形和指标汇报结果公开讨论和报告里比较常见的 SHAP 展示方式有summary plot整体看哪些基因对预测贡献大。bar plot看基因贡献排名。dependence plot看某个基因表达量和 SHAP 值的关系。force plot看单个细胞的解释。这几种图对应的输出结果不太一样建议汇报时至少保留一份基因贡献排名表因为表格比图片更适合后续筛选和归档。表格格式可以如下排名基因名平均绝对 SHAP 值主要影响类别对应扰动分数是否进入候选验证列表1NKX2-10.042上皮细胞0.87是2FOXA10.031上皮细胞0.82是3GATA30.025管腔细胞0.61是这里的数值只是示例格式实际数值和你用的模型、数据、细胞类型直接相关不要把它当成固定阈值。5.4 SHAP 结果要谨慎对待的三件事第一SHAP 值高的基因不一定有因果关系它只代表模型预测路径里贡献较大。第二如果输入特征之间存在高度共线性SHAP 的分配可能会在相关基因之间波动。第三单细胞数据零值比例高很多基因表达量低但依然有功能意义SHAP 只能解释模型不能解释全部生物学。因此在结论里建议把 SHAP 和虚拟扰动结果都作为“候选依据”不要说成“该基因是疾病驱动基因”。6. 结果验证、异常现象和排查路径6.1 怎么判断虚拟扰动分析结果可信一个可信的虚拟扰动分析通常要满足几个条件已知细胞类型标签能通过原始状态向量正确区分。扰动已知关键转录因子后细胞类型预测显著改变。阴性对照基因扰动后细胞状态几乎不变。同一实验在不同随机种子下结果排序基本稳定。阴性对照很关键。建议构造一组“不太可能影响细胞状态”的基因比如在许多细胞类型中普遍低表达的基因然后观察它们的扰动分数是否明显低于阳性基因。如果没有对照你很难判断所有结果是不是噪声。6.2 常见问题排查表问题现象可能原因检查方式处理建议模型加载失败模型路径错误或词表文件缺失检查模型路径、config.json、tokenizer 文件确认下载完整目录不要只放权重文件token 化后大部分基因被丢弃基因名格式和词表不一致打印词表前 20 个基因对比输入基因名统一基因名格式例如转换为标准 HUGO 符号所有扰动分数都接近 0扰动输入没有真正改变 token 序列打印扰动前后 token 列表确认删除生效检查 keep_mask 逻辑重复 token 是否都删除不同批次结果波动大模型有 dropout或测试数据量太少固定随机种子多次重复试验推理前调用 model.eval()并关闭梯度SHAP 计算非常慢特征维度过高且使用 KernelExplainer检查样本数和特征数降维到细胞类型分类任务或者只对 top 基因做解释GPU 显存不足批量大小或序列长度过大观察显存占用日志减小 batch size截断序列长度使用梯度检查敲除基因后细胞类型预测不变该基因本来就对分类贡献小查看 SHAP 贡献排名不等于基因无功能只说明模型没有学到相关关系6.3 从现象倒推的三个优先排查顺序先检查输入。先确认原始序列的基因排序是否正常token ID 映射是否正确attention mask 是否覆盖了所有有效 token。输入错了后面所有结果都没有意义。再检查扰动构造。删除基因后长度是否重新 paddingpadding 位置是否被 mask同一个基因是否有多个 tokenID。扰动后序列和原始序列完全一样是最常见的低级错误。最后检查模型输出。同一个细胞重复推理两次结果是否稳定池化方向是否正确比较的是细胞向量还是最后一层某个位置的状态。很多离奇结论都来自比较错了对象。6.4 关于“shap 图怎么画”的实践建议很多初学者从头开始学 SHAP 时第一步就卡在画图上。实际建议是先不画图先看懂shap_values的维度结构和feature_names对应关系。维度对不上图怎么画都是错的。print(shap_values.shape) # 查看输出维度 print(len(gene_names)) # 查看基因名数量两者如果维度和特征数不一致需要回到池化向量构建阶段修正。画图本身不是难点难点是特征名和特征列有时经过筛选后不再对齐。7. 可复现实验清单从零开始跑通一次完整分析7.1 环境检查清单[ ] Python 版本为 3.9 或 3.10。[ ] PyTorch 已安装GPU 可用性或 CPU 性能已知。[ ] Transformers、datasets、Scanpy 已安装。[ ] Geneformer 模型目录完整包括配置文件、权重、词表。[ ] 单细胞数据基因名格式与模型词表一致。[ ] 已准备一个包含已知细胞类型标签的验证子集。7.2 分析复现清单[ ] 读取 h5ad 文件确认矩阵维度。[ ] 过滤低质量细胞和低表达基因。[ ] 将基因按表达量排序并 token 化。[ ] 用原始输入提取每个细胞的池化向量。[ ] 在验证子集上检查细胞类型聚类结果。[ ] 确定候选基因集合优先选择差异表达和通路核心基因。[ ] 对每个候选基因构造虚拟敲除输入。[ ] 计算扰动前后细胞向量差异分数。[ ] 按细胞类型聚合扰动分数。[ ] 对细胞分类器应用 SHAP 解释。[ ] 汇总 SHAP 排名和扰动分数排名。[ ] 选择排名靠前基因进入真实实验验证计划。7.3 学习环境和生产环境差异如果只是为了学习建议使用公开数据集的一个子集直接用预训练模型做推理即可。数据量控制在几千个细胞单卡即可完成。如果要在正式课题或企业项目中使用还需要补充以下内容将数据预处理和模型推理写到独立脚本中避免在 Jupyter Notebook 里长期运行。记录每次实验的模型版本、数据版本、随机种子和参数。增加完整的日志包括基因丢弃比例、token 长度分布、扰动前后状态差异分布。对候选基因做多次重复扰动输出均值、标准差和置信区间而不是只给一个点估计。将 SHAP 和扰动结果同时保存为 CSV 和图片方便后续归档和复查。如果是敏感临床数据要先完成脱敏和权限审批再进入分析流程。7.4 对新手的学习路线建议如果你是刚接触这部分内容不建议一上来就复现完整流程。建议按下面顺序做几个小练习先学会用 Scanpy 读取单细胞数据并理解表达矩阵结构。再手动实现一个简单的 rank value 编码函数理解排序如何影响 token 序列。然后只对一个细胞对比原始序列和敲除一个基因后模型的输出差异。再做 20 个候选基因的批量扰动观察排序稳定性。最后才加入 SHAP 解释和可视化。练习和实际项目的关系是先把“单一细胞输入输出差异”这条最小验证链路跑通再扩展到全数据集。很多项目失败不是模型不好而是最底层的数据输入早就错了上层扰动和解释自然跟着错。如果之后继续深入学习可以关注这几条扩展方向在自有数据集上对 Geneformer 做领域自适应微调提升特定组织或疾病场景下的表现。把虚拟扰动结果和 CRISPR 公共数据库做定量比较评估哪些预测能被真实实验复现。尝试对整个基因集做分组扰动模拟观察核心转录因子、信号通路下游基因之间的协同影响。将 SHAP 值与差异表达、共表达网络和表型关联打分共同构建一个多证据融合的候选基因排序方式。Geneformer 虚拟扰动分析不是一套能一次跑完就出最终结论的现成工具它更像是一个假设生成框架。合理的使用方式是把模型扰动结果、SHAP 解释和实验验证组合在一起让机器学习帮你缩小候选基因范围再由真实生物学检验结论。这样既能发挥深度学习方法在单细胞数据上的整合能力也不会把模型的预测结果误当成实验事实。对于刚入门的读者最重要的事情不是追新模型而是先把数据格式、token 化、输入输出验证这层基本功打扎实后续无论模型怎么迭代这套分析框架都能复用。
返回列表