Skip to content

CellMSA: Context Modeling for Single-Cell Representation Learning

Conference: NeurIPS2026
arXiv: 2609.38908
Code: https://github.com/PharMolix/CellMSA
Area: Computational Biology
Keywords: single-cell representation learning, multi-cell context, gene-pair representations, batch integration, cell state classification

TL;DR

CellMSA aligns cross-batch and related-type cells by gene identity, extracts gene-pair dependencies from low-dimensional context, and uses them to guide target-cell encoding, achieving strong results in label-informed integration, classification without test-label retrieval, and perturbation prediction combined with STATE-ST.

Background & Motivation

Single-cell RNA sequencing provides a high-dimensional expression vector for each cell, but a zero can indicate either genuine absence of expression or missing signal caused by limited sequencing depth or stochastic dropout. Foundation models such as scGPT and Geneformer mainly encode cells independently, inferring cell identity and gene relationships from one sparse observation; technical differences across donors, platforms, and batches can further obscure biological variation. CellPLM, STATE, and Stack introduce multiple cells, but compressing cells into vectors or gene modules can discard fine-grained gene relationships, while restricting context to same-type cells from the same batch often supplies primarily local denoising evidence.

The question is not simply how many neighbors to add, but which differences to compare and how those comparisons should influence target-cell modeling. Same-type cells across batches expose stable identity signals, whereas related but different cell types provide a background for distinguishing functional differences; the model must preserve both source distinctions and individual gene positions. Protein multiple sequence alignment (MSA) offers a modeling analogy: consistency and coordinated variation across related samples can be summarized into pairwise relationships and then used to process a target sample. Here, alignment follows gene identity rather than evolutionary sequence homology, so the borrowed element is a structured-context inductive bias, not an interpretation of expression associations as protein co-evolution or causal regulation.

Core idea: first extract state-dependent gene-pair relationships from gene-aligned multi-cell context, then inject those relationships into target-cell encoding as attention biases, allowing neighbors to change how the target is interpreted rather than merely averaging their cell vectors.

Method

Overall Architecture

The inputs are a target cell and its retrieved related cells; outputs include a 512-dimensional target-cell embedding and gene-pair representations available for analysis. The pipeline proceeds through Relation-Grouped Retrieval, CellMSA-Module, and GenePairformer: the low-dimensional module processes all context rows, while the higher-dimensional encoder processes only the target row. Three-Objective Pretraining supplies training supervision; extracting embeddings does not require reconstruction labels or training losses, but still requires context retrieval.

%%{init: {'flowchart': {'rankSpacing': 24, 'nodeSpacing': 28, 'padding': 6, 'wrappingWidth': 400}}}%%
flowchart TD
    A["Target cell and candidate pool"] --> B["Relation-Grouped Retrieval"]
    B --> C["CellMSA-Module"]
    C -->|Gene-pair representations| D["GenePairformer"]
    A -->|Target gene embeddings| D
    D --> E["Target-cell embedding"]
    D -.->|Training only: gene prediction| F["Three-Objective Pretraining"]
    E -.->|Training only: reconstruction and contrast| F
    E -->|Inference| G["Classification or STATE-ST"]

Expression is discretized into 11 categories, with zero expression separated and nonzero values binned by within-cell quantiles; the maximum input length is 2,048, and the gene vocabulary contains 61,982 entries. The target and its neighbors must share corresponding gene columns: independently sorting each cell by expression and treating matching column positions as matching genes would invalidate the alignment. Gene-identity and expression-value embeddings are added together, with neighbor rows also receiving an embedding of their relation to the target; a <cls> token aggregates the final cell representation. Gene-pair states are initialized from the identities of the two genes and subsequently refined by context and target-cell features, rather than supplied as a fixed external regulatory graph.

Key Designs

1. Relation-Grouped Retrieval: expose local repetition, cross-batch stability, and related-type differences

During pretraining, each target retrieves 16 same-batch same-type cells, 16 cross-batch same-type cells, and 8 related but different-type cells, giving 40 neighbors and 41 input rows including the target. The first group identifies locally shared expression, the second helps separate biological consistency from batch-specific features, and the third supplies contrast among nearby cell identities. Learnable relation embeddings explicitly distinguish these sources, so the model need not treat every neighbor as an equivalent positive example.

