最近,Nature 上出现了一个很有代表性的研究方向:让一个基础模型同时处理蛋白、细胞和肿瘤微环境三类跨尺度生物数据,并在多个独立队列之间构建一张可迁移的“虚拟地图”。这不是单纯的生信工具,而是一次从“单任务模型”向“多模态基础模型”迁移的范式变化。对于做 AI 算法、生物信息、肿瘤研究或者医学图像处理的人来说,这个方向既涉及大模型架构设计,又涉及多模态数据对齐、批次效应消除、跨队列泛化等工程难题。
这篇文章不打算做新闻式复述,而是从工程实践角度做一个系统拆解。我会先解释什么是生物医学基础模型、为什么要跨尺度和跨队列建模;然后梳理整体技术架构;接着给出一个最小可运行的 PyTorch 原型,把蛋白序列、单细胞表达谱和空间位置信息统一到一个对比学习框架中;最后补充常见问题排查思路与工程落地建议。代码部分尽量保持完整,方便对照思路去复现和扩展。
1. 背景与核心概念
1.1 为什么需要生物医学基础模型
传统生物信息学模型通常是“一个任务一个模型”。比如识别细胞类型,就训练一个细胞分类器;预测蛋白功能,就单独训练一个蛋白模型;分析病理切片,又需要另一个图像模型。每个模型都在自己的数据集上独立训练,数据标注成本高,模型之间无法共享知识,遇到新的队列或者新的实验平台,往往需要重新训练和调参。
基础模型则不一样。它先在大量无标注或者弱标注数据上进行预训练,学习到通用的生物学表征,然后通过少量下游数据微调就能适配多个任务。这种“预训练 + 微调”的模式在 NLP 和 CV 领域已经非常成熟,现在逐步被引入生物医学场景。
在肿瘤微环境研究中,数据天然是多模态的:蛋白序列决定分子功能,单细胞转录组反映细胞状态,空间转录组和病理图像则提供组织结构信息。如果能让模型同时学习这些信息,就可以把“分子层面”和“组织层面”连接起来,形成更完整的生物学图景。
1.2 跨尺度数据指的是什么
跨尺度是指数据来自不同的生物学层级,通常包括三个层次:
- 蛋白尺度:氨基酸序列、蛋白结构、蛋白互作网络。
- 细胞尺度:单细胞转录组、蛋白质组、表观基因组,描述细胞类型和状态。
- 组织微环境尺度:空间转录组、病理切片、免疫组化图像,描述细胞在组织中的空间分布和周围环境。
这三个尺度之间存在内在逻辑联系:蛋白表达变化会影响细胞状态,细胞状态变化会改变组织微环境。传统方法往往只关注其中一个尺度,而基础模型的目标是同时建模这些尺度之间的关联。
比如,一个基因突变可能改变蛋白结构,进而影响细胞信号通路,最终改变肿瘤微环境中的免疫细胞浸润程度。这种跨尺度的因果链条很难用单一数据模态捕捉,需要模型具备多模态对齐能力。
1.3 跨队列建模的难点
跨队列是指在不同的患者队列、不同的医院、不同的测序平台之间迁移模型。现实中的生物数据存在严重的批次效应:同样是肿瘤组织样本,不同平台的基因表达量分布可能差异很大;不同实验室的病理切片染色条件也不一样。
这就导致一个在训练集上表现很好的模型,到了新队列上性能会显著下降。基础模型的优势在于,通过大规模预训练学习到更稳健的生物学特征,而不是过拟合到特定平台的噪声上。
所以,这篇文章讨论的“虚拟地图”,本质上是希望构建一个统一的表征空间,让蛋白、细胞和微环境数据都映射到同一个向量空间里。在这个空间里,相似的生物学状态距离相近,不同尺度的数据可以通过向量运算建立关联。
2. 整体技术架构设计
2.1 数据模态与统一表征
构建跨尺度基础模型,首先要解决“数据格式不统一”的问题。蛋白序列是字符串,单细胞表达谱是高维稀疏向量,病理图像是像素矩阵。要让模型同时处理它们,必须为每种模态设计编码器,把原始数据转换成固定维度的向量。
常见的做法是:
- 蛋白序列:使用氨基酸词汇表进行 token 化,通过 ESM 或者 ProtBERT 风格的 transformer 编码器提取序列特征。
- 单细胞表达谱:筛选高变基因后,把表达量向量输入 MLP 或自编码器,得到细胞 embedding。
- 空间转录组 / 病理图像:使用空间位置编码和图像 patch 编码器,得到组织微环境 embedding。
这些编码器的输出向量维度必须一致,比如统一到 256 维或 768 维,才能在同一个空间里做对齐。
2.2 对比学习与多模态对齐
多模态对齐的基础思路是对比学习。核心思想是:来自同一生物学样本的不同模态表示应该相近,来自不同样本的表示应该互相远离。
在实现时,通常构造正样本对和负样本对:
- 正样本对:来自同一细胞或同一组织区域的蛋白序列和表达谱。
- 负样本对:来自不同样本或不同细胞类型的组合。
然后使用 InfoNCE 或 NT-Xent 损失函数拉近正样本、推开负样本。
图神经网络也常被用来建模空间关系。将细胞作为节点,空间距离或表达相似性作为边,通过图卷积或者图注意力机制,让模型学习细胞之间的交互关系。这样,模型不仅能识别单个细胞类型,还能理解细胞在微环境中的组织方式。
2.3 跨队列泛化策略
要让模型在多个队列之间泛化,通常需要组合多种策略:
- 大规模预训练:在多个公开数据集上联合训练,让模型看到更多平台差异。
- 数据增强:对表达谱添加噪声、随机 mask 部分基因,模拟平台差异。
- 对抗域适应:增加一个域判别器,让编码器学习去掉批次信息。
- 标准化:在输入阶段对表达量做标准化处理,减少批次效应。
其中,域对抗训练是工程中比较有效的手段。它会让编码器尽量提取与批次无关的生物学特征,从而提升模型在新队列上的表现。
2.4 虚拟地图的含义
所谓虚拟地图,指的是把高维表征空间理解为一张地图。横纵坐标可以代表不同的生物学状态轴,例如细胞类型轴、功能状态轴、空间位置轴。
在这个地图上,点与点之间的距离代表生物学相似度,路径代表状态转换过程。研究者可以通过地图发现新的细胞亚群、寻找新的生物标志物、预测药物响应。
构建虚拟地图的核心就是训练好编码器,让表征空间具有生物学意义。
3. 环境准备与原型设计
3.1 开发环境
为了便于演示,我们使用 Python 和 PyTorch 构建一个简化原型。整体结构如下:
project/ ├── config.py ├── data/ │ ├── protein_seq.csv │ ├── cell_expr.csv │ └── spatial_pos.csv ├── models/ │ ├── encoders.py │ └── alignment.py ├── train.py └── evaluate.py依赖库包括:
torch>=1.12 pandas numpy scikit-learn transformers版本需要根据你的项目实际情况调整,本文示例以常见环境为例,重点演示配置思路。
3.2 数据准备与模拟数据生成
在真实场景中,数据来自单细胞测序、空间转录组和蛋白数据库。为了让读者能够直接运行代码,我们先生成一组模拟数据,模拟 500 个细胞、每个细胞关联一段蛋白序列、一个表达向量和一个空间坐标。
# file: config.py import torch # 模拟数据参数 NUM_CELLS = 500 NUM_GENES = 1000 NUM_PROTEIN_TOKENS = 128 PROTEIN_VOCAB_SIZE = 30 # 20种氨基酸 + 特殊token EMBED_DIM = 128 # 训练参数 BATCH_SIZE = 32 EPOCHS = 30 LEARNING_RATE = 1e-3 TEMPERATURE = 0.07 # 设备 DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")这里将蛋白序列 token 化后的长度设置为固定 128。真实场景中,蛋白序列长度差异较大,需要做 padding 和 mask。
模拟数据生成脚本如下,它会构造三类张量,保存在同一个数据容器中:
# file: make_data.py import numpy as np import pandas as pd import torch np.random.seed(42) torch.manual_seed(42) # 1. 模拟蛋白序列 token(整数索引) protein_seq = np.random.randint(1, 21, size=(500, 128), dtype=np.int64) # 2. 模拟 500 个细胞的基因表达量(500 x 1000),使用对数化稀疏分布 cell_expr = np.random.poisson(lam=1.0, size=(500, 1000)).astype(np.float32) cell_expr = np.log1p(cell_expr) # 3. 模拟空间坐标(每个细胞在二维空间中的位置) spatial_pos = np.random.rand(500, 2).astype(np.float32) # 4. 模拟细胞类型标签,用于监督评估 cell_types = np.random.choice(["T cell", "B cell", "Macrophage", "Fibroblast"], size=500) np.savez("data/sim.npz", protein_seq=protein_seq, cell_expr=cell_expr, spatial_pos=spatial_pos, cell_types=cell_types) print("模拟数据已生成,保存在 data/sim.npz") print("蛋白序列 shape:", protein_seq.shape) print("表达谱 shape:", cell_expr.shape) print("空间坐标 shape:", spatial_pos.shape)这段代码生成的数据虽然不包含真实生物学信息,但可以完整验证整个训练流程。
4. 核心模型实现
4.1 三模态编码器设计
我们先实现三个编码器,分别处理蛋白序列、表达谱和空间位置。为了让代码简洁,这里都使用较轻量的网络:
- 蛋白编码器:Embedding + 两层 1D 卷积 + Global Max Pooling。
- 细胞编码器:两层全连接,对表达量降维。
- 空间编码器:两层全连接,把空间坐标映射到同样维度。
核心代码如下:
# file: models/encoders.py import torch import torch.nn as nn class ProteinEncoder(nn.Module): """蛋白序列编码器,输入 token 序列,输出固定维度向量。""" def __init__(self, vocab_size, embed_dim, seq_len, hidden_dim=256): super().__init__() self.embedding = nn.Embedding(vocab_size, embed_dim) self.conv1 = nn.Conv1d(embed_dim, hidden_dim, kernel_size=5, padding=2) self.conv2 = nn.Conv1d(hidden_dim, hidden_dim, kernel_size=5, padding=2) self.pool = nn.AdaptiveMaxPool1d(1) self.proj = nn.Linear(hidden_dim, embed_dim) self.gelu = nn.GELU() def forward(self, x): # x: [B, L] x = self.embedding(x) # [B, L, D] x = x.transpose(1, 2) # [B, D, L] x = self.gelu(self.conv1(x)) x = self.gelu(self.conv2(x)) x = self.pool(x).squeeze(-1) # [B, hidden_dim] x = self.proj(x) # [B, D] return x class CellEncoder(nn.Module): """细胞表达谱编码器,输入基因表达向量,输出固定维度向量。""" def __init__(self, n_genes, embed_dim, hidden_dim=256): super().__init__() self.net = nn.Sequential( nn.Linear(n_genes, hidden_dim), nn.BatchNorm1d(hidden_dim), nn.GELU(), nn.Dropout(0.2), nn.Linear(hidden_dim, hidden_dim), nn.GELU(), nn.Linear(hidden_dim, embed_dim), ) def forward(self, x): # x: [B, N_GENES] return self.net(x) class SpatialEncoder(nn.Module): """空间坐标编码器,输入二维坐标,输出固定维度向量。""" def __init__(self, embed_dim, hidden_dim=64): super().__init__() # 增加位置编码,把二维坐标映射到更高维空间 self.net = nn.Sequential( nn.Linear(2, hidden_dim), nn.GELU(), nn.Linear(hidden_dim, hidden_dim), nn.GELU(), nn.Linear(hidden_dim, embed_dim), ) def forward(self, x): # x: [B, 2] return self.net(x)这里的关键点是,三个编码器的输出维度都对齐到embed_dim。如果后续想换用更强的预训练蛋白模型,比如 ESM-2,只需要替换ProteinEncoder,保留输出层即可。
4.2 对比学习模型
有了三个编码器之后,我们需要把它们组合起来,形成多模态对齐模型。在训练阶段,模型会拉近同一细胞的蛋白表征、表达谱表征和空间表征;在推理阶段,三个模态可以互相检索。
组合模型代码如下:
# file: models/alignment.py import torch import torch.nn as nn import torch.nn.functional as F class CrossScaleModel(nn.Module): """连接蛋白、细胞与空间位置的多模态对齐模型。""" def __init__(self, vocab_size, n_genes, embed_dim, seq_len, temperature=0.07): super().__init__() self.protein_encoder = ProteinEncoder( vocab_size, embed_dim, seq_len ) self.cell_encoder = CellEncoder(n_genes, embed_dim) self.spatial_encoder = SpatialEncoder(embed_dim) self.temperature = temperature def forward_protein(self, seq): return F.normalize(self.protein_encoder(seq), dim=-1) def forward_cell(self, expr): return F.normalize(self.cell_encoder(expr), dim=-1) def forward_spatial(self, pos): return F.normalize(self.spatial_encoder(pos), dim=-1) def contrastive_loss(self, seq, expr, pos): """ 计算三模态对比损失: 同一个样本的三个模态互相为正样本对。 """ p = self.forward_protein(seq) # [B, D] c = self.forward_cell(expr) # [B, D] s = self.forward_spatial(pos) # [B, D] # 三组模态两两计算 InfoNCE 损失 loss_pc = self._info_nce(p, c, self.temperature) loss_ps = self._info_nce(p, s, self.temperature) loss_cs = self._info_nce(c, s, self.temperature) return (loss_pc + loss_ps + loss_cs) / 3.0 def _info_nce(self, z1, z2, temperature): """ z1, z2: [B, D] 已经归一化 返回对称的 InfoNCE 损失。 """ logits = z1 @ z2.T / temperature # [B, B] labels = torch.arange(logits.shape[0], device=logits.device) loss = F.cross_entropy(logits, labels) + F.cross_entropy(logits.T, labels) return loss / 2.0在对比损失中,我们把同一样本的蛋白表示和细胞表示作为正样本对,同一样本的细胞表示和空间表示也作为正样本对。这样模型会学到“如果一个细胞具有某种表达状态,它应该与对应的蛋白功能特征和空间位置特征一致”的映射关系。
4.3 训练流程
训练流程分三步:
- 加载模拟数据。
- 构建模型和优化器。
- 循环训练,记录损失变化。
完整代码如下:
# file: train.py import numpy as np import torch import torch.optim as optim from torch.utils.data import DataLoader, TensorDataset from config import * from models.alignment import CrossScaleModel def load_sim_data(): data = np.load("data/sim.npz") protein_seq = torch.LongTensor(data["protein_seq"]) cell_expr = torch.FloatTensor(data["cell_expr"]) spatial_pos = torch.FloatTensor(data["spatial_pos"]) cell_types = data["cell_types"] return protein_seq, cell_expr, spatial_pos, cell_types def main(): # 1. 加载数据 protein_seq, cell_expr, spatial_pos, _ = load_sim_data() dataset = TensorDataset(protein_seq, cell_expr, spatial_pos) loader = DataLoader(dataset, batch_size=BATCH_SIZE, shuffle=True) # 2. 构建模型 model = CrossScaleModel( vocab_size=PROTEIN_VOCAB_SIZE, n_genes=NUM_GENES, embed_dim=EMBED_DIM, seq_len=NUM_PROTEIN_TOKENS, temperature=TEMPERATURE, ).to(DEVICE) optimizer = optim.Adam(model.parameters(), lr=LEARNING_RATE) # 3. 训练循环 model.train() for epoch in range(1, EPOCHS + 1): total_loss = 0.0 for batch in loader: seq, expr, pos = [x.to(DEVICE) for x in batch] loss = model.contrastive_loss(seq, expr, pos) optimizer.zero_grad() loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) optimizer.step() total_loss += loss.item() * seq.size(0) avg_loss = total_loss / len(dataset) if epoch % 5 == 0 or epoch == 1: print(f"Epoch {epoch:03d} | Contrastive Loss: {avg_loss:.4f}") torch.save(model.state_dict(), "checkpoints/cross_scale_model.pt") print("模型训练完成,已保存至 checkpoints/cross_scale_model.pt") if __name__ == "__main__": main()4.4 下游任务验证
训练完成后,我们需要验证学习到的表征是否具有生物学意义。一个常用的验证方式是:把细胞表达谱输入模型得到 cell embedding,然后训练一个简单的分类器,看能不能根据 embedding 区分细胞类型。
# file: evaluate.py import numpy as np import torch import torch.nn as nn from sklearn.linear_model import LogisticRegression from sklearn.model_selection import train_test_split from sklearn.metrics import accuracy_score, f1_score from config import * from models.alignment import CrossScaleModel def extract_cell_embeddings(model, loader): model.eval() embeddings, labels = [], [] with torch.no_grad(): for seq, expr, pos in loader: p = model.forward_protein(seq.to(DEVICE)) c = model.forward_cell(expr.to(DEVICE)) s = model.forward_spatial(pos.to(DEVICE)) # 三种embedding拼接,作为更丰富的细胞表示 emb = torch.cat([p, c, s], dim=-1).cpu().numpy() embeddings.append(emb) return np.concatenate(embeddings, axis=0) def main(): data = np.load("data/sim.npz", allow_pickle=True) cell_types = data["cell_types"] protein_seq = torch.LongTensor(data["protein_seq"]) cell_expr = torch.FloatTensor(data["cell_expr"]) spatial_pos = torch.FloatTensor(data["spatial_pos"]) model = CrossScaleModel( vocab_size=PROTEIN_VOCAB_SIZE, n_genes=NUM_GENES, embed_dim=EMBED_DIM, seq_len=NUM_PROTEIN_TOKENS, ).to(DEVICE) model.load_state_dict(torch.load("checkpoints/cross_scale_model.pt")) dataset = TensorDataset(protein_seq, cell_expr, spatial_pos) loader = DataLoader(dataset, batch_size=128, shuffle=False) embeddings = extract_cell_embeddings(model, loader) X_train, X_test, y_train, y_test = train_test_split( embeddings, cell_types, test_size=0.2, random_state=42 ) clf = LogisticRegression(max_iter=1000) clf.fit(X_train, y_train) y_pred = clf.predict(X_test) acc = accuracy_score(y_test, y_pred) f1 = f1_score(y_test, y_pred, average="macro") print(f"细胞类型分类准确率: {acc:.4f}") print(f"Macro F1: {f1:.4f}") if __name__ == "__main__": main()这里把三种模态的 embedding 拼接起来,相当于用多模态信息增强细胞表征,再用逻辑回归验证表征的可分性。真实项目中,通常还会用 UMAP 降维可视化,观察不同细胞类型是否在虚拟地图上自然聚类。
4.5 结果说明
在模拟数据上,由于我们人为构造了随机标签,逻辑回归准确率不会高到离谱,但训练流程可以完整跑通。这个原型的主要作用是验证框架正确性,而不是追求精度。
如果把它替换成真实数据,预期可以看到:
- 同一细胞类型在表征空间中聚成一簇。
- 蛋白相似性高的细胞,其细胞表征也更接近。
- 空间邻近的细胞在嵌入空间中也彼此靠近。
这就是“跨尺度虚拟地图”的基本形态。
5. 常见问题与排查思路
在复现和扩展这类模型时,容易遇到下面几类问题,我把排查思路整理成表:
| 问题现象 | 常见原因 | 解决思路 |
|---|---|---|
| 训练损失不下降 | 学习率过大或过小;数据未归一化 | 使用学习率预热;检查输入数据是否标准化 |
| 对比学习坍塌 | 正样本对质量差;负样本太少 | 增加负样本数量;增大 batch size;使用梯度裁剪 |
| 跨队列迁移效果差 | 批次效应严重 | 加入域对抗训练;增加域混合增强;使用 Harmonization 方法 |
| 显存溢出 | 蛋白序列过长或 batch size 过大 | 缩短序列;减小 batch;使用梯度累积 |
| 细胞类型分类准确率异常高 | 模拟数据或标签泄漏 | 检查数据划分;确保测试集独立 |
| 蛋白与细胞模态不对齐 | 蛋白编码器太弱 | 换用预训练蛋白模型,例如 ESM 的表示作为初始化 |
5.1 对比损失不下降
这是最常见的训练问题。先检查数据是否做了标准化。表达谱数据通常分布极不均匀,建议做 log1p 归一化或 z-score 归一化。其次检查温度参数是否设置合理,常见范围是 0.05 到 0.1。温度太低会导致 logits 过大,梯度容易爆炸;温度太高则无法有效区分正负样本。
5.2 模型坍塌
对比学习中最怕模型把所有样本映射到同一个点,也就是“表征坍塌”。一旦出现这种情况,损失会很低,但下游任务完全不可用。解决方案通常是:
- 增大 batch size,让负样本更丰富。
- 加入预测头(projection head)。
- 使用 SimCLR、MoCo 等成熟的对比学习框架。
- 定期用 UMAP 可视化表征分布。
5.3 跨队列泛化差
真实数据中,训练集和验证集来自同一个医院时表现很好,换一个医院就明显下降,这种情况大多是批次效应造成的。除了在训练阶段加入域对抗外,也可以在预处理阶段使用 Harmony、ComBat 等工具消除批次效应,再进行模型训练。
6. 最佳实践与工程建议
6.1 数据质量优先
基础模型非常依赖数据质量。在投入算力之前,先花时间检查数据分布、缺失率、标注一致性和批次效应。单细胞数据的质量控制需要关注线粒体基因比例、基因检出数、双细胞比例等指标。蛋白序列数据则要注意序列长度分布和物种来源。
一个常见误区是只关注数据“数量”而忽略“多样性”。如果预训练数据全部来自同一平台,模型就很难在不同平台间泛化。建议预训练阶段混合多个公开数据集,并记录每个样本的来源批次,方便后续做 domain adaptation。
6.2 特征标准化与数据增强
跨尺度数据在数值范围上差异极大。蛋白序列是离散 token,表达谱是连续数,空间坐标是物理位置。在进入模型前,需要对连续数据做标准化,对离散数据做 embedding。
数据增强方面,单细胞表达谱可以通过以下方式增加鲁棒性:
- 添加高斯噪声。
- 随机 mask 部分基因。
- 使用不同批次的数据混合(mixup)。
空间数据可以通过旋转、翻转等几何增强来提升空间编码器的泛化能力。
6.3 模型架构选择与预训练权重
如果计算资源有限,不建议从零训练蛋白编码器。目前已有大量开源蛋白语言模型可以复用,它们的 embedding 已经包含丰富的进化与结构信息。工程上可以采用“冻结预训练蛋白模型 + 训练轻量投影层”的方案,既降低显存占用,又加快训练速度。
细胞表达谱编码器也可以用基因模块或通路先验来初始化,例如基于 KEGG、GO 数据库构建基因分组,再通过注意力机制聚合通路信息。
6.4 评估体系构建
跨尺度模型的评估不能只看单个任务指标。建议构建多层评估体系:
- 模态内评估:蛋白功能预测、细胞类型分类、空间区域识别。
- 跨模态评估:蛋白-细胞检索、细胞-空间对齐。
- 跨队列评估:在外部独立队列上验证零样本或少样本迁移性能。
6.5 隐私、合规与可解释性
生物医学数据涉及患者隐私,使用前必须完成脱敏和合规审查。模型训练和部署应遵循最小权限原则,访问控制要严格。公开发布模型权重前,需要确认训练数据中不包含可识别患者身份的信息。
可解释性方面,建议使用注意力权重可视化、表征聚类、差异基因富集分析等手段,帮助研究者理解模型在关注哪些生物学特征,而不仅仅把模型当作黑盒。
7. 总结
本文从一个 Nature 研究主题出发,拆解了基础模型如何连接蛋白、细胞与肿瘤微环境,以及如何构建跨尺度、跨队列的虚拟地图。重点介绍了多模态编码器设计、对比学习对齐、跨队列泛化策略,并给出一个可运行的 PyTorch 原型,覆盖数据生成、模型构建、训练、评估全流程。
从工程角度看,真正落地这类模型,核心工作并不在模型本身,而在于数据治理、批次效应处理、评估体系设计和算力规划。对于刚接触这个方向的读者,可以先跑通本文原型,再用公开的单细胞数据集替换模拟数据,观察 UMAP 可视化结果,逐步理解跨尺度对齐的含义。后续可以继续学习蛋白语言模型、空间转录组建模、域适应算法等内容,不断扩展模型的生物学覆盖面。