1. scGPT-spatial 是什么:复现前必须弄懂的底层设计
1.1 从单细胞大模型到空间版:它解决什么问题
scGPT 是 2024 年初发表在 Nature Methods 上的单细胞基础模型,本质上是一个以 Transformer 为核心的生成式架构。它把每个基因当作一个"token",把每个细胞当作一个"句子",通过自注意力机制学习基因之间的共现关系。这个设计思路跟 NLP 里的 BERT 几乎同构:预训练阶段用大规模 scRNA-seq 数据去学习基因表达的内在规律,下游任务阶段再通过微调适配到具体的生物学问题。
scGPT-spatial 就是在 scGPT 基础上针对空间转录组数据专门设计的一个变体分支。空间转录组和普通单细胞测序最大的区别在于:每个细胞或 spot 不仅有自己的转录组表达谱,还带有一组空间坐标。这组坐标引入了"邻近关系"——相邻位置上的细胞通常共享相似的微环境和功能状态。scGPT-spatial 的改进点在于,它把空间邻近关系融合进了注意力机制的编码过程,让模型在预测基因表达、识别组织结构的时候,既能用到表达谱信息,又能用到空间位置信息。
这个设计上的微妙之处,恰恰是复现时需要特别关注的。很多人在跑 scGPT-spatial 的时候把它当普通 scGPT 用,只是换了数据输入,结果空间模块根本没有被激活。后面我会专门讲到模型加载和输出这块的坑。
1.2 空间域识别任务到底是啥
scGPT-spatial 最常见的应用是空间域识别(spatial domain identification)。所谓空间域,本质上就是组织切片上一些功能相似、转录组特征相近的连续区域,比如大脑皮层里的不同层、肿瘤组织里的不同生态区。传统做法是用聚类算法(如 Leiden)先在表达谱上聚类,再叠加到空间坐标上看分布。但这种方法容易忽略空间连续性,导致同一片功能性区域被切得七零八落。
scGPT-spatial 的做法不一样。它利用空间注意力把相邻 spot 的表达信息相互增强,再做聚类或分类时,原本边界模糊的区域会因为空间先验而变得更加一致。用大白话说:相邻的点如果长得像,它们就更倾向被归到同一类。复现的时候你会发现,用 scGPT-spatial 得到的分区结果,在边界处的连续性和生物可解释性通常比纯表达聚类好不少。
1.3 复现前要问自己的三个问题
动手之前,先想清楚三个问题,能帮你省下大量无效时间:
第一,你要复现的是哪个版本?官方仓库里 scGPT 主线代码和空间分支代码混在一起。如果你只是 clone 下来跑main.py,默认走的是通用预训练路线,空间任务需要单独进入scgpt_spatial相关脚本或指定特定参数,上错车的结果就是跑了一堆 baseline 还摸不到空间模型的边。
第二,你的硬件条件是否够?空间转录组数据虽然比动辄百万细胞的 scRNA-seq 小一些,但 Transformer 的显存开销还是实打实的。我自己实测,单张 16G 显存(比如 v100 或者 4090)跑一个小鼠脑切片数据勉强可以,如果上更大数据集,要么梯度累积,要么切块,要么老老实实租卡。
第三,你的数据格式对不对?scGPT-spatial 官方示例数据用的是 h5ad 格式(AnnData),数据里必须同时包含表达矩阵和空间坐标。很多公开数据集给的是filtered_feature_bc_matrix目录或spaceranger_out,如果预处理不到位,模型读到的是残缺的坐标信息,训练出来的结果就是在瞎猜。
这三个问题想清楚之后,复现的过程就有了明确的目标:搭环境、整数据、跑微调、验结果,每一步你都知道自己在干什么,而不是跟着 README 无脑敲命令。
2. 环境准备:这是全流程里最不该省的步骤
2.1 依赖清单:为什么每个包都要锁版本
复现开源项目,最痛苦的不是代码逻辑,而是环境依赖。scGPT-spatial 的官方仓库依赖清单大致包括 Python 3.8 以上、PyTorch、PyTorch-Geometric(pyg)、scanpy、anndata、tqdm、transformers 等。这个组合里有两个天然的"版本雷区":
第一个雷区是 PyTorch 与 CUDA 的匹配。scGPT 核心训练代码基于 PyTorch,如果你的 CUDA 版本和 PyTorch 编译版本不一致,轻则警告,重则CUDA error: no kernel image is available on the device。这种情况下再好的模型也跑不动。
第二个雷区是 PyTorch-Geometric。pyg 对 torch 的版本极其敏感,torch-geometric的 2.x 版本必须与 torch 1.x/2.x 严格对应,一旦版本错位,import 阶段就会直接报错。很多新手在这里栽跟头,误以为是代码问题,折腾半天结果发现是 pyg 装错了。
另外还有一个小坑:官方仓库的 requirements 往往不会完全列清所有传递依赖。我的建议是不要贪快,老老实实按下面的流程建一个独立 conda 环境。
2.2 一步步建一个干净可复用的环境
我建议用 miniconda 管理环境,整个搭建过程大概 15 分钟。执行下面的命令前,先确认你的机器上已经安装好了合适的 NVIDIA 驱动和 CUDA 工具包。这里有个经验:查看nvidia-smi顶部的 CUDA Version,选 PyTorch 版本时,PyTorch 要求的 CUDA 版本必须小于等于驱动支持的 CUDA 版本。
conda create -n scgpt python=3.9 -y conda activate scgpt pip install torch==2.0.1 torchvision==0.15.2 torchaudio==2.0.1 --index-url https://download.pytorch.org/whl/cu118 pip install scanpy==1.9.3 anndata==0.8.0 pip install torch-geometric==2.3.1 pip install scikit-learn pandas tqdm pip install scgpt # 或者 clone 源码,推荐源码模式方便调试如果你走源码模式,克隆之后记得执行pip install -e .,这样后续修改源码里某些自定义函数能即时生效,不需要重复 install。
装完以后,建议先做一个空的导入测试:
import torch import torch_geometric import scanpy as sc import anndata import scgpt print("torch:", torch.__version__) print("pyg:", torch_geometric.__version__) print("scanpy:", sc.__version__) print("scgpt:", scgpt.__version__ if hasattr(scgpt, "__version__") else "ok")如果import torch_geometric这步就报错,八成是 pyg 和 torch 版本不匹配。此时不要硬调,直接去 pyg 官网查当前 torch 版本对应的 wheel 包,或者改用源码编译。
2.3 显存监控建议
环境搭好还只是第一步。训练 scGPT-spatial 时,显存占用会随 batch size 和数据规模剧烈波动。建议在正式开始训练前,用一个小数据子集先跑 1-2 个 step,通过nvidia-smi -l 1实时观察显存变化。如果接近显存上限,优先调低 batch size,其次考虑梯度累积,最后才考虑换小模型。这个顺序很关键:直接换小模型会改变隐层维度,导致预训练权重没法加载,反而更麻烦。
3. 数据管线:从 10x Visium 原始数据到 h5ad
3.1 数据集怎么选:第一次复现别选用太"野"的
scGPT-spatial 官方提供的示例数据包括小鼠大脑前部切片(10x Visium)和鸡胚心脏等公开空间转录组数据。第一次复现,我强烈建议就用官方示例数据,别一开始就上自己的数据。原因很简单:官方数据在通道兼容性上做过验证,跑通了再迁移到自己的数据,排查问题时能排除"数据格式不对"这个变量。
去 10x Genomics 官网或者 GEO 下载数据时,优先选filtered_feature_bc_matrix格式,这个目录下有barcodes.tsv.gz、features.tsv.gz和matrix.mtx.gz三个文件,是标准的稀疏矩阵格式。另外还需要spatial/tissue_positions_list.csv(对应每个 spot 的坐标)和scalefactors_json.json(用于后续可视化)。如果下载的是spaceranger完整输出,也能用,但需要额外的解析步骤。
3.2 预处理的关键操作:把坐标和表达对接上
拿到原始数据后,第一步是把表达矩阵读成 AnnData 对象。这里有一个容易踩的坑:scanpy 的read_10x_h5函数可以直接读 h5 文件,但如果你下载的是 mtx 目录,就得用read_mtx或手动拼接。建议用 scanpy 的底层函数统一处理:
import scanpy as sc import pandas as pd import numpy as np adata = sc.read_mtx("filtered_feature_bc_matrix/matrix.mtx.gz").T features = pd.read_csv("filtered_feature_bc_matrix/features.tsv.gz", sep="t", header=None) barcodes = pd.read_csv("filtered_feature_bc_matrix/barcodes.tsv.gz", sep="t", header=None) adata.var_names = features[1].values adata.obs_names = barcodes[0].values adata.var_names_make_unique()注意read_mtx默认行是基因、列是细胞,所以要做一次.T转置,否则后续维度全反了。
坐标信息的导入是空间数据的核心。把tissue_positions_list.csv读进来,按 barcode 对齐到adata.obs里:
positions = pd.read_csv("spatial/tissue_positions_list.csv", header=None) positions.columns = ["barcode", "in_tissue", "row", "col", "x", "y"] positions.set_index("barcode", inplace=True) adata.obs["x"] = positions.loc[adata.obs_names, "x"].values adata.obs["y"] = positions.loc[adata.obs_names, "y"].values如果某个 barcode 在坐标文件里找不到,先别急着删,看一下是不是数据类型不一致(比如str和int混了),大概率需要做一次astype(str)对齐。
接下来做基础的质量过滤。过滤标准没有绝对的对错,按我的习惯:
sc.pp.filter_genes(adata, min_counts=1) sc.pp.filter_cells(adata, min_counts=1)这一步的意思是去掉完全没有表达的基因和 spot,降低无效计算量。但千万不要做标准化的 scale 操作。scGPT 的处理逻辑是基于原始计数分箱(binning),它内部会自己处理归一化。如果你在前面提前做了对数化和标准化,会破坏模型的输入分布。
最后保存:
adata.write_h5ad("mouse_brain_raw.h5ad")4. 微调实操:让 scGPT-spatial 在自己的数据上跑起来
4.1 预训练权重和基因词典:从哪里来、怎么放
scGPT 官方在 Hugging Face 上发布了不同组织层次的预训练权重,包括scGPT_human、scGPT_mouse等。空间任务通常建议选用与数据物种一致的权重。下载后你会看到best_model.pt、vocab.json、args.json等文件。
vocab.json是基因词典,相当于把基因名映射成 token id。模型在训练时只认识 vocab 里存在的基因,不认识的基因会被忽略。如果你的数据集里基因名不是标准的 gene symbol(比如带着ENSEMBL前缀的),覆盖率会非常低,严重影响训练效果。所以数据预处理阶段最好把基因名统一成 UCSC gene symbol 格式。
预训练权重的加载方式有两种:一种是直接指定--model-file指向权重路径,另一种是在 Python 脚本里用torch.load手动加载。命令行方式更方便,但手动加载更灵活,利于调试。手动加载的骨架大致是这样:
model = scgpt.model.SpatialTransformer( ntoken=len(vocab), d_model=512, nhead=8, nlayers=12, d_hid=512, dropout=0.1, n_bins=51, pad_token=0, spatial_dim=2, ) state_dict = torch.load("best_model.pt", map_location="cpu") model.load_state_dict(state_dict, strict=False) model = model.to(device)注意strict=False很有讲究。预训练模型的权重结构里,大部分层可以直接匹配,但空间注意力引入的某些参数(比如坐标编码层)或者微调任务新增的分类头,在预训练权重里并不存在。用strict=False可以跳过这些缺失的层,避免加载直接报错。当然,这只是一种常见的微调实践,具体还要看你用的权重文件和模型结构的一致性。
4.2 主命令与参数解读:每个参数背后的逻辑
如果你习惯用官方仓库的训练入口,一个典型的命令长这样:
python main.py \ --data-path ./data/ \ --data-name mouse_brain_raw.h5ad \ --input-style raw \ --output ./output/ \ --model-file ./ckpt/scGPT_mouse/ \ --vocab-file ./ckpt/scGPT_mouse/vocab.json \ --n-bins 51 \ --batch-size 32 \ --epochs 30 \ --lr 1e-4 \ --optimizer adamw \ --grad-norm 1.0 \ --log-interval 100这里几个参数值得展开说:
--n-bins 51:模型把基因表达量离散化成 51 个分箱,这是 scGPT 延续 BERT 的做法——把连续的表达值变成 token 候选集。分箱数越大,模型对表达量的区分度越高,但训练难度也会增加。官方预训练模型基本都用 51,切换这个值意味着输出头的维度会变,预训练权重可能无法直接加载。
--lr 1e-4:微调阶段学习率普遍比预训练低,通常可以再试着调低到 5e-5。Transformer 对学习率很敏感,太大会导致 loss 冲高到 NaN,太小则会陷入漫长的收敛过程。建议用余弦退火调度器配合 warmup,前几个 epoch 让模型先稳定下来。我不会跟你说这是"最优解"——它只是一个经过验证非常稳的起点,你可以在这个基础上调到适合自己的数据集。
--batch-size 32:空间转录组数据一个 spot 就是一个样本,但 spot 之间本身有空间关联,抽样时如果完全随机,会破坏空间语义。更合理的做法是按区域采样,或者直接把整个切片作为一个图处理。官方实现里为了方便,通常还是用随机 batch,但你心里要清楚,batch 内的空间关系可能因为随机采样而部分丢失,这也是为什么空间任务上 batch size 太大反而不一定好。
4.3 训练监控:怎么判断模型在变好
训练开始后,不要只盯着 loss。空间域识别这类任务,loss 下降不代表分区结果变好。我的经验是至少额外盯三个指标:
一是基因表达重建的 accuracy。scGPT 自监督任务的核心是 mask 一些基因 token,让模型根据上下文预测这些基因的表达分箱。在验证集上,这个预测准确率是一个比 loss 更直观的质量信号。
二是空间域的连通性。训练过程中定期保存模型,然后跑到验证切片上做分区,肉眼观察分区结果的边界是否平滑、是否出现很多"孤岛"点。如果频繁出现零零碎碎的小区域,大概率模型的表达特征主导了分区,空间约束没有发挥出作用,可以考虑增大空间注意力层的权重。
三是显存和训练速度。如果每个 epoch 时间越来越长,可能是后期梯度累积 / 内存碎片等问题导致,及时停掉排查,别硬等。
5. 我踩过的五个坑,希望你一个都别踩
5.1 坑一:CUDA kernel 不匹配,模型搬不上 GPU
这个坑我复现时踩过。症状是:torch.cuda.is_available()返回 True,但模型.to(device)后一前向传播就报:
RuntimeError: CUDA error: no kernel image is available on the device排查链路是这样:先nvidia-smi看驱动支持的 CUDA 版本,然后python -c "import torch; print(torch.__version__)"看 PyTorch 编译用的 CUDA 版本。如果驱动版本过低而 PyTorch 编译版本过高(比如驱动只支持 CUDA 11.2 但 PyTorch 要求 CUDA 11.8),就会出这个错。解决办法很简单:重新安装适配当前驱动的 PyTorch,或者升级驱动。我当时是卸了 cu118 的 torch 换上 cu113 才跑通的。
5.2 坑二:pyg 版本错位,import 直接崩
torch_geometric和 PyTorch 版本错位的问题,最常见的报错是:
ModuleNotFoundError: No module named 'torch_geometric'不对,这个常见,但如果装的是不兼容版本,往往是ImportError或者 undefined symbol 类的底层错误。我的建议是,装 pyg 的时候别用pip install torch-geometric一把梭,而是先去 pyg 官网查与你 torch 版本对应的安装命令。pyg 为了兼容不同 torch 版本,用了一套特殊的编译分发机制,装错是常态。
还有一个小点:如果你装了torch-sparse或者torch-scatter,这几个扩展也跟 pyg 主版本强关联。一旦出现_sparse_cuda.so找不到的报错,基本就是这几个扩展和 pyg 版本不匹配导致的,干脆卸载重装。
5.3 坑三:基因词典覆盖率太低,模型等于白跑
这是个隐蔽的坑。我把一个鼠脑空间数据集的基因名拿来和vocab.json比对,发现匹配率只有 60% 出头。这时候模型其实只用了 60% 的基因信息在训练,剩下 40% 的基因被直接忽略。
为什么会这样?因为vocab.json里收录的是经过筛选的、在单细胞数据中高变的基因,而不是全基因组所有基因。你的数据如果不做高变基因筛选,大量低表达基因会稀释覆盖率。解决思路有两种:数据侧做sc.pp.highly_variable_genes筛选,只保留高变基因进模型;或者自定义一个覆盖更全的 vocab,但这样预训练权重就未必能对上。实操中我更倾向前者,毕竟微调的意义在于利用预训练知识。
5.4 坑四:坐标单位不一致,空间图建错
空间转录组的坐标文件里,有row/col(格点索引)和x/y(微米坐标)两套,除此之外还有px_col、px_row(像素坐标)这种别称。如果脚本里用的是像素坐标,但模型期望的是微米坐标,或者反过来,空间邻近关系的计算就会整体失真。典型症状是训练时 loss 能正常降,但可视化时空间域边界全部歪掉。
排查方式是把坐标画出来,散点图的形态应该和组织的实际形态一致。如果发现图像被拉伸、翻转或者明显错位,先检查坐标列选的是不是同一个系统。我在预处理脚本里加了一步:统一转成微米坐标,并做了空间单位归一化(除以切片的最大尺寸),这样既保留相对距离信息,又能让坐标数值在模型能消化的范围内。
5.5 坑五:OOM 之后没清缓存,显存被吃干
训练中途遇到 OOM,直接改 batch size 重启,结果发现显存并没有完全释放。偶尔还会出现CUDA out of memory但nvidia-smi显示显存占用一大片的情况。这是由于之前的失败进程可能还没完全退出,或者 PyTorch 的显存缓存机制导致显存碎片化。
我的做法是:遇到 OOM 先kill -9掉所有残留的 Python 训练进程,然后用nvidia-smi --gpu-reset复位 GPU(注意只有该 GPU 上没有其他任务时才能用)。训练代码里也建议在 dataloader 侧开启pin_memory=False,并调低num_workers,减少显存拷贝压力。
6. 结果怎么看:复现的终点是理解而不是跑通
6.1 定量指标:别被 ARI 忽悠
空间域识别任务最常用的定量指标是 ARI(Adjusted Rand Index)和 NMI(Normalized Mutual Information),它们衡量模型分出的空间域和手工注释/参考注释之间的吻合度。复现时看到官方结果里的 ARI 很高,但自己跑出来低不少,不要慌,先检查以下几个方面。
一是是否用了同一套参考注释。官方示例数据带了手工注释结果,比如小鼠大脑各层标注,这部分标注本身存在主观性。你在不同脚本或不同协议里拿到的注释版本可能不一致,直接对比 ARI 没有意义。
二是是否做了分层采样评估。空间域识别模型在训练时如果用了验证切片的数据做微调,验证 ARI 会虚高。更公平的做法是用一个切片做微调,在另一个切片上评估,或者至少做交叉验证。
三是 ARI 本身对离散区域数量敏感。区域数量越多,随机分区的 ARI 基线越低,模型效果好但 ARI 数值可能不如区域少的任务高。所以当你对比不同模型时,一定要确保分区数量一致,或者统一用同一份聚类数。
6.2 可视化:这才是空间模型证明自己的地方
定量的数值得看,但空间任务很大程度上是靠可视化说话的。我复现完之后,通常会画两张图放在一起对比:第一张是普通表达聚类(Leiden)叠加空间坐标,第二张是 scGPT-spatial 的分域结果叠加空间坐标。这个方法非常直观——你一眼就能看到,是不是传统聚类那边出现了很多小碎块,而空间模型那边区域更完整、边界更符合解剖学预期。
可选的可视化方式还有:
- 空间热图:按 spot 坐标画表达量均值热图,检查模型重建的表达分布是否保留了组织特异性。
- 结构熵分析:对每个 spot 计算熵值,熵越高说明那个位置越混乱。scGPT-spatial 通常会在组织边界处产生更高熵值,这是合理的,因为它识别到了不该强行合并的分界。
- 基因程序富集:把分出的每个空间域做差异表达分析,看富集到的基因程序是不是符合该区域的已知功能。
6.3 复现之后的扩展方向:这套流程还能怎么用
跑通一次 scGPT-spatial 之后,整个数据处理和训练流程基本就成了一个模板,后续可以做的扩展方向不少。
最直接的是换数据集跑。从 10x Visium 换到 Stereo-seq 或者其他空间平台,数据格式可能不同(像素级 spot、大尺寸矩阵),但核心的数据管线和微调流程是通用的。遇到高分辨率空间数据时,注意显存压力会显著增大,需要做区域切片(crop)处理,把大图切成若干个小 patch 分别输入模型,最后融合结果。
另一个方向是做批次整合。空间转录组实验里,多个切片之间存在批次效应,直接合并训练会让空间域识别收到批次信号干扰。用 scGPT-spatial 的思路扩展一下,可以在微调阶段加入批次 token 的嵌入,或者用对抗训练方式去掉切片来源信息。这个思路我在自己的数据上试过,效果比直接合并跑通用 embedding 好不少。
最后,如果你想深挖一下模型内部,可以尝试可视化注意力权重。空间注意力权重可以告诉我们哪些位置对预测某个 spot 的表达贡献最大。我第一次跑出这个图的时候,发现模型学到的"参考邻域"是各向异性的——它更依赖组织主要走向上的邻居信息,而不是简单的圆形邻域。这个观察对理解空间转录组的信号结构挺有价值。
复现整个 scGPT-spatial,说到底不是为了在 GitHub 上点个 star 或者截一张 ARI 表格发推。这个过程帮你理清了空间转录组数据到底应该怎么建模、预训练权重怎么复用、Transformer 在空间数据上的优势从哪里来。
最后分享一个小细节:我在跑完小鼠大脑之后,试着把同一个权重迁移到另一个物种(大鼠)的空间数据上,结果发现微调时如果保持较低学习率并冻结前几层,迁移效果居然也还不错。这说明 scGPT 在单细胞层面学到的基因关系是跨物种部分保守的,空间位置信号则更像是附加任务层的"微调专用"信息。这个理解可以指引你在自己的数据上做迁移实验时,优先尝试冻结主干、只调空间模块的配置,训练起来更省显存,收敛也更快。