Related types are not selected arbitrarily: the authors aggregate expression across batches for each type, construct type-level mean vectors, and use cosine distance with Ward hierarchical clustering to form 28 clusters, retaining up to 50 most similar related types within the same cluster. The appendix reports mean within-cluster and between-cluster cosine similarities of 0.845 and 0.581, respectively, supporting expression-based similarity without establishing experimentally validated lineage relationships. This construction uses CELLxGENE type and batch metadata during pretraining, so the entire framework should not be characterized as fully label-free self-supervised learning.

Retrieval permissions change with downstream tasks: the main integration experiment uses existing type annotations, whereas classification validation and test retrieval use highly variable genes (HVGs), PCA, and KNN, without the type or state labels being predicted. Deployment therefore requires an explicit candidate pool, metadata policy, and retrieval rule; there is no universal set of “40 best neighbors” independent of task-specific data organization.

2. CellMSA-Module: compress cross-cell evidence into gene pairs without first discarding gene resolution

Gene embeddings for all rows are projected down to 128 dimensions before four layers alternating Outer-product Mean and Pair-weighted Averaging. For each gene pair, the former separately projects the features at the two gene columns, computes their outer product within each row, averages across cell rows, and writes the result back into the pair state. A group of genes showing consistent or coordinated changes across neighbors can thus accumulate cross-cell evidence in its pair representation; this is not a direct Pearson correlation of raw expression values.

\[ \mathbf{p}_{uv}^{l+1}=\mathbf{p}_{uv}^{l}+W_p^l\operatorname{Flatten}\left(\frac{1}{S}\sum_{s=1}^{S}\mathbf{a}_{su}^{l}\otimes\mathbf{b}_{sv}^{l}\right). \]

Here, \(S\) includes target and context rows, and the two projected vectors come from the low-dimensional features of genes \(u\) and \(v\) in the corresponding row. The outer product is flattened and projected back into an 8-dimensional gene-pair space, preserving pairwise structure while limiting the channel width of context computation.

Pair-weighted Averaging projects the updated pair states into head-specific gene-to-gene weights, applies softmax over the other gene position, and uses these weights to aggregate value features within each row. A gene-specific sigmoid gate in each row controls the update strength; cells share pair weights inferred from the context but retain their own expression features. Extracting relationships from row features and refining row features with those relationships therefore forms an iterative exchange rather than a single neighbor average.

