ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

Geneformer虚拟扰动分析:AI驱动单细胞基因功能预测

Geneformer虚拟扰动分析:AI驱动单细胞基因功能预测 单细胞测序已经积累了大量数据但很多分析还停留在“描述细胞是什么状态”的阶段。大家更想知道的是如果把某个基因敲掉细胞会往哪个方向变化这个基因对细胞命运的影响到底有多大以前这类问题只能靠真实实验去做成本高、周期长而且很难批量验证。Geneformer这类单细胞基础模型出现之后我们可以在计算机里做虚拟扰动预测先筛出值得做湿实验验证的基因再回到实验中去确认。本文就是一套围绕 Geneformer 虚拟扰动分析的图文教程包含 AI 扰动模型设计、虚拟基因敲除流程、机器学习 SHAP 归因分析和可视化方法希望给正在做单细胞数据挖掘的读者一条可以直接上手的路径。我最近也把整个流程整理成了视频版本图文版把代码、原理和踩坑点展开写清楚。文章涉及的代码比较多建议收藏后按章节逐步跑实验。1. 背景与核心概念1.1 什么是 GeneformerGeneformer 是一个基于 Transformer 架构的单细胞转录组基础模型由 Gladstone Institutes 团队提出2023 年发表于 Nature。它的核心思路非常直接把单细胞 RNA 测序数据中每个细胞的基因表达情况按表达量从高到低排序构造成一段“基因序列”然后像自然语言处理里的 BERT 一样去做自监督预训练。传统单细胞分析侧重聚类、差异表达、拟时序分析这些方法擅长描述数据但难以建模基因之间的调控关系。Geneformer 不一样它在数千万个单细胞数据上学到了基因共表达模式、细胞类型特异性以及部分调控网络信息。简单说它不是在“看细胞”而是在“读基因之间的语法”。在实际项目中Geneformer 主要有三类用途细胞类型分类利用预训练模型提取的特征做下游分类。基因调控网络推断通过注意力机制分析基因之间的关联。虚拟扰动分析在模型输入层面模拟基因过表达或敲除预测细胞状态变化。本文重点讲第三类。1.2 什么是 AI 扰动模型与虚拟扰动分析虚拟扰动分析又叫 in silico perturbation analysis指的是不进行真实实验而是对训练好的模型输入做人为修改模拟某个基因的表达量改变再观察模型预测结果的变化。AI 扰动模型的“扰动”不是直接操作细胞而是操作模型输入。以 Geneformer 为例输入是基因表达排序后的 token 序列。当我们把某个基因对应的 token 掩码掉、删除掉或加强其权重模型内部的注意力分布和最终隐藏状态就会发生改变。通过比较扰动前后的输出我们可以推测这个基因在特定细胞类型中的功能重要性。这种分析方式非常适合做基因筛选。真实 CRISPR 筛选成本高、周期长而且很难覆盖所有细胞类型。AI 扰动模型可以在数分钟内完成成千上万个基因的虚拟筛查帮助研究人员锁定候选基因再用真实实验验证。需要强调的是虚拟扰动是预测不是实验替代。它在机制研究中更多承担“假设生成器”的职责。1.3 什么是虚拟基因敲除虚拟基因敲除virtual knockout是虚拟扰动分析中最常见的一种操作。真实实验中的基因敲除是通过 CRISPR/Cas9 等工具让某个基因不表达而虚拟敲除是在模型输入中让某个基因“不出现”或“被掩码”。在 Geneformer 的输入表达序列中每个基因的 token 是按表达量排序出现的。虚拟敲除某个基因时通常有两种做法将该基因的 token 替换为掩码 token让模型去预测这个位置最可能的基因。将该基因的 token 从序列中移除保持其他基因相对顺序不变观察模型输出变化。通过对比敲除前后的细胞表示cell embedding或下游预测结果可以判断该基因对细胞状态的贡献程度。这也是“AI 扰动模型”中比较标准的一种操作。1.4 SHAP 在基因分析与扰动解释中的作用SHAPSHapley Additive exPlanations是一种基于博弈论的机器学习可解释性方法。它把每个输入特征视为一个“玩家”通过计算 Shapley 值来量化每个特征对模型预测结果的边际贡献。在基因分析场景中SHAP 可以回答两类问题对于某一个细胞的预测哪些基因起了关键作用对于虚拟敲除引起的状态变化哪些基因的贡献最大SHAP 值是加性的正负号能表示促进或抑制作用非常适合基因表达这类连续性特征。结合基因名做可视化就得到了我们常说的 SHAP 图。需要注意SHAP 分析和虚拟敲除是两个层面的概念。虚拟敲除是修改输入观察输出变化SHAP 是追溯已有预测结果的归因。两者可以互补使用先用虚拟敲除产生扰动假设再用 SHAP 解释模型预测的关键驱动基因。2. 核心原理拆解2.1 Geneformer 模型架构与输入构建Geneformer 没有使用传统的基因表达矩阵送入全连接网络而是把每个细胞建模成一段 token 序列。具体流程是对一个细胞中所有基因按表达量从高到低排序。构建一个基因词典每个基因分配一个唯一 ID。把排序后的基因映射成 ID 序列作为模型的输入 token。这种设计的优势在于保留基因表达排序信息的同时避免了基因数量对输入维度的影响更重要的是Transformer 可以捕捉序列中不同位置基因之间的依赖关系也就是基因调控的“上下文”。下图是输入构建的最小逻辑原始细胞表达 GeneA 表达量 120 GeneB 表达量 85 GeneC 表达量 40 排序后 token 序列 [GeneA_id, GeneB_id, GeneC_id, ...]Geneformer 在预训练阶段采用掩码语言建模随机遮挡部分基因 token让模型根据上下文重建被遮挡的基因。经过大量单细胞数据训练后模型对“哪些基因倾向于共表达”“哪些基因状态改变会引发连锁反应”具备了一定的建模能力。2.2 虚拟敲除的两种实现思路在实现虚拟敲除时我建议区分两种操作第一种是 Mask 方式。将目标基因 token 替换为mask_token_id让模型基于上下文重新预测该位置应该出现的基因。如果模型预测出的新基因与原基因差异很大说明该基因的位置上存在较强的调控主导性。第二种是 Delete 方式。直接从序列中移除目标基因 token其他基因的相对顺序保持不变用 padding token 补齐长度。这种方式更贴近真实的“基因不存在”状态。两种方式各有侧重方式模拟效果优点缺点Mask基因缺失且模型要推断缺失内容可以观察模型对缺失基因的补偿预测预测结果受上下文影响大Delete基因从序列中消失更接近真实敲除后的表达谱序列长度和位置偏移会影响模型在实际分析中我一般两种都跑比较结果的一致性。如果某个基因在两种扰动方式下都导致显著变化那这个基因值得重点关注。2.3 SHAP 归因原理简介SHAP 值源于合作博弈论中的 Shapley 值。对于一个机器学习模型 f某一个特征的贡献值通过计算该特征加入不同特征组合时带来的边际贡献均值来获得。在基因分析中如果我们拿一个细胞类型分类器作为解释对象那么 SHAP 可以告诉我们某个基因表达量升高或降低如何推动模型将细胞预测为某一类型。针对 Geneformer 场景SHAP 通常用在以下环节对模型隐藏层输出接入一个简单分类头如细胞类型分类器。对待解释样本进行 SHAP 归因得到每个输入 token 对应的贡献值。将 token 的归因值映射回基因名。SHAP 支持多种解释器例如TreeExplainer适合树模型DeepExplainer和GradientExplainer适合神经网络。Geneformer 是 Transformer通常使用梯度类解释器。3. 环境准备与数据格式3.1 安装依赖Geneformer 基于 Hugging Face Transformers 框架推荐使用 Linux 服务器环境并配备 NVIDIA GPU。如果只学习代码流程CPU 也可以跑通但模型加载和推理会非常慢。建议通过 conda 创建一个独立环境conda create -n geneformer python3.9 conda activate geneformer pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu118 pip install transformers datasets pip install geneformer pip install shap pip install scanpy anndata说明geneformer包版本更新较快具体 API 以项目官方说明为准。上方的cu118是 PyTorch CUDA 版本需要根据你本机显卡驱动版本调整。3.2 数据格式准备Geneformer 期望的输入数据不是普通 CSV 表达矩阵而是 tokenize 后的数据格式。整体转换步骤为读取单细胞表达矩阵格式一般为h5ad。使用基因词典对基因名做映射。对每个细胞按表达量降序排序得到 token 序列。将序列写入 Hugging FaceDataset格式。基因命名建议统一使用 Ensembl ID因为它标准、唯一在不同版本之间稳定性好。Symbol 名称容易出现别名冲突。核心转换逻辑如下import scanpy as sc import anndata as ad adata sc.read_h5ad(your_data.h5ad) # 假设基因名是 Symbol需要先转换为 Ensembl ID # 这一步取决于你的数据来源建议提前处理 gene_to_id { GAPDH: ENSG00000111640, # ... } # 构建表达排序后的基因序列 # 这里省略了完整实现核心是 # 对每个细胞取表达量前 N 个基因映射为 token id这里强调一下每个项目中基因表示的标准化非常关键否则会出现训练和推理时基因名对不上的问题。4. 完整实战案例基于 Geneformer 的虚拟基因敲除与 SHAP 分析这一节我们分 5 步完成一个最小可运行的流程。说明以下代码以理解流程为主。实际运行时Geneformer 预训练权重需要从 Hugging Face Hub 下载建议先下载到本地再做离线加载。4.1 加载预训练模型import torch from geneformer import GeneformerTokenizer, GeneformerForMaskedLM # 官方模型 ID model_name ctheodoris/Geneformer # 如果网络条件允许直接加载 tokenizer GeneformerTokenizer.from_pretrained(model_name) model GeneformerForMaskedLM.from_pretrained(model_name) # 切换到评估模式并移动到 GPU device torch.device(cuda if torch.cuda.is_available() else cpu) model model.to(device) model.eval()如果服务器无法直接访问 Hugging Face Hub需要在一台有网络的机器上提前下载模型目录然后通过本地路径加载model GeneformerForMaskedLM.from_pretrained(/data/geneformer_model)这个方法也适用于公司内网环境。4.2 构造示例输入我们没有在示例中使用大型数据集而是构造一个模拟的“细胞表达序列”用来演示虚拟敲除流程。# 假设某个细胞的表达排序前 20 个基因对应的 token id input_ids torch.tensor( [[101, 502, 830, 1201, 332, 990, 150, 400, 760, 203, 880, 610, 300, 905, 420, 555, 130, 999, 200, 505]], dtypetorch.long, devicedevice ) # 构造注意力掩码长度与 input_ids 一致 attention_mask torch.ones_like(input_ids, devicedevice)真实项目中这里的input_ids来自对单细胞表达矩阵做 tokenize 后的结果不能直接使用随机数字。上面这段代码只是为了演示模型 API。4.3 定义虚拟敲除函数我们实现 Mask 和 Delete 两种扰动方式。def virtual_knockout( input_ids, target_gene_id, tokenizer, modemask, ): 对目标基因执行虚拟敲除。 参数: input_ids: 细胞的 token 序列 target_gene_id: 目标基因 token id tokenizer: Geneformer tokenizer mode: mask 或 delete 返回: perturbed_ids: 扰动后的 token 序列 mask_positions: 目标基因在序列中的位置 input_ids input_ids.clone() positions (input_ids target_gene_id).nonzero(as_tupleTrue)[1] if len(positions) 0: print(目标基因不在此细胞的表达序列中跳过。) return input_ids, positions if mode mask: # 将目标基因 token 替换为掩码 token input_ids[:, positions] tokenizer.mask_token_id elif mode delete: # 将目标基因 token 删除序列末尾用 padding 补位 mask input_ids ! target_gene_id filtered_ids torch.masked_select(input_ids, mask) pad_len input_ids.shape[1] - filtered_ids.shape[0] pad_ids torch.full((1, pad_len), tokenizer.pad_token_id, dtypetorch.long) filtered_ids filtered_ids.unsqueeze(0) input_ids torch.cat([filtered_ids, pad_ids], dim1) return input_ids, positions这里需要注意tokenizer.pad_token_id和tokenizer.mask_token_id是否存在且正确赋值。如果加载的模型没有设置这两个特殊 token需要手动指定。4.4 执行扰动预测并比较结果扰动完成后我们对原始序列和扰动序列分别做一次前向传播观察目标位置预测的新基因。def predict_masked_gene(model, input_ids, attention_mask, mask_positions): with torch.no_grad(): outputs model(input_idsinput_ids, attention_maskattention_mask) logits outputs.logits # 形状: [batch, seq_len, vocab_size] predicted_ids [] for pos in mask_positions.tolist(): pos_logits logits[0, pos, :] pred_id torch.argmax(pos_logits).item() predicted_ids.append(pred_id) return predicted_ids # 示例对第 5 个基因做 Mask 式虚拟敲除 target_gene_id input_ids[0, 4].item() perturbed_ids, positions virtual_knockout( input_ids, target_gene_id, tokenizer, modemask ) predicted_ids predict_masked_gene( model, perturbed_ids, attention_mask, positions ) # 将 token id 映射回基因名 # tokenizer 需要提供 id2gene 字典不同版本名称不同 # 这里假设存在 tokenizer.id2gene if hasattr(tokenizer, id2gene): original_gene tokenizer.id2gene.get(target_gene_id, unknown) predicted_gene tokenizer.id2gene.get(predicted_ids[0], unknown) print(f原始基因: {original_gene}) print(f模型预测替代基因: {predicted_gene})如果模型预测的新基因与原始基因不一致说明在模型看来该位置的基因缺失后可以由其他基因补偿或该位置对后续细胞状态的影响不大。如果预测结果仍然维持原基因说明该基因的表达信号在上文语境中非常强。4.5 使用 SHAP 进行基因归因分析真实项目里SHAP 更适合解释“下游分类任务”而不是直接解释掩码语言模型的输出。因此我们在 Geneformer 顶部增加一个细胞类型分类头用 SHAP 对这个分类器做归因。假设我们有一个训练好的分类器classifier输入是基因序列输出是各类别概率。import shap import numpy as np # 1. 定义包装函数 def model_predict(token_ids): 输入: token id 序列 输出: 类别概率 token_ids torch.tensor(token_ids, dtypetorch.long, devicedevice) with torch.no_grad(): outputs model(input_idstoken_ids) cell_embedding outputs.last_hidden_state.mean(dim1) probs torch.softmax(classifier(cell_embedding), dim-1) return probs.cpu().numpy() # 2. 选择背景数据 background input_ids.cpu().numpy().repeat(10, axis0) # 实际项目中background 应该从训练集中随机选择若干样本 # 3. 使用 GradientExplainer explainer shap.GradientExplainer(model_predict, background) # 4. 对单个样本进行解释 shap_values explainer.shap_values(input_ids.cpu().numpy()) # 5. 将 SHAP 值映射到基因名 gene_names [tokenizer.id2gene.get(i, ftoken_{i}) for i in input_ids[0].tolist()]注意GradientExplainer会计算输入 token 对输出概率的梯度近似因此解读时更关注相对大小而非绝对数值。4.6 绘制 SHAP 图针对细胞类型二分类问题我们可以画出传统的 SHAP summary 图和 bar 图。import matplotlib.pyplot as plt # 方式一summary_plot适合看全局特征影响 shap.summary_plot( shap_values[1], # 二分类中解释类别 1 featuresinput_ids.cpu().numpy(), feature_namesgene_names, max_display15, showFalse ) plt.tight_layout() plt.savefig(shap_summary.png, dpi150) plt.show() # 方式二bar plot 展示平均绝对 SHAP 值 shap.summary_plot( shap_values[1], featuresinput_ids.cpu().numpy(), feature_namesgene_names, plot_typebar, max_display15, showFalse ) plt.tight_layout() plt.savefig(shap_bar.png, dpi150) plt.show()如果使用新版shap也可以使用shap.plots.bar和shap.plots.beeswarm。两种 API 绘图逻辑类似。绘制完成后优先关注 SHAP 绝对值排名靠前的基因这些基因通常是模型判断细胞类型的关键驱动因子。把它们和虚拟敲除结果中的敏感基因取交集会得到更有价值的候选基因列表。5. 常见问题与排查思路实际操作中常见的报错和问题集中在模型加载、基因 ID 匹配、显存和 SHAP 计算效率几个方面。这里整理了一份排查表问题现象常见原因解决思路加载模型时网络超时Hugging Face Hub 网络不稳定在有外网的机器上下载模型权重然后本地离线加载加载模型时缺少依赖包transformers 或 tokenizers 版本不匹配先升级pip install --upgrade transformers datasets tokenizersGPU 显存不足输入序列过长或 batch 过大减小max_length将 batch size 设为 1开启梯度检查点目标基因不在词典中基因名与模型训练时的基因命名不一致统一使用 Ensembl ID检查基因名版本虚拟敲除前后无变化目标基因没有出现在序列中或模型对该细胞不敏感打印positions确认目标基因位置尝试改变扰动方式SHAP 计算速度极慢GradientExplainer 本身计算量较大减少 background 样本数量减少待解释的 token 数量SHAP 图报错Invalid feature shape输入维度与特征名长度不一致确保feature_names长度等于输入 token 数模型输出 logits 形状异常预训练模型和微调模型接口不一致打印outputs.logits.shape确认是[batch, seq_len, vocab]排查时建议写一个“最小复现代码”只用一条细胞数据跑通完整流程再逐步扩展数据量。6. 最佳实践与工程建议6.1 明确扰动分析的边界虚拟基因敲除是预测工具不是实验验证。所有虚拟扰动结果只能作为筛选依据在发表论文或指导临床研究时必须结合真实实验验证。写作和汇报时要使用“模型预测”“虚拟筛选”等表述避免把 in silico 结果描述为实验事实。6.2 统一基因标识体系建议从数据准备阶段就统一使用 Ensembl ID。Symbol 名称虽然可读性好但存在一对多、多对一的情况容易造成模型输入错位。如果必须使用 Symbol务必提供可靠的基因名映射表并做去重。6.3 多次扰动与结果稳定性单个细胞的虚拟敲除结果可能受输入排序影响。建议对同一基因在不同细胞中执行多次扰动并统计结果的稳定性。# 伪代码批量统计多次扰动 result {} for cell_index in range(num_cells): original_genes get_cell_genes(cell_index) for target_gene in candidate_genes: change_score run_virtual_knockout(original_genes, target_gene) result[target_gene].append(change_score) # 取平均和标准差筛选高置信度基因6.4 SHAP 与虚拟敲除结合使用单独做 SHAP 解释有时会得到大量基因无法聚焦。建议先用虚拟敲除缩小候选范围再对候选基因做 SHAP 归因解释模型为什么会把某个细胞预测为某个状态。两者结合时重点关注在两类分析中表现一致的基因。6.5 计算资源管理Geneformer 模型权重较大加载和推理都需要一定显存。建议线上训练用 GPU 实例线下用 CPU 做代码调试。大批量扰动预测时使用torch.no_grad()减少显存占用。可以考虑批量推理后再统一保存结果不要每做一次扰动就保存一次中间状态。6.6 结果可视化规范化SHAP 图和扰动结果图要统一风格建议在分析脚本中固定随机种子、字体和配色方便组内合作时对图进行对比。import random import numpy as np import torch SEED 42 random.seed(SEED) np.random.seed(SEED) torch.manual_seed(SEED) torch.cuda.manual_seed_all(SEED)7. 总结与下一步到这一步你应该已经理解了 Geneformer 虚拟扰动分析的整体流程从单细胞表达谱构建基因 token 序列通过掩码或删除方式模拟基因敲除观察模型输出变化再用 SHAP 归因定位关键驱动基因。这套流程特别适合做三件事大批量候选基因的排序筛选。特定细胞类型的关键调控基因发现。与真实敲除实验数据做对比验证。如果接下来想深入可以从三个方向继续第一阅读 Geneformer 原始论文及其后续版本理解掩码语言建模在基因数据上的设计细节。第二把自己的h5ad数据完整转成 Geneformer 输入格式跑通单细胞扰动预测的 pileline。第三学习 SHAP 更多可视化方式例如 force plot 和 waterfall plot提升单个细胞预测的解释说服力。虚拟扰动分析在单细胞领域仍然处于快速发展阶段模型预测准确度、跨物种迁移能力、与真实实验的一致性都还有不少提升空间。对做生信分析和计算生物学的人来说现在正是掌握这套方法的好时机。
返回列表