Skip to content

Latest commit

 

History

1 Commit

Folders and files

NameName
Last commit message
Last commit date
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 

Repository files navigation

scRNA-seq Neural Dynamics

Reconstructing cell differentiation as a continuous-time stochastic process, learned entirely from single-cell RNA-seq snapshots — no individual cell is ever tracked over time, only population snapshots at different days, so the pipeline infers the most likely "flow" between them and then learns a generative model of that flow.

Built on the Schiebinger et al. 2019 iPSC reprogramming dataset (236,285 cells, 39 real timepoints spanning day 0–18), using diffusion maps, entropic optimal transport, and a neural stochastic differential equation (SDE) trained via the adjoint sensitivity method.


Table of contents


What this project actually does

Imagine a photo album of thousands of cells, snapped at random moments as they reprogram from one cell type into another. The problem: no single cell is ever tracked twice. Every snapshot is a different, unlabeled group of cells — there's no way to know which cell in day 3's photo grew out of which cell in day 2's photo.

This project answers two questions using only those disconnected snapshots:

  1. What is the most likely "path" a typical cell takes as it reprograms?
  2. Can a model simulate new, realistic cell trajectories along that path — including the natural randomness real biology has?

It does this in four stages:

Stage What it does Analogy
Diffusion map Compresses ~2,000 noisy gene-expression dimensions into a smooth, low-dimensional map Flattening a crumpled ball of paper into a clean 2D map
Optimal transport (OT) Infers the most likely correspondence between two timepoints' cell populations, without tracking individuals Inferring traffic flow from two aerial photos of a city an hour apart
Neural SDE Learns the actual rules of motion — drift (direction) and diffusion (randomness) — from the OT-inferred flow Fitting a weather model that can now generate new, realistic forecasts
Simulation Uses the trained SDE to generate brand-new, realistic cell trajectories Hitting "play" on a differentiation movie that was never actually filmed

The math, stage by stage

This isn't a black-box pipeline — every stage is real, inspectable mathematics:

  1. Graph Laplacian / diffusion maps — build a k-nearest-neighbor graph over cells, compute the normalized graph Laplacian, and use its eigenvectors as the low-dimensional embedding. Diffusion distance in this space approximates geodesic distance on the underlying cell-state manifold, justified via the connection between the graph Laplacian and the continuous Laplace-Beltrami operator (heat kernel: e^{-tλᵢ}).
  2. Entropic optimal transport — for each pair of consecutive real timepoints, solve a regularized Kantorovich problem (Sinkhorn algorithm) to get the most probable coupling between the two cell populations, without ever assuming individual-cell correspondence.
  3. Neural SDE — dX_t = f_θ(X_t, t) dt + g_θ(X_t, t) dW_t, where f_θ (drift) and g_θ (diffusion) are small MLPs, trained end-to-end via torchsde's adjoint sensitivity method so simulated trajectories match the OT-derived coupling across all 38 consecutive day-pairs jointly.
  4. Validation — trustworthiness (embedding quality), coupling entropy (OT confidence), and simulated-vs-real distance (does the trained SDE actually land where real cells went) are all computed and plotted per stage.

At approximate/big-data scale, the diffusion map step uses a Nyström-extended eigensolve (exact eigendecomposition on a landmark subset, extended to all cells) instead of a full O(n³) eigendecomposition, and OT uses moscot's low-rank, GPU-accelerated Sinkhorn instead of a dense O(n²) cost matrix.


Architecture

A layered pipeline architecture with the Strategy pattern for the two swappable science stages:

raw scRNA-seq data
        │
        ▼
preprocess & normalize        (data/preprocessing.py)
        │
        ▼
diffusion map embedding       (embedding/diffusion_map.py)     [EmbeddingStrategy]
        │
        ▼
optimal transport pseudotime  (transport/optimal_transport.py) [TransportStrategy]
        │
        ▼
train neural SDE               (dynamics/neural_sde.py, trainer.py)
        │
        ▼
simulated cell trajectories    (visualization/plots.py, app/streamlit_app.py)
  • embedding/base.py — EmbeddingStrategy interface. Current implementation: DiffusionMapEmbedding (exact or Nyström-approximated).
  • transport/base.py — TransportStrategy interface. Current implementation: OptimalTransportCoupling, backend-switchable between POT (exact, small-data) and moscot (GPU, low-rank, big-data).
  • Config-driven, dependency-injected — no hyperparameter is hardcoded inside any stage. Every stage receives a typed config dataclass at construction time, all loaded from one YAML file (utils/config.py).
  • The notebooks/ and app/ layers only ever call into pipeline/ and visualization/ — they never touch science-layer internals directly, so those internals stay independently swappable and testable.

Project structure

scrna-neural-dynamics/
├── README.md
├── requirements.txt
├── configs/
│   └── default.yaml                       # every hyperparameter, one file
├── data/
│   ├── raw/datasets/                      # downloaded .h5ad files go here
│   └── processed/                         # pipeline outputs (models, embeddings)
├── app/
│   └── streamlit_app.py                   # interactive dashboard
├── src/scrna_dynamics/
│   ├── data/
│   │   ├── loader.py                      # backed mode, sparse-aware ingestion
│   │   └── preprocessing.py               # normalize, HVG selection, PCA
│   ├── embedding/
│   │   ├── base.py                        # EmbeddingStrategy interface
│   │   └── diffusion_map.py               # exact + Nystrom-approximated
│   ├── transport/
│   │   ├── base.py                        # TransportStrategy interface
│   │   └── optimal_transport.py           # POT / moscot backends
│   ├── dynamics/
│   │   ├── neural_sde.py                  # drift + diffusion MLPs
│   │   └── trainer.py                     # adjoint-method training loop
│   ├── pipeline/
│   │   └── pipeline.py                    # orchestrates every stage via config
│   ├── evaluation/
│   │   └── metrics.py                     # trustworthiness, coupling entropy, SDE validation
│   ├── visualization/
│   │   └── plots.py                       # embeddings, coupling heatmaps, drift fields, animations
│   └── utils/
│       └── config.py                      # typed config dataclasses + YAML loader
├── notebooks/
│   └── 00_all_visualizations.ipynb        # every visualization, one notebook
├── scripts/
│   └── run_pipeline.py                    # CLI entry point
└── tests/                                  # unit tests per module

Setup

Python version

Python 3.10 or 3.11. Not 3.12+ (JAX/moscot dependencies lag on 3.12 support). Not below 3.10 (recent scanpy drops older Python).

Environment

python3.11 -m venv .venv
source .venv/bin/activate        # Windows: .venv\Scripts\activate
pip install -r requirements.txt

If you're behind a network that blocks Anaconda's default channels (this affects users in some regions due to sanctions on repo.anaconda.com), skip conda entirely and use the venv route above — it relies on PyPI and package-specific hosts instead.

GPU support (optional, needed for moscot/big-data mode)

requirements.txt installs CPU-only wheels for JAX and PyTorch by default. For GPU:

pip install --upgrade "jax[cuda12]"
pip install torch --index-url https://download.pytorch.org/whl/cu121

(match the CUDA version to your system — check pytorch.org for the current command)


Getting the data

This project uses the Schiebinger et al. 2019 iPSC reprogramming time-course. It is not a scanpy built-in — download it manually:

mkdir -p data/raw/datasets
curl -L -o data/raw/datasets/reprogramming_schiebinger.h5ad \
  "https://ndownloader.figshare.com/files/28618734"

This is a ~1.4 GB file (236,285 cells × 19,089 genes). A smaller serum-only subset (165,892 cells) is available at https://ndownloader.figshare.com/files/35858033 if you want a faster first pass.

Verify the download:

ls -lh data/raw/datasets/reprogramming_schiebinger.h5ad

A file a few KB in size means the download failed/was interrupted — delete and retry.

The dataset's real timepoint column is .obs['day'] — this is the time_key used throughout the pipeline.


Configuration

Every hyperparameter lives in configs/default.yaml:

dataset_name: schiebinger19
random_seed: 42
backed_mode: true              # memory-map instead of loading fully into RAM
use_sparse: true
max_total_cells: 10000         # stratified subsample cap across all timepoints
max_cells_per_timepoint: 3000  # per-segment cap for OT/SDE stages

preprocessing:
  n_top_genes: 2000
  n_pca_components: 50

diffusion_map:
  n_neighbors: 15
  n_components: 10
  diffusion_time: 1.0
  approximate: true            # Nystrom-approximated eigensolve
  nystrom_sample_size: 5000

optimal_transport:
  epsilon: 0.05                # entropic regularization strength
  backend: moscot               # 'pot' (exact, small-data) or 'moscot' (GPU, big-data)
  rank: 100
  device: gpu

neural_sde:
  hidden_dim: 64
  n_layers: 3
  learning_rate: 0.001
  n_epochs: 200
  solver: euler
  device: gpu

Rerun any experiment with different settings by editing this file — no code changes required.


Running the pipeline

python scripts/run_pipeline.py --config configs/default.yaml --time-key day

Outputs are written to data/processed/:

  • sde_model.pt — trained neural SDE weights
  • embedding.pt — diffusion map embedding for every cell
  • coupling.pt — list of OT transport matrices, one per consecutive day-pair
  • segments.pt — the exact (t_start, t_end, source, target) splits used during training, saved so downstream analysis never needs to rerun preprocessing/embedding

Note on memory: the full dataset is 236k cells. max_total_cells in the config stratifies a subsample across all 39 days before any expensive stage runs — increase this once you've confirmed the pipeline runs cleanly end to end.


Interactive dashboard

streamlit run app/streamlit_app.py

Four tabs, each calling only into pipeline/ and visualization/:

  1. Pipeline Runner — load a config, run the full pipeline, watch progress
  2. Embedding — inspect the diffusion map, colorable by any .obs column
  3. OT Diagnostics — live epsilon slider, recompute coupling, watch entropy change in real time
  4. SDE Training & Landscape — learned drift field as arrows over the embedding, plus animated simulated trajectories

Notebooks

notebooks/00_all_visualizations.ipynb — every visualization in the project, in one file:

  • Sections 1–5 run on a fast ~2,000-cell subsample of the real dataset (data exploration, preprocessing, diffusion map with from-scratch eigendecomposition verification, OT epsilon sweep, toy SDE training)
  • Section 6 loads the real, full-scale saved outputs from run_pipeline.py and validates them: coupling entropy across all 38 real segments, real drift field, real trajectory animation, per-segment simulated-vs-real accuracy curve, and a summary readout

Scaling to big data

Every stage has a small-data and big-data mode, switched entirely via config:

Stage Small-data Big-data
Load/preprocess in-memory, dense backed_mode: true + sparse matrices
Diffusion map exact eigendecomposition approximate: true — Nystrom-extended, chunked to bound peak memory
OT pseudotime backend: pot — exact Sinkhorn, O(n²) cost matrix backend: moscot — GPU, low-rank, no dense cost matrix
Neural SDE CPU, small batches device: gpu, larger batches

Testing

pytest tests/ -v

Covers config loading, preprocessing, diffusion map (exact + Nystrom), OT coupling (both backends), the neural SDE's drift/diffusion outputs, and pipeline orchestration (timepoint splitting, error handling on missing columns).


Evaluation metrics

Implemented in evaluation/metrics.py:

  • trustworthiness — does the diffusion map preserve local neighborhood structure from the original PCA space? (O(n²), small-subset only)
  • pseudotime_correlation — Spearman correlation between recovered pseudotime and known biological staging, when available
  • coupling_entropy — how confident vs. diffuse is a given OT transport plan (diagnoses epsilon misconfiguration)
  • simulated_vs_real_distance — does the trained SDE's simulated endpoint distribution actually resemble where real cells ended up at the next timepoint

Known limitations

  • paul15 (scanpy's built-in myeloid progenitor dataset) is a single snapshot, not a time-course — it has no real timepoints and cannot be used with the OT-coupling stage. Only schiebinger19 (or another real multi-timepoint dataset) works for the full pipeline.
  • The Nystrom-approximated diffusion map trades some accuracy for scalability — always sanity-check against an exact run on a subset before trusting results on unfamiliar data.
  • trustworthiness requires the original PCA coordinates, which are not among run_pipeline.py's saved artifacts — computing it post-hoc requires rerunning preprocessing.

References

  • Schiebinger, G. et al. (2019). Optimal-Transport Analysis of Single-Cell Gene Expression Identifies Developmental Trajectories in Reprogramming. Cell.
  • Coifman, R. R. & Lafon, S. (2006). Diffusion maps. Applied and Computational Harmonic Analysis.
  • Cuturi, M. (2013). Sinkhorn Distances: Lightspeed Computation of Optimal Transport. NeurIPS.
  • Li, X. et al. (2020). Scalable Gradients for Stochastic Differential Equations. AISTATS. (the torchsde adjoint method)

About

scRNA-seq Neural Dynamics: Reconstructing cell differentiation as a continuous-time stochastic process from single-cell snapshots using diffusion maps, optimal transport, and neural SDEs.

Resources

Stars

1 star

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages