CellPLM#
CellPLM 是首个编码细胞间关系的单细胞预训练语言模型,在多种下游任务中持续优于现有的预训练和非预训练模型,推理速度比现有预训练模型快 100 倍。
您可以使用 omicverse.llm.SCLLMManager(model_type="cellplm") 直接调用此模型。
引用:Wen, H., Tang, W., Dai, X., Ding, J., Jin, W., Xie, Y., & Tang, J. (2023). CellPLM: Pre-training of cell language model beyond single cells. BioRxiv, 2023-10.
import scanpy as sc
import omicverse as ov
ov.plot_set(font_path='Arial')
# Enable auto-reload for development
%load_ext autoreload
%autoreload 2
🔬 Starting plot initialization...
Downloading Arial font from GitHub...
Arial font downloaded successfully to: /tmp/omicverse_arial.ttf
Registered as: Arial
🧬 Detecting CUDA devices…
✅ [GPU 0] NVIDIA H100 80GB HBM3
• Total memory: 79.1 GB
• Compute capability: 9.0
____ _ _ __
/ __ \____ ___ (_)___| | / /__ _____________
/ / / / __ `__ \/ / ___/ | / / _ \/ ___/ ___/ _ \
/ /_/ / / / / / / / /__ | |/ / __/ / (__ ) __/
\____/_/ /_/ /_/_/\___/ |___/\___/_/ /____/\___/
🔖 Version: 1.7.6rc1 📚 Tutorials: https://omicverse.readthedocs.io/
✅ plot_set complete.
加载示例数据集#
本教程使用 NeurIPS 2021 单细胞竞赛数据集中的三个批次,这是批次整合和细胞类型注释的绝佳测试案例。
adata1=ov.read('data/neurips2021_s1d3.h5ad')
adata1.obs['batch']='s1d3'
adata2=ov.read('data/neurips2021_s2d1.h5ad')
adata2.obs['batch']='s2d1'
adata3=ov.read('data/neurips2021_s3d7.h5ad')
adata3.obs['batch']='s3d7'
adata=sc.concat([adata1,adata2,adata3],merge='same')
adata
AnnData object with n_obs × n_vars = 27423 × 13953
obs: 'GEX_n_genes_by_counts', 'GEX_pct_counts_mt', 'GEX_size_factors', 'GEX_phase', 'ADT_n_antibodies_by_counts', 'ADT_total_counts', 'ADT_iso_count', 'cell_type', 'batch', 'ADT_pseudotime_order', 'GEX_pseudotime_order', 'Samplename', 'Site', 'DonorNumber', 'Modality', 'VendorLot', 'DonorID', 'DonorAge', 'DonorBMI', 'DonorBloodType', 'DonorRace', 'Ethnicity', 'DonorGender', 'QCMeds', 'DonorSmoker', 'is_train'
var: 'feature_types', 'gene_id'
obsm: 'ADT_X_pca', 'ADT_X_umap', 'ADT_isotype_controls', 'GEX_X_pca', 'GEX_X_umap'
layers: 'counts'
adata=ov.pp.preprocess(adata,mode='shiftlog|pearson',
n_HVGs=3000,batch_key=None,target_sum=1e4)
adata
adata.raw = adata
adata = adata[:, adata.var.highly_variable_features]
adata
View of AnnData object with n_obs × n_vars = 27423 × 3000
obs: 'GEX_n_genes_by_counts', 'GEX_pct_counts_mt', 'GEX_size_factors', 'GEX_phase', 'ADT_n_antibodies_by_counts', 'ADT_total_counts', 'ADT_iso_count', 'cell_type', 'batch', 'ADT_pseudotime_order', 'GEX_pseudotime_order', 'Samplename', 'Site', 'DonorNumber', 'Modality', 'VendorLot', 'DonorID', 'DonorAge', 'DonorBMI', 'DonorBloodType', 'DonorRace', 'Ethnicity', 'DonorGender', 'QCMeds', 'DonorSmoker', 'is_train'
var: 'feature_types', 'gene_id', 'n_cells', 'percent_cells', 'robust', 'means', 'variances', 'residual_variances', 'highly_variable_rank', 'highly_variable_features'
uns: 'log1p', 'hvg', 'status', 'status_args', 'REFERENCE_MANU'
obsm: 'ADT_X_pca', 'ADT_X_umap', 'ADT_isotype_controls', 'GEX_X_pca', 'GEX_X_umap'
layers: 'counts'
初始化 CellPLM 模型#
CellPLM 需要包含 Transformer 权重和标记化词汇表的预训练模型检查点。该模型支持多种流水线任务,包括嵌入提取、细胞类型注释和数据填补。
从以下地址下载 CellPLM 模型:https://www.dropbox.com/scl/fo/i5rmxgtqzg7iykt2e9uqm/h/ckpt?dl=0&subfolder_nav_tracking=1
manager = ov.llm.SCLLMManager(
model_type="cellplm",
model_path="llm_model/models/cellplm",
pretrain_version="20231027_85M"
)
📥 Loading CellPLM model from llm_model/models/cellplm...
[Loaded] CellPLM model loaded from llm_model/models/cellplm
- Loaded pipelines: annotation, embedding, imputation
- Pretrain version: 20231027_85M
零样本嵌入生成#
零样本嵌入利用 CellPLM 的预训练知识,无需任何数据集特定训练即可生成有意义的细胞表征。
512 维嵌入为每个细胞的转录状态提供了压缩但信息丰富的表示。
embeddings = manager.get_embeddings(adata)
[🔬Cells] Data Summary:
Cells: 27,423
Genes: 3,000
Batches: 3
s3d7: 11,230 cells
s2d1: 10,258 cells
s1d3: 5,935 cells
[Embedding] Starting get_embeddings...
cells: 27,423
genes: 3,000
[Embedding] Extracting embeddings for 27,423 cells...
Automatically converting gene symbols to ensembl ids...
After filtering, 2339 genes remain.
[✅Complete] get_embeddings completed successfully!
[✅Complete] Results summary:
embedding_shape: (27423, 512)
embedding_dim: 512
embeddings.shape
(27423, 512)
print(f"embedding: {embeddings.shape}")
adata.obsm['X_cellplm'] = embeddings
sc.pp.neighbors(adata, use_rep='X_cellplm')
sc.tl.umap(adata)
ov.pl.embedding(
adata,
basis='X_umap',
color=['batch', 'cell_type']
)
embedding: (27423, 512)
computing neighbors
finished: added to `.uns['neighbors']`
`.obsp['distances']`, distances for each pair of neighbors
`.obsp['connectivities']`, weighted adjacency matrix (0:00:22)
computing UMAP
finished: added
'X_umap', UMAP coordinates (adata.obsm)
'umap', UMAP parameters (adata.uns) (0:00:17)
微调以提升性能#
微调将 CellPLM 的预训练权重适配到数据集的具体特征,显著提升下游任务的性能。这种监督学习方法利用参考批次(s1d3)的高质量细胞类型注释来:
reference_adata=adata[adata.obs['batch']=='s1d3']
reference_adata.obs['celltype']=reference_adata.obs['cell_type'].copy()
fine_tune_results = manager.model.fine_tune(
train_adata=reference_adata,
epochs=1500, #
batch_size=32, #
lr=1e-4, #
)
🚀 Starting CellPLM fine-tuning for annotation task...
Training parameters: epochs=500, batch_size=32, lr=0.0001
📊 Preparing cell type mapping...
Found 30 cell types: ['B1 B IGKC+', 'B1 B IGKC-', 'CD4+ T activated', 'CD4+ T activated integrinB7+', 'CD4+ T naive', 'CD8+ T CD49f+', 'CD8+ T CD57+ CD45RA+', 'CD8+ T CD69+ CD45RA+', 'CD8+ T CD69+ CD45RO+', 'CD8+ T TIGIT+ CD45RO+', 'CD8+ T naive', 'CD14+ Mono', 'CD16+ Mono', 'Erythroblast', 'G/M prog', 'HSC', 'Lymph prog', 'MAIT', 'MK/E prog', 'NK', 'Naive CD20+ B IGKC+', 'Naive CD20+ B IGKC-', 'Normoblast', 'Plasma cell IGKC+', 'Proerythroblast', 'Reticulocyte', 'T reg', 'Transitional B', 'cDC2', 'pDC']
🔄 Preparing training and validation data...
Split data: 4748 train, 1187 validation
🏋️ Starting training with CellPLM pipeline...
📈 Training for 500 epochs with real-time metrics...
After filtering, 2339 genes remain.
📊 Final Training Results (Epoch 499):
🎯 Train ACC: 0.7077
✅ Valid ACC: 0.6934
📈 Train F1: 0.7066
📈 Valid F1: 0.6846
✓ CellPLM annotation fine-tuning completed successfully!
使用微调模型进行批次整合#
微调完成后,我们执行批次整合以消除技术变异,同时保留生物学差异。这一关键步骤确保来自不同批次的细胞能够被正确比较和共同分析。
zero_shot_results = manager.model.integrate(
adata,
batch_key="batch",
correction_method="mnn",
)
adata.obsm['X_cellplm_fine'] = zero_shot_results['embeddings']
🔗 Performing batch integration for 27423 cells...
🧬 Extracting embeddings for 27423 cells...
Automatically converting gene symbols to ensembl ids...
After filtering, 2339 genes remain.
Applying MNN correction...
MNN correction applied to 3 batches
sc.pp.neighbors(adata, use_rep='X_cellplm_fine')
sc.tl.umap(adata)
ov.pl.embedding(
adata,
basis='X_umap',
color=['batch', 'cell_type']
)
使用微调模型进行细胞类型注释#
经过微调的 CellPLM 模型现在可以预测数据集中所有细胞的类型,包括那些来自训练中未使用的批次的细胞。这展示了模型将学习到的模式泛化到新数据的能力,同时充分利用了微调所获得的改进判别能力。
prediction_results = manager.model.predict_celltypes(
adata,
)
adata.obs['predicted_celltype'] = prediction_results['predicted_celltypes']
adata.obs['predicted_celltype_id'] = prediction_results['predictions']
🔮 Predicting cell types for 27423 cells...
Automatically converting gene symbols to ensembl ids...
After filtering, 2339 genes remain.