从H5AD到空间感知scGPT:手把手复现与多任务训练实战
1. 空间转录组与scGPT技术背景
单细胞转录组测序技术近年来快速发展,为生物医学研究提供了前所未有的细胞分辨率。传统的单细胞RNA测序(scRNA-seq)虽然能够揭示细胞间的基因表达差异,但丢失了细胞在组织中的空间位置信息。而空间转录组技术(如10x Visium、MERFISH等)的出现,使得我们能够在保留空间位置信息的同时获取全转录组数据。
H5AD文件格式已成为单细胞和空间转录组数据的标准存储格式,它基于HDF5二进制格式,能够高效存储大规模的稀疏矩阵数据。一个典型的H5AD文件包含:
X矩阵:基因表达数据(细胞×基因)obs:细胞注释信息(如细胞类型、样本来源)var:基因注释信息(如基因名称、特征选择标记)obsm:细胞维度嵌入数据(如空间坐标、UMAP坐标)
scGPT是基于Transformer架构设计的单细胞数据分析模型,其核心优势在于:
- 基因词表映射:将基因名称转换为离散token,类似自然语言处理中的单词编码
- 多任务学习:同时处理基因表达预测(MLM)、细胞表达补全(MVC)和空间信息辅助补全(MVC Impute)
- 空间感知:通过坐标信息增强模型对组织结构的理解
2. 数据预处理实战
2.1 H5AD文件读取与基础处理
首先使用scanpy加载H5AD文件并进行基础质控:
import scanpy as sc # 读取H5AD文件 adata = sc.read_h5ad("spatial_data.h5ad") # 基础质控 print(f"原始数据维度: {adata.shape}") sc.pp.filter_cells(adata, min_genes=200) # 过滤低质量细胞 sc.pp.filter_genes(adata, min_cells=3) # 过滤低频基因 print(f"质控后维度: {adata.shape}")关键预处理步骤包括:
- 归一化处理:使用CPM(Counts Per Million)校正测序深度差异
- 对数变换:稳定方差,使数据更适合下游分析
- 高变基因选择:减少计算量,聚焦信息量丰富的基因
from scgpt.preprocess import Preprocessor preprocessor = Preprocessor( normalize_total=1e4, log1p=True, subset_hvg=2000 # 选择2000个高变基因 ) preprocessor(adata)2.2 空间坐标处理
空间转录组数据的核心价值在于其坐标信息,通常存储在obsm['spatial']中:
import pandas as pd # 检查空间坐标 if 'spatial' in adata.obsm: coords = pd.DataFrame(adata.obsm['spatial'], columns=['x', 'y'], index=adata.obs_names) print("空间坐标示例:") print(coords.head()) else: raise ValueError("未找到空间坐标信息!")对于坐标数据的特殊处理:
- 坐标归一化:不同样本间坐标尺度可能不同,需统一到相同范围
- 邻域构建:基于空间坐标计算细胞邻域关系,用于后续的KNN补全
# 坐标归一化 coords = (coords - coords.min()) / (coords.max() - coords.min()) adata.obsm['spatial_norm'] = coords.values3. 基因词表构建与数据转换
3.1 基因到token的映射
scGPT使用预定义的基因词表将基因名称映射为数字ID:
from scgpt.tokenizer import GeneVocab # 加载预训练词表 vocab = GeneVocab.from_file("vocab.json") # 基因名称统一为大写 adata.var_names = [gene.upper() for gene in adata.var_names] # 建立映射关系 adata.var['gene_id'] = [vocab[gene] for gene in adata.var_names] valid_genes = adata.var['gene_id'] >= 0 adata = adata[:, valid_genes] # 过滤不在词表中的基因3.2 构建训练样本
每个细胞需要转换为模型可接受的输入格式:
import torch import numpy as np def create_cell_example(adata, cell_idx): # 获取基因ID和表达值 gene_ids = adata.var['gene_id'].values expressions = adata.X[cell_idx].toarray().flatten() # 构建样本字典 example = { "genes": torch.tensor(gene_ids, dtype=torch.long), "expressions": torch.tensor(expressions, dtype=torch.float32), } # 添加空间坐标(如果存在) if 'spatial_norm' in adata.obsm: example["coordinates"] = torch.tensor( adata.obsm['spatial_norm'][cell_idx], dtype=torch.float32 ) return example4. 模型训练与多任务整合
4.1 scGPT模型架构
scGPT的核心是一个多层Transformer编码器:
from scgpt.model import TransformerModel model = TransformerModel( ntoken=len(vocab), # 词表大小 d_model=512, # 嵌入维度 nhead=8, # 注意力头数 d_hid=2048, # 前馈网络维度 nlayers=6, # Transformer层数 n_cls=1, # 分类头数量 vocab=vocab, # 基因词表 do_mvc=True, # 启用表达补全任务 do_mvc_impute=True # 启用空间辅助补全 ).to(device)4.2 多任务损失函数
scGPT同时优化三个任务的损失:
def compute_loss(outputs, inputs): # Masked Language Modeling损失 loss_mlm = F.mse_loss( outputs['mlm_output'], inputs['expr_values'] ) # Masked Value Completion损失 loss_mvc = F.mse_loss( outputs['mvc_output'], inputs['expr_values'] ) # 空间辅助补全损失 loss_mvci = masked_mse_loss( outputs['impute_pred'], inputs['expr_values'], ~inputs['padding_mask'] ) # 加权组合 total_loss = loss_mlm + 0.2 * loss_mvc + 0.1 * loss_mvci return total_loss4.3 训练循环实现
完整的训练流程包括数据加载、前向传播和参数更新:
from torch.utils.data import DataLoader from tqdm import tqdm # 构建DataLoader dataset = SpatialDataset(adata, vocab) dataloader = DataLoader(dataset, batch_size=32, shuffle=True) # 优化器设置 optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4) for epoch in range(30): model.train() total_loss = 0 for batch in tqdm(dataloader): # 数据转移到设备 genes = batch['genes'].to(device) exprs = batch['expressions'].to(device) coords = batch.get('coordinates', None) # 前向传播 outputs = model( src=genes, values=exprs, coordinates=coords ) # 计算损失 loss = compute_loss(outputs, batch) # 反向传播 loss.backward() optimizer.step() optimizer.zero_grad() total_loss += loss.item() print(f"Epoch {epoch+1}, Loss: {total_loss/len(dataloader):.4f}")5. 常见问题与解决方案
5.1 基因名称映射失败
问题现象:大量基因无法匹配到词表中的ID
可能原因:
- 基因命名方式不一致(ENSEMBL ID vs Symbol)
- 大小写不匹配
- 物种不匹配(人类vs小鼠)
解决方案:
# 使用mygene进行ID转换 import mygene mg = mygene.MyGeneInfo() query_result = mg.querymany( adata.var_names.tolist(), scopes='ensembl.gene', fields='symbol', species='human' ) # 构建映射字典 ensg_to_symbol = { item['query']: item.get('symbol', None) for item in query_result if not item.get('notfound') }5.2 训练损失不收敛
典型表现:
- Loss初始值很高(>100)
- 训练过程中波动大
- 后期下降缓慢
调试策略:
- 检查数据归一化:确保表达值经过log1p变换
print(f"表达值范围: {np.min(adata.X)} - {np.max(adata.X)}")- 调整学习率:尝试1e-5到1e-3之间的不同值
- 验证损失计算:单独检查每个任务的损失
print(f"MLM loss: {loss_mlm.item():.4f}") print(f"MVC loss: {loss_mvc.item():.4f}") print(f"Impute loss: {loss_mvci.item():.4f}")5.3 空间信息利用不足
优化方向:
- 增强坐标特征:将原始坐标转换为相对位置特征
def enhance_coordinates(coords): # 计算细胞间距离矩阵 dist = torch.cdist(coords, coords) # 添加局部密度特征 k = min(5, coords.shape[0]-1) knn_dist = torch.topk(dist, k=k+1, largest=False).values density = 1.0 / (knn_dist[:, 1:].mean(dim=1) + 1e-6) return torch.cat([coords, density.unsqueeze(1)], dim=1)- 调整损失权重:提高空间相关任务的损失系数
6. 进阶技巧与性能优化
6.1 多GPU训练加速
使用PyTorch的DistributedDataParallel实现多卡训练:
# 启动命令 torchrun --nproc_per_node=4 train.py对应的训练脚本修改:
# 初始化分布式环境 def setup_ddp(): dist.init_process_group("nccl") local_rank = int(os.environ["LOCAL_RANK"]) torch.cuda.set_device(local_rank) return torch.device(f"cuda:{local_rank}") # 模型包装 model = DDP(model, device_ids=[device.index])6.2 混合精度训练
通过自动混合精度(AMP)减少显存占用:
from torch.cuda.amp import GradScaler, autocast scaler = GradScaler() for batch in dataloader: optimizer.zero_grad() with autocast(): outputs = model(**batch) loss = compute_loss(outputs, batch) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()6.3 模型保存与加载
保存检查点时包含完整训练状态:
checkpoint = { 'model_state': model.state_dict(), 'optimizer_state': optimizer.state_dict(), 'epoch': epoch, 'loss': best_loss } torch.save(checkpoint, "model_checkpoint.pt")加载时恢复训练:
checkpoint = torch.load("model_checkpoint.pt") model.load_state_dict(checkpoint['model_state']) optimizer.load_state_dict(checkpoint['optimizer_state']) start_epoch = checkpoint['epoch']在实际项目中,空间感知的scGPT模型训练通常需要20-50个epoch才能达到理想效果。关键是要监控各个任务的损失变化,当MVC损失开始稳定时,可以适当降低学习率继续微调。