The cost remains quadratic in gene count: Outer-product Mean costs approximately \(O(SG^2d_p^2)\), and Pair-weighted Averaging approximately \(O(G^2d_p+SG^2d')\). Dimensionality reduction narrows each pairwise operation but does not make the dense pair matrix linear in size; gene truncation is consequently both an engineering budget and a boundary on biological coverage.

3. GenePairformer: inject context relationships into target attention and refine them with target expression

After context extraction, pair states are projected into attention-head space, while the original target-row gene embeddings initialize six layers of target encoding. This division matters: context supplies a pairwise prior, but the final embedding still comes from the target cell rather than from a higher-dimensional Transformer applied equally to all 41 rows.

At each layer, target query–key similarities are added to the existing pair states, and softmax over the updated states determines value aggregation. The pair bias is therefore not a one-time addition: it absorbs target-sequence attention logits and progressively reflects the joint influence of current target features and cross-cell evidence.

\[ R_{uv}^{L+1,H}=R_{uv}^{L,H}+\frac{Q_u^{L,H}(K_v^{L,H})^\top}{\sqrt{d_H}}. \]

\(H\) denotes an attention head and \(d_H\) its dimension; the updated \(R\) determines gene-to-gene attention weights, and the output <cls> is projected into the cell embedding. Unlike a standard Transformer with a relationship graph attached only at the output, this encoder already constrains gene communication with contextual relationships inside every layer.

4. Three-Objective Pretraining: prevent neighbor copying while retaining identity and cell-specific variation

Masked gene expression modeling (MGM) predicts the original expression bins at masked positions, with the same gene positions masked simultaneously in the target and all context rows. Masking only the target would let the model read expression at the corresponding position in a neighbor; synchronized masking instead encourages reliance on other genes and their pairwise relationships, reducing that shortcut. Expression reconstruction recovers expression at randomly selected genes from the final cell embedding and gene identity, requiring the embedding to retain global transcriptomic information rather than only type identity.

Cell-level contrastive learning (CCE) uses a same-type cell as a positive and performs InfoNCE contrast against candidates from the same batch, using cosine similarity and a temperature of 0.05. Type metadata is weak supervision here as well; it helps establish biologically discriminative embeddings despite noise, but excessive emphasis can erase disease-state differences within a type. The appendix also indicates that excessive reconstruction emphasis can retain more input noise, motivating the use of both objectives rather than type contrast or expression reconstruction alone.

A Worked Example

For a proximal tubule (PT) target cell, pretraining retrieves 16 same-batch same-type neighbors, 16 same-type neighbors from other batches, and 8 related-type neighbors. Each gene occupies the same column across all 41 rows, allowing comparison of expression features stable across batches and relationships prominent in cellular states near the target. CellMSA-Module converts these comparisons into pair states, GenePairformer uses the states to interpret target expression, and <cls> provides the cell vector needed by a classifier.

For PT state classification at test time, neighbors instead come from HVG–PCA–KNN retrieval; using the true aPT, dPT, or dPT/DTL labels first would not constitute a valid prediction protocol. This example illustrates the mechanism rather than an additional experimental result; disease-associated pair activation is a state-association clue, not a demonstrated causal regulatory edge.

Loss & Training

The three objectives are combined as:

\[ \mathcal{L}=\mathcal{L}_{\mathrm{MGM}}+0.01\mathcal{L}_{\mathrm{Rec}}+10\mathcal{L}_{\mathrm{CCE}}. \]

MGM and reconstruction both use cross-entropy over expression bins; the MGM selection probability is 0.15, and the reconstruction gene-selection probability is 0.1. The model has 47.13M parameters and uses AdamW, 1,000 steps of linear warmup followed by a constant learning rate of \(10^{-5}\), and an effective batch size of 32. One epoch on four A800 GPUs takes approximately 20 days; pretraining draws on 2,090 datasets and 818 cell types, with downstream evaluation datasets excluded before training.

“109M” counts cell observations, including 65.6M primary observations, rather than 109M independent biological cells. The authors retain some non-primary observations from re-aggregation or atlas integration because a cell can have different contexts, but this does not eliminate the risk of overweighting repeated cells or atlas-specific structure.

Key Experimental Results

Main Results

The table retains results that address distinct questions; classification values are means and standard deviations over five random donor splits, whereas perturbation prediction uses one fixed split. Integration Total is computed as \(0.6S_{\mathrm{bio}}+0.4S_{\mathrm{batch}}\); all listed metrics are higher-is-better, but scores from different tasks are not directly comparable.

Dataset / Task Metric CellMSA Comparison Comparison Boundary
Tabula Sapiens integration Total 0.736 Stack 0.663 Both use label-informed type retrieval
Tabula Sapiens integration Bio / Batch 0.850 / 0.566 Stack 0.740 / 0.547 Does not imply superiority on every batch metric
Blood type classification Macro-F1 \(0.912\pm0.005\) CellPLM \(0.906\pm0.010\) Validation and test retrieval do not use true type labels
Kidney Atlas PT state classification Macro-F1 \(0.931\pm0.006\) STATE-SE \(0.912\pm0.009\) 3 states within one lineage
Replogle + STATE-ST Pearson \(\Delta\) 0.433 Expression + STATE-ST 0.398 Shared state-transition framework; one fixed split
Replogle + STATE-ST PRAUC / DE Overlap 0.334 / 0.215 STATE-SE + STATE-ST 0.287 / 0.180 Recovery of differentially expressed genes

In the main integration benchmark, scGPT, Geneformer, scVI, and other baselines do not use type labels for representation extraction, so the full ranking is not a label-free comparison with equal information access. The appendix separately reports label-free retrieval on Bladder only: CellMSA achieves Batch / Bio / Total of 0.509 / 0.807 / 0.688, versus 0.451 / 0.699 / 0.600 for Stack. This supports an advantage not entirely attributable to type labels, but does not replace label-free replication across all tissues or isolate an individual module's contribution.

Classification uses donor-ID splits in a 7:1:2 ratio and a three-layer MLP trained on embeddings from frozen foundation models; Blood includes only types accounting for at least 1% of cells. PT labels aPT, dPT, and dPT/DTL denote normal-like, injured, and severely injured states with transcriptional identity drift; CellMSA achieves Accuracy of \(0.962\pm0.002\) and \(0.958\pm0.004\) on the two classification tasks, respectively. These results do not establish performance on arbitrary rare cell types or clinical disease diagnosis.

The perturbation dataset contains four cell lines and 100 shared perturbations with relatively many cells, totaling approximately 132k cells with controls; 45% of HepG2 perturbations are held out for testing and 5% for validation, with the remainder and other cell lines used for training. Context is restricted to the same cell line, and cross-type means different perturbation-defined states rather than ontology-level cell types; post-perturbation expression from validation and test sets cannot serve as context. Pearson \(\Delta\) is the gene-wise Pearson correlation between predicted and real “perturbation mean minus control mean” vectors; 0.433 represents approximately an 8.8% relative gain over 0.398. CellMSA + STATE-ST achieves Spearman-FC of 0.431 versus 0.404 for Expression + STATE-ST, measuring fold-change rank agreement on truly significant differentially expressed genes.

Ablation Study

Config PT Accuracy PT Macro-F1 Perturbation Pearson \(\Delta\) Perturbation Spearman-FC
Remove CellMSA-Module and pair representations 0.854 0.811 0.355 0.366
Retain architecture without context input 0.878 0.817 0.389 0.367
Same-batch context only 0.929 0.893 0.400 0.397
Full model 0.958 0.931 0.433 0.431

Relative to same-batch neighbors only, the full model gains 0.038 in PT Macro-F1 and 0.033 in perturbation Pearson \(\Delta\), supporting the additional value of cross-batch and related-type context. Removing the entire module and pair states reduces PT Macro-F1 by 0.120, but changes several factors simultaneously, so the entire difference cannot be attributed to one operator. Context-size gains saturate beyond 40 neighbors; exact values for individual curve points are not provided in the main text and are therefore not reconstructed into a table.

Key Findings

  • Gene-pair representations supply state-associated interpretation clues: in the PT disease analysis, heads 6 and 7 align more with disease markers and heads 4 and 5 with homeostatic markers, but head indices are not fixed causal-function labels.
  • In the external STRING comparison, 26 of KDR's top 50 partners have functional-association support, or 52%; without a random-background significance test, this is not significant enrichment or a measure of GRN recovery accuracy.
  • After randomly removing 50% of nonzero genes in Blood cells, original and corrupted embeddings retain a mean Pearson correlation of 0.972 and a median of 0.984; this measures embedding stability, not classification accuracy on corrupted inputs.

Highlights & Insights

  • Context estimates how genes jointly participate in a state instead of serving as a mean of final vectors. Multi-cell information therefore operates inside target-cell attention.
  • Synchronized masking at matching positions is an important companion to the architecture: withholding the predicted position from neighbors more strongly encourages learning cross-gene dependencies.
  • Low-dimensional multi-row modeling followed by higher-dimensional target-only modeling may transfer to other omics data, but requires feature-identity alignment rather than treating arbitrary neighbor collections as valid MSAs.

Limitations & Future Work

  • Pair activations are statistical dependencies and may reflect shared upstream factors, cell composition, or technical bias; even a directed matrix does not establish causal regulatory direction.
  • Information access differs in the main integration benchmark, and complete label-free evidence covers only Bladder; evaluation should extend to more tissues and candidate pools with unreliable annotations.
  • Expression binning, the 2,048-token limit, and dense quadratic pair computation constrain expression precision and gene coverage; low-rank or sparse alternatives require validation of their state-discrimination capacity.
  • Non-primary corpus observations may be repeated, and perturbation prediction uses only one split; current results do not establish stability across random seeds, species, or clinical domains.
  • Independent perturbation experiments could test the functional meaning of pair relationships; current outputs are better suited to generating research hypotheses than to direct clinical decisions.
  • vs scGPT / Geneformer: independent-cell encoders directly extract representations, whereas CellMSA additionally retrieves and uses context relationships; the gains bring candidate-pool and computation requirements.
  • vs CellPLM / Stack: all go beyond single-cell inputs, but CellMSA maintains gene-column alignment and explicitly extracts gene pairs; its Blood Macro-F1 advantage over CellPLM is small, so comparisons only to weaker baselines would exaggerate the general gain.
  • vs AlphaFold-style MSA: the transferable element is forming pair representations from consistency and variation across samples, not the biological meaning of residue-contact prediction.
  • vs STATE: perturbation experiments use STATE-ST as the state-transition model, showing that enhanced embeddings help that framework rather than that CellMSA alone replaces the complete perturbation-prediction system.

Rating

  • Novelty: 4/5. Effectively connects gene-pair structure, source-aware context, and target-cell encoding, while borrowing core operators from established architectures.
  • Experimental Thoroughness: 4/5. Covers integration, classification, perturbation, and interpretation, with unequal information access and a single perturbation split limiting comparison strength.
  • Writing Quality: 4/5. The main pipeline and appendix settings are relatively clear, but association-based interpretation must remain distinct from causal mechanisms.
  • Value: 4/5. Offers a reusable route to context-enhanced single-cell representations, with deployment requiring retrieval-quality assessment and independent biological validation.