拓冰建站拓冰建站
首页 / 资讯中心 / 正文

单细胞响应预测:迈向虚拟细胞的关键一步

近年来单细胞测序技术让研究者能够看清细胞与细胞之间的差异但“看清”并不等于“算准”。尤其在药物研发早期我们迫切需要回答一个很难的问题一种从未见过的药物用在某种疾病细胞上单细胞层面的转录组会如何变化这一预测难题如果能够被计算模型解决将极大缩短药物筛选周期、降低实验成本。近期发表在《Nature Machine Intelligence》上的一项工作把目光投向这个方向提出了一种迈向“真正虚拟细胞”的新框架。这篇文章不会去复述论文里的每个公式而是拆解这类框架的设计思路并用一个最小可运行的示例帮助你理解单细胞响应预测模型到底在做什么。1. 背景与核心概念1.1 单细胞响应预测从“观测”走向“预测”传统转录组测序通常得到的是组织或者细胞群体的平均表达信号反映不出细胞异质性。单细胞 RNA 测序scRNA-seq则把分辨率推进到单个细胞能够看到同一组织中不同细胞状态构成的“图谱”例如肿瘤微环境里的免疫细胞亚群、耐药细胞亚群等。但有一个现实问题测序只能观测给药之前的静态状态。真正想知道的是“如果我把某种药加到这群细胞里哪些细胞亚群会发生变化、哪些基因会上调、哪些通路会被激活”。对这种“扰动后的响应”进行预测就是单细胞响应预测任务。用一句话概括就是输入一个细胞当前的基因表达状态 一个候选药物的分子结构输出给药后的细胞状态或者一组差异表达基因或者一个药物敏感性标签。如果模型能够可靠地完成这个映射那它就不再只是一个数据处理工具而是一个可以“做实验”的计算仿真器也就是虚拟细胞的雏形。1.2 “未知药物”为什么难药物响应预测最困难的地方并不是对已知药物做拟合而是对未知药物做到可靠预测。这里要区分两种“未知”训练集中没有出现过这种药但存在结构类似的药。这种情况下模型可以借助分子结构相似性完成一定程度的迁移。训练集中连结构类别都非常少见甚至药物分子所处的化学空间和训练数据差异很大。这种外推本质上非常困难模型容易退化成“对已知药物响应取平均”。细胞层面的响应又有更多的复杂性。同样是 500 个细胞可能分属 5 个不同的细胞状态对同一个药物的响应方向和强度都不相同。平均表达信号掩盖了这种异质性所以只在细胞系级别做预测到了单细胞层面往往失真。因此单细胞响应预测框架需要同时处理好两个层次的泛化问题细胞状态层面的泛化对不同分化状态、不同微环境下的细胞都能判断。药物结构层面的泛化对训练数据里没有出现过的分子都能给出合理响应。这也是《Nature Machine Intelligence》这一新框架试图解决的核心难题。1.3 什么是“虚拟细胞”“虚拟细胞”目前没有一个完全统一的定义但可以理解为一个可计算的细胞数学模型。这个模型能够根据细胞的基因表达状态和外部扰动输入推演细胞走向。你可以把虚拟细胞想象成一台“细胞模拟器”接收一份细胞表达谱相当于获取细胞当前状态接收一个药物分子相当于给细胞施加扰动运行模型输出扰动后的表达状态或细胞命运变化再通过通路富集等方法解释模型输出结果。单细胞响应预测是虚拟细胞的核心能力之一。没有这个能力虚拟细胞只能描述“是什么”不能回答“如果怎样会怎样”。所以当一个新框架声称“迈向真正的虚拟细胞”它本质上是在说我们开始尝试用学习模型去逼近细胞对扰动的响应函数。2. 框架的整体设计思路2.1 把响应预测拆成三要素无论论文中的框架写得多么复杂单细胞响应预测模型的骨架通常都可以拆成三个部分细胞编码器把高维基因表达向量映射成一个紧凑的细胞表示。药物编码器把药物分子结构映射成一个可参与计算的药物表示。响应预测头融合细胞表示和药物表示输出响应结果。这种“双编码器 融合预测头”的结构最早常见于药物敏感性预测和化学-蛋白质相互作用预测后来被迁移到单细胞响应预测中。它的好处是模块清晰细胞编码器和药物编码器可以独立预训练也可以分别替换成更先进的模型。细胞表达向量 ── 细胞编码器 ── 细胞表示 ─┐ ├─ 融合 ─ 响应预测头 ─ 预测结果 药物分子 ───── 药物编码器 ── 药物表示 ─┘2.2 细胞表示基因表达如何向量化单细胞测序得到的原始数据是基因表达计数矩阵行是细胞列是基因。一个常见处理流程是质量过滤去掉低质量细胞和低表达基因。归一化对每个细胞的测序深度做归一化常用log1p(CPM/10000)一类变换。特征选择选取高变基因HVG通常几千个减少噪声和维度。批次校正合并多个样本时校正测序批次效应。经过这些处理后每个细胞得到一个几千维的向量。细胞编码器负责从这个向量中提取更有信息量的表示可能是一个 128 维或 256 维的向量。在很多单细胞预训练模型里细胞编码器已经可以学习到细胞类型、细胞周期、应激状态等隐含特征。这层表示如果足够好下游的药物响应预测头就不需要重新学习太多细胞生物学知识只需要学习“药物如何改变这个细胞状态”。2.3 药物表示从 SMILES 到可计算向量药物分子通常用 SMILES 字符串表示例如阿司匹林的 SMILES 是CC(O)Oc1ccccc1C(O)O。模型不能直接处理字符串所以需要把分子编码成向量。常见方案有三种分子指纹用 RDKit 将分子转换为固定长度的二进制向量例如 Morgan 指纹。优点是简单、无需训练缺点是会丢失部分三维结构信息。分子图神经网络把原子看成节点、化学键看成边用 GNN 学习分子图表示。表达能力更强是当前主流做法之一。分子描述符计算分子量、LogP、氢键供体数等理化性质适合作为辅助特征。在响应预测任务中药物编码器的目标不是判断“药物是否有效”而是学到药物对细胞状态的影响方式。因此药物表示最好保留结构-活性关系信息而不是只输出一个药物有效性标签。2.4 响应头预测差异表达还是二分类常见的药物响应标签有三种形式对应三种不同的预测目标预测目标输出内容典型损失函数信息量二分类标签敏感/耐药BCE低连续IC50/疗效值一个数值MSE中全转录组扰动向量每个基因的表达变化MSE/L1高单细胞响应预测框架中最有信息量的做法是预测“差异表达向量”也就是给药后每个基因相对于对照的变化量。虽然这种任务更难训练但它能告诉研究者哪些基因受影响、哪些通路被激活而不仅仅是给一个有效/无效的判断。近期的框架通常使用生成式思路给定细胞状态和药物表示预测一个 delta 表达向量再叠加到初始表达上。这样做的好处是模型既能用于预测药物治疗后状态也能用于虚拟筛选挑选出能够把疾病细胞状态“拉回”正常状态的候选药物。3. 环境准备与数据组织3.1 依赖安装本文的示例代码以 Python PyTorch 为基础核心依赖如下Python 3.9 或更高版本PyTorch 2.xnumpy、pandas、scikit-learnRDKit用于真实药物分子指纹转换建议使用 conda 创建独立环境conda create -n virtual-cell python3.9 -y conda activate virtual-cell pip install torch numpy pandas scikit-learn pip install rdkitRDKit 在部分平台上需要较长时间安装如果安装失败可以尝试使用conda install -c conda-forge rdkit。不过本文的最小示例为了降低门槛会直接使用随机生成的指纹矩阵作为药物特征你完全可以先不安装 RDKit等需要处理真实 SMILES 时再安装。3.2 数据约定为了便于理解我们约定如下数据结构base_expr形状为(N, G)的矩阵表示 N 个细胞在给药前的基因表达向量。drug_fp形状为(K, FP_DIM)的矩阵表示 K 个药物的分子指纹二进制向量。delta_target形状为(N, G)的矩阵表示每个细胞在对应药物处理后的真实差异表达向量。真实场景中delta_target需要依靠配对对照实验获得例如同一批细胞分成对照组和给药组测序后计算差异表达。这是非常昂贵的数据所以模型才需要学会泛化尽量降低对新药做实验的需求。3.3 项目结构为了方便维护我们把代码拆分成模块virtual-cell-demo/ ├── data.py # 模拟数据生成与数据集定义 ├── models.py # 细胞编码器、药物编码器、响应预测模型 ├── train.py # 训练与评估脚本 └── utils.py # 评估指标工具如果你只是在本地跑通流程也可以把所有代码写进一个main.py。这里拆分模块主要是为了符合工程化习惯也为后续替换真实数据留出结构空间。4. 实战一个最小响应预测框架这一节我们会实现一个简化版框架。请注意这是一个教学示例用于说明单细胞响应预测的训练链路而不是对《Nature Machine Intelligence》论文的复现。真实模型会复杂得多。4.1 生成模拟数据我们先用随机方式生成一份“模拟单细胞响应数据”。为了让“未知药物”的评估有一点意义我们给药物-基因响应施加一个线性低秩假设# data.py import torch from torch.utils.data import Dataset G 2000 # 基因数量 N_TRAIN 512 # 训练细胞数量 N_TEST 128 # 测试细胞数量 K_TRAIN 16 # 训练集已知药物数量 K_TEST 4 # 测试集“未知药物”数量 FP_DIM 256 # 药物指纹位点数 SEED 42 torch.manual_seed(SEED) def generate_simulated_data(): # 1. 生成药物指纹模拟二值分子指纹 n_drugs K_TRAIN K_TEST drug_fps torch.randint(0, 2, (n_drugs, FP_DIM)).float() # 2. 构造真实的药物-基因作用矩阵 # 这里假设响应机制是线性的delta_expr drug_fp W_true W_true torch.randn(FP_DIM, G) * 0.05 # 3. 生成训练细胞和测试细胞的基线表达 base_train torch.randn(N_TRAIN, G) * 0.5 base_test torch.randn(N_TEST, G) * 0.5 # 4. 为每个细胞随机分配一个药物 train_drug_ids torch.randint(0, K_TRAIN, (N_TRAIN,)) test_drug_ids torch.randint(K_TRAIN, K_TRAIN K_TEST, (N_TEST,)) # 5. 计算真实差异表达加入少量噪声模拟生物学和测序噪声 train_delta drug_fps[train_drug_ids] W_true torch.randn(N_TRAIN, G) * 0.02 test_delta drug_fps[test_drug_ids] W_true torch.randn(N_TEST, G) * 0.02 return { base_train: base_train, base_test: base_test, drug_fps: drug_fps, train_drug_ids: train_drug_ids, test_drug_ids: test_drug_ids, train_delta: train_delta, test_delta: test_delta, }这里值得说明的是“测试集包含未知药物”这一设计。test_drug_ids是从K_TRAIN到K_TRAIN K_TEST - 1之间的药物这些药物指纹完全没有出现在训练集中。模型在训练时只能通过药物指纹向量本身来推断药物的作用规律这正是未知药物泛化的核心场景。4.2 数据集封装使用 PyTorch 的Dataset封装训练数据# data.py 续 class ResponseDataset(Dataset): def __init__(self, base_expr, drug_fps, drug_ids, delta_target): self.base_expr base_expr self.drug_vec drug_fps[drug_ids] self.delta_target delta_target def __len__(self): return len(self.base_expr) def __getitem__(self, idx): return self.base_expr[idx], self.drug_vec[idx], self.delta_target[idx]在模型训练过程中我们输入的是给药前的基线表达和药物指纹输出的目标是差异表达delta。这里的base_expr并不直接参与损失计算但它是细胞编码器的输入相当于告诉了模型“这个细胞原本处于什么状态”。4.3 构建编码器与预测模型模型部分包括细胞编码器、药物编码器和响应预测头。核心代码如下# models.py import torch import torch.nn as nn class CellEncoder(nn.Module): 将高维基因表达向量压缩为低维细胞表示。 def __init__(self, input_dim, hidden_dim256): super().__init__() self.net nn.Sequential( nn.Linear(input_dim, hidden_dim), nn.LayerNorm(hidden_dim), nn.ReLU(), nn.Dropout(0.1), ) def forward(self, cell_expr): return self.net(cell_expr) class DrugEncoder(nn.Module): 将药物分子指纹映射为低维药物表示。 def __init__(self, fp_dim, hidden_dim256): super().__init__() self.net nn.Sequential( nn.Linear(fp_dim, hidden_dim), nn.ReLU(), nn.Dropout(0.1), ) def forward(self, drug_fp): return self.net(drug_fp) class ResponsePredictor(nn.Module): 双编码器 融合预测头。 输入基线细胞表达、药物指纹。 输出预测的差异表达向量。 def __init__(self, gene_dim, fp_dim, hidden_dim256): super().__init__() self.cell_encoder CellEncoder(gene_dim, hidden_dim) self.drug_encoder DrugEncoder(fp_dim, hidden_dim) self.fusion nn.Sequential( nn.Linear(hidden_dim * 2, hidden_dim), nn.ReLU(), nn.Linear(hidden_dim, gene_dim), ) def forward(self, cell_expr, drug_fp): h_cell self.cell_encoder(cell_expr) h_drug self.drug_encoder(drug_fp) h torch.cat([h_cell, h_drug], dim-1) delta_pred self.fusion(h) return delta_pred模型结构里有一点值得注意细胞表达维度是 2000这个数字在真实场景中可能是几万个基因直接输入全连接网络会带来严重的参数爆炸。所以在真实项目中通常会先做高变基因筛选或者用自编码器/预训练模型把表达向量压缩到几百维再作为细胞编码器的输入。4.4 训练脚本训练逻辑本身并不复杂重点是学习率、损失函数和数据划分# train.py import torch import torch.nn as nn from torch.utils.data import DataLoader from data import generate_simulated_data, ResponseDataset, N_TRAIN, N_TEST, G, FP_DIM from models import ResponsePredictor def pearson_r(x, y): 按行计算两个矩阵样本之间的平均皮尔逊相关系数 x x - x.mean(dim1, keepdimTrue) y y - y.mean(dim1, keepdimTrue) xy (x * y).sum(dim1) x_norm torch.sqrt((x * x).sum(dim1)) y_norm torch.sqrt((y * y).sum(dim1)) return (xy / (x_norm * y_norm 1e-8)).mean().item() def main(): data generate_simulated_data() train_dataset ResponseDataset( data[base_train], data[drug_fps], data[train_drug_ids], data[train_delta] ) test_dataset ResponseDataset( data[base_test], data[drug_fps], data[test_drug_ids], data[test_delta] ) train_loader DataLoader(train_dataset, batch_size64, shuffleTrue) test_loader DataLoader(test_dataset, batch_size64, shuffleFalse) model ResponsePredictor(gene_dimG, fp_dimFP_DIM, hidden_dim256) optimizer torch.optim.Adam(model.parameters(), lr1e-3) loss_fn nn.MSELoss() epochs 30 for epoch in range(epochs): model.train() total_loss 0.0 for cell_expr, drug_fp, delta_true in train_loader: delta_pred model(cell_expr, drug_fp) loss loss_fn(delta_pred, delta_true) optimizer.zero_grad() loss.backward() optimizer.step() total_loss loss.item() * len(cell_expr) train_loss total_loss / len(train_dataset) # 验证 model.eval() test_loss 0.0 preds, trues [], [] with torch.no_grad(): for cell_expr, drug_fp, delta_true in test_loader: delta_pred model(cell_expr, drug_fp) loss loss_fn(delta_pred, delta_true) test_loss loss.item() * len(cell_expr) preds.append(delta_pred) trues.append(delta_true) test_loss test_loss / len(test_dataset) preds torch.cat(preds, dim0) trues torch.cat(trues, dim0) r pearson_r(preds, trues) print(fEpoch {epoch1:02d} | train MSE: {train_loss:.4f} | test MSE: {test_loss:.4f} | test Pearson r: {r:.4f}) if __name__ __main__: main()这个示例中的测试集非常关键它里面的药物指纹在训练中从未出现。如果测试 Pearson r 明显高于随机水平说明模型至少学到了某种“从指纹到基因变化”的泛化规律如果测试指标很差通常说明模型只是在记忆训练集里的药物响应模式。4.5 运行与结果说明运行训练脚本python train.py由于数据规模较小CPU 环境也能较快速跑完。预期你会看到训练损失逐渐下降测试损失也呈下降趋势测试集 Pearson r 会收敛到一个中等水平。当然由于模拟数据过于理想化这个结果不能代表真实场景的效果但它完整验证了框架的输入输出链路。真实项目中delta_target来自对比实验计算数据噪声远大于随机模拟基因之间的相关性和调控关系也更为复杂。因此真实工作的重点通常放在如何构建高质量的配对单细胞响应数据如何让细胞编码器具有更好的迁移能力如何设计损失函数让模型关注关键基因而非所有基因的平均误差。4.6 在真实场景中替换 RDKit 分子指纹如果要从模拟数据切换到真实药物数据你可以用 RDKit 把 SMILES 转成二进制指纹from rdkit import Chem from rdkit.Chem import AllChem import numpy as np def smiles_to_fingerprint(smiles, n_bits2048, radius2): mol Chem.MolFromSmiles(smiles) if mol is None: return None fp AllChem.GetMorganFingerprintAsBitVect(mol, radiusradius, nBitsn_bits) return np.array(fp, dtypenp.float32)然后把上面的模拟drug_fps替换成真实药物指纹矩阵即可。需要注意指纹位点数量FP_DIM也要相应修改为 2048。5. 从 Baseline 到“真虚拟细胞”5.1 加入基因调控网络简单模型把基因看作相互独立的输出维度这忽略了基因之间的调控关系。真实细胞中基因表达受转录因子调控基因与基因之间存在共表达模块和调控网络。一个基因的变化会通过调控网络影响下游基因形成级联效应。更接近虚拟细胞的模型应该在输出层或隐空间中加入基因调控关系约束。例如将基因表达输出层替换为图神经网络在基因相互作用图上做消息传递对预测的差异表达向量做可解释性分析看受影响基因是否富集在特定调控模块引入通路先验知识作为正则项约束模型预测结果符合已知生物学通路。5.2 用单细胞预训练模型替代随机初始化的细胞编码器我们示例中的细胞编码器是从零开始训练的。但真实单细胞数据样本量往往不足以让编码器学到稳健的细胞状态表示。近期的趋势是使用单细胞大模型作为底座这些模型在数亿级单细胞转录组上预训练得到能够区分细胞类型、细胞周期状态等基础特征。在这种框架中细胞基础模型参数通常被冻结或轻量微调药物编码器单独训练响应预测头则在较少的配对抗动数据上训练。这样做的好处是模型对细胞状态的初值判断更准确对稀有细胞类型也有更好的泛化能力。5.3 多任务预测与多模态整合单细胞响应不只是基因表达变化还包括染色质可及性、蛋白质表达、细胞形态变化等。虚拟细胞最终要整合多模态信息但目前最重要也最成熟的数据仍然是转录组。多任务学习也是一种有效路径同一个模型同时预测基因表达变化、细胞增殖抑制、细胞凋亡比例等多个指标。这些任务共享细胞表示和药物表示可以用多任务损失联合训练通常要比单独训练一个表达变化任务更稳定。5.4 需要理性看待的边界即使论文使用了很强的模型和大量的数据单细胞响应预测仍然存在不可忽视的边界外推能力有限训练数据覆盖的细胞类型和药物类别越窄对新场景的预测可信度越低。批次效应干扰不同实验室、不同测序平台的数据差异可能大于生物学差异。验证成本高计算预测最终需要湿实验验证不能因为“模型预测敏感”就直接推进临床。虚拟细胞是一个值得竞逐的长期目标但现阶段更应该把它定位为“实验假设生成器”而不是“实验结果替代品”。一个预测结果只有在后续实验中不断被检验和修正模型本身才会往真实虚拟细胞逼近。6. 常见问题与排查思路在实际搭建单细胞响应预测框架时你会遇到各种问题。下面整理几个高频问题。问题现象常见原因解决思路训练损失下降验证集指标很差过拟合模型记住了训练药物响应增加正则化、扩大数据量、使用按药物划分的交叉验证测试集“未知药物”预测结果接近随机药物表示信息量不足或数据划分不当使用更丰富的分子表示如预训练分子向量或图神经网络模型输出所有基因几乎不变差异表达信号稀疏MSE 被大量未变化基因主导改用加权损失重点关注显著变化基因训练不收敛损失震荡学习率过高、目标尺度差异大降低学习率对目标做标准化或对梯度做裁剪细胞输入特征维度差异不一致训练和测试使用不同基因集合统一基因列表缺失基因补 0同一药物同时出现在训练集和验证集验证指标虚高数据泄露按照药物 ID 而不是细胞样本 ID 划分数据集真实多批次数据合并后效果差批次效应未校正在预处理阶段加入批次校正或把批次作为协变量输入模型这里最容易被忽略的是“数据泄露”。如果训练集和验证集来自同一个药模型完全可以通过记忆药物 ID 而不是学习泛化规律来取得高分。这在单细胞响应预测里尤其需要警惕因为验证集中不同样本可能来自同一个给药实验。好的做法是始终按药物进行分组划分确保验证集药物 ID 从未在训练集出现。7. 最佳实践与工程建议最后这部分我会从工程落地角度给出一些具体建议帮助你在自己的项目中少走弯路。第一设计清晰的模块边界。细胞编码器、药物编码器、融合层、输出头要能够独立替换。这样当更好的单细胞预训练模型出现时你不需要重写整个框架只需要换掉对应模块。第二把数据版本管理当作一等公民。单细胞响应数据涉及表达矩阵、药物注释、实验批次、预处理参数等多个层面非常容易混乱。建议为每个数据集打上版本号并记录完整的预处理脚本和参数。第三评估指标不能只用 MSE。建议同时关注显著变化基因的召回率预测差异表达方向是否与真实一致Top 变化基因的重叠程度通路层面的富集一致性。这些指标比单一平均误差更能反映生物学意义。第四重视不确定性估计。对于未知药物模型除了给出预测响应最好同时给出置信区间。实现上可以用 MC Dropout 或者训练多个模型组成集成输出预测均值和方差。这样实验人员能知道哪些预测值得验证哪些预测可能不太可靠。第五训练时把“未知药物泛化”写进验证流程。不要只报告一个随机划分的测试集结果而应该报告“已知药物响应预测”和“未知药物响应预测”两组指标。这两组指标的差距才是框架真实泛化能力的体现。第六关注可解释性。响应预测模型不能只是个黑盒。建议在输出后对预测的差异表达基因做通路富集分析例如使用 GSEA 或富集分析工具检查模型预测到的通路是否与药物已知的机制一致。一个预测结果如果通路层面完全不合理即使数值指标看起来很好也不能用于实际决策。第七计算资源要合理规划。单细胞表达矩阵通常很大如果不做特征筛选直接在全部基因上训练显存和训练时间都会非常可观。实际工程中通常先选 3000 到 5000 个高变基因或者在预训练自编码器的低维隐空间上训练再映射回全基因表达空间。如果你准备在这个方向做进一步探索建议从复现一个公开的单细胞药物响应数据集开始先把数据预处理、数据划分、评估脚本跑通再逐步替换模型组件。算力允许的话可以尝试在预训练单细胞模型基础上做轻量微调观察响应预测指标相比随机初始化的提升幅度。这条路很长但从计算模型的角度看能够提前预测未知药物的单细胞响应本身就值得投入精力去逐步逼近。希望这篇文章能把“虚拟细胞”这个概念落到一个具体可运行的最小框架上也让你在面对单细胞响应预测任务时有一个清晰的起点。文中的示例代码虽然简单但其中的数据划分、编码器设计、评估思路都可以复用到更复杂的框架里。如果你觉得有帮助可以先收藏备用后续也能在此基础上继续扩展自己的模型。
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门