A simple deep learning framework for single-cell RNA-seq embeddings.
Single-cell RNA sequencing generates expression profiles for thousands of genes across thousands of cells. To make sense of this data, we need to reduce it to a manageable number of dimensions while preserving the biological signal. PCA is the standard approach, but it's linear and may miss complex relationships.
We built scembed to explore whether autoencoders—simple neural networks that learn to compress and reconstruct data—can do better.
The framework has a few core pieces:
-
Encoders and Decoders: The building blocks. An encoder compresses gene expression down to a small number of dimensions; a decoder tries to reconstruct the original from that compressed form.
-
AutoEncoder (AE): Combines encoder and decoder. The bottleneck in the middle becomes our embedding.
-
Variational AutoEncoder (VAE): Like AE, but adds randomness to the bottleneck. This encourages smoother embeddings where similar cells land near each other.
-
Conditional VAE (CVAE): VAE that takes batch information as input, enabling batch effect correction for data integration.
-
Classifiers: MLPClassifier for direct classification, CellClassifier for using pretrained encoder + classification head.
-
Trainer: Handles the training loop with validation tracking and early stopping so we don't overfit.
Everything is modular. You can swap in different architectures or losses without rewriting the whole thing.
We used the PBMC 3K dataset—about 2,700 blood cells with known cell types. After preprocessing (normalization, log transform, selecting 2,000 highly variable genes), we trained both AE and VAE models and compared them to PCA.
| Experiment | Hidden Layers | Latent Dim | VAE Silhouette |
|---|---|---|---|
| Baseline | [512, 256] | 32 | 0.111 |
| Deeper | [512, 256, 128] | 32 | 0.207 |
| Deeper + wider | [512, 256, 128] | 64 | 0.221 |
The deeper network with 64 latent dimensions worked best.
| Method | Silhouette Score |
|---|---|
| PCA | 0.222 |
| AutoEncoder | 0.194 |
| VAE | 0.221 |
The VAE essentially matches PCA for cluster separation.
The VAE automatically discovered axes corresponding to major cell type programs:
| Dimension | Top Genes | Interpretation |
|---|---|---|
| 0 | CD74, HLA-DRA, HLA-DRB1 | Antigen presentation |
| 1 | S100A8, S100A9 | Monocyte markers |
| 2, 4, 6, 9 | NKG7, GZMB, GNLY | Cytotoxicity |
| 8 | LYZ, CST3 | Monocyte/macrophage signature |
Can VAE embeddings improve cell type classification over simpler approaches?
- Logistic Regression - Linear baseline (sklearn)
- Random Forest - Non-linear baseline (sklearn)
- MLP on Expression - Neural network on raw 1838 genes
- MLP on PCA - Neural network on 50 PCA components
- MLP on VAE - Neural network on 64 VAE latent dimensions
- CellClassifier - Frozen VAE encoder + classification head
| Method | Accuracy | F1 (macro) |
|---|---|---|
| Logistic Regression | 93.2% | 0.890 |
| Random Forest | 92.2% | 0.865 |
| MLP (Expression) | 94.7% | 0.948 |
| MLP (PCA) | 94.7% | 0.916 |
| MLP (VAE) | 94.1% | 0.917 |
| CellClassifier (frozen) | 93.0% | 0.914 |
VAE did NOT beat PCA for classification. This is an important negative result:
- For PBMC 3K, linear dimensionality reduction is sufficient
- VAE optimizes for reconstruction, not classification
- The dataset is small and cell types are well-separated
- Simpler methods win when the problem is simple
The errors are biologically sensible:
- CD4 T ↔ CD8 T confusion (both T cells, share markers)
- CD14+ ↔ FCGR3A+ Monocytes (both monocyte subtypes)
- NK → CD8 T (both cytotoxic, share genes like NKG7)
No B cell is ever classified as a Monocyte—the model learns real biology.
When combining data from different experiments/labs, technical "batch effects" can dominate:
Dataset A (Lab 1) ──┐
├──→ Combined ──→ Cells cluster by BATCH, not CELL TYPE
Dataset B (Lab 2) ──┘
The key insight: hide batch from the encoder, provide it to the decoder.
Encoder: genes → z (doesn't see batch)
Decoder: z + batch → reconstruction
This forces z to be batch-invariant—the encoder can't encode batch information because it doesn't see it.
We split PBMC 3K into two batches and added realistic effects to Batch B:
- Gene-specific scaling (0.5x to 2x per gene)
- Gene-specific shift (-0.2 to 0.3)
- Dropout (20% of values become zero)
- Gaussian noise (std=0.1)
| Method | Cell Type Silhouette (↑) | Batch Silhouette (→0) |
|---|---|---|
| PCA (no correction) | 0.075 | 0.067 |
| CVAE | 0.038 | 0.036 |
- CVAE reduces batch effect (0.067 → 0.036)
- But also reduces cell type separation (0.075 → 0.038)
- There's a trade-off between batch mixing and preserving biology
Our simple CVAE works but isn't perfect. State-of-the-art tools like scVI add:
- Negative binomial likelihood (for count data)
- Library size normalization
- More careful hyperparameter tuning
VAE/CVAE shine in scenarios our small, clean dataset doesn't represent:
| Scenario | PCA | VAE/CVAE |
|---|---|---|
| Small clean data | ✅ | ≈ |
| Large noisy data | ≈ | ✅ |
| Transfer learning | ❌ | ✅ |
| Batch integration | ❌ | ✅ |
| Few labels (semi-supervised) | ≈ | ✅ |
| Generate synthetic cells | ❌ | ✅ |
git clone https://github.com/narayanar/scembed.git
cd scembed
pip install -e .Requires Python 3.11+, PyTorch 2.0+, scanpy, and anndata.
from scembed.models import VAE
from scembed.trainer import Trainer
model = VAE(input_dim=2000, latent_dim=64, hidden_dims=[512, 256, 128])
trainer = Trainer(model, lr=1e-3, is_vae=True, kl_weight=0.001)
trainer.fit(train_loader, epochs=50)
# Get embeddings
embeddings = model(X)["mu"].detach().numpy()from scembed.models import CellClassifier
# Use pretrained VAE encoder
classifier = CellClassifier(
encoder=vae.encoder,
n_classes=8,
freeze_encoder=True
)from scembed.models import ConditionalVAE
cvae = ConditionalVAE(
input_dim=1838,
latent_dim=64,
n_batches=2,
batch_in_encoder=False # key for batch correction
)
# Forward pass includes batch info
out = cvae(expression, batch_onehot)
embeddings = out['mu'] # batch-correctedscembed/
├── scembed/
│ ├── encoders.py # MLPEncoder
│ ├── decoders.py # MLPDecoder
│ ├── losses.py # Reconstruction, KL divergence
│ ├── data.py # Dataset wrapper for AnnData
│ ├── trainer.py # Training loop
│ ├── metrics.py # Silhouette, ARI, NMI, accuracy, F1
│ └── models/
│ ├── ae.py # AutoEncoder
│ ├── vae.py # VAE
│ ├── cvae.py # Conditional VAE
│ ├── classifier.py # ClassificationHead, CellClassifier
│ └── mlp.py # MLPClassifier
├── tests/ # Unit tests (31 passing)
├── examples/
│ ├── pbmc_embedding.py
│ ├── pbmc_classification.py
│ ├── batch_simulation.py
│ └── batch_realistic.py
└── pyproject.toml
- Simple models work—no fancy architectures needed
- Architecture matters more than training time
- VAE matches PCA for clustering
- VAE doesn't beat PCA for well-separated cell types
- Errors are biologically sensible
- Simpler methods win on simple problems
- Hiding batch from encoder is key for CVAE
- There's a trade-off between batch mixing and cell type separation
- Real tools (scVI) add complexity for good reason
- Try negative binomial likelihood for raw counts
- Test on larger datasets where VAE may show advantage
- Compare with established tools (scVI, Harmony)
MIT