Skip to content

Repository files navigation

RhinoNet

Overview of the RhinoNet framework

Geometry-aware representation learning and conditional generation for variable-size single-cell populations

This repository contains the implementation of RhinoNet, a framework for learning population-level representations of single-cell point clouds and generating condition-specific cellular populations. RhinoNet combines graph or simplicial wavelet features with an autoencoding objective, preserves externally estimated population geometry through PhenoGS regularization, and generates new point clouds using conditional flow matching in the learned latent space.

Installation

RhinoNet requires Python 3.11 or newer. The locked environment can be installed with uv:

git clone <ANONYMOUS_REPOSITORY_URL>
cd RhinoNet
uv sync
source .venv/bin/activate

To use a pre-existing environment, install the project dependencies from pyproject.toml. GPU training requires a CUDA-compatible PyTorch installation. The supplied Slurm launchers request one generic GPU and contain no site-specific account, partition, or hardware names.

Input data format

RhinoNet accepts a cohort of variable-size point clouds. A population is a matrix

X_i: [N_i, D]

where N_i is the number of cells in population i and D is the shared feature dimension. Different populations may contain different numbers of cells, but every population must use the same features in the same column order. Inputs should be finite, numeric, and stored as float32. RhinoNet uses list-based batching, so populations do not need to be padded or subsampled to a common size.

Non-spatial populations: NPZ input

unsupervised_main.py, latent export, conditional generation, and interpolation consume a NumPy .npz population cache. Its fields are:

Key Shape/type Required for Description
populations Object array of length M; element i is float32 [N_i,D] All workflows The ordered collection of non-empty cell populations.
labels Integer array [M] Optional Population labels. If omitted, the representation loader assigns zeros.
num_labels Scalar or one-element integer array Optional Number of label classes. If omitted, it is inferred from labels.
group_keys String array [C] Generation/interpolation Ordered names of categorical metadata fields, such as Patient and Treatment.
group_names String array [M] Generation/interpolation One metadata record per population. Values follow group_keys order and are joined by __.

For example, if

group_keys = np.array(["Patient", "Treatment"])

then a corresponding entry in group_names can be "21__DMSO". The first component is interpreted as Patient=21 and the second as Treatment=DMSO.

A minimal cache can be created as follows:

import numpy as np

clouds = [
    np.random.randn(800, 44).astype(np.float32),
    np.random.randn(1250, 44).astype(np.float32),
    np.random.randn(600, 44).astype(np.float32),
]

populations = np.empty(len(clouds), dtype=object)
populations[:] = clouds

np.savez(
    "data/populations.npz",
    populations=populations,
    labels=np.array([0, 1, 1], dtype=np.int64),
    num_labels=np.array([2], dtype=np.int64),
    group_keys=np.array(["Patient", "Treatment"]),
    group_names=np.array([
        "P01__Control",
        "P01__DrugA",
        "P02__DrugA",
    ]),
)

The archive can be checked before training:

import numpy as np

with np.load("data/populations.npz", allow_pickle=True) as data:
    clouds = [np.asarray(x, dtype=np.float32) for x in data["populations"]]
    assert clouds
    assert all(x.ndim == 2 and x.shape[0] > 0 for x in clouds)
    assert len({x.shape[1] for x in clouds}) == 1
    assert all(np.isfinite(x).all() for x in clouds)

    if "group_names" in data:
        assert len(data["group_names"]) == len(clouds)
        assert len(data["group_keys"]) > 0

print(f"{len(clouds)} populations; feature dimension {clouds[0].shape[1]}")

Spatial populations: H5AD input

spatial_main.py expects one .h5ad file per tissue region or population. Files are discovered in lexical filename order from

<data_root>/<dataset>/*.h5ad

Each file must contain:

AnnData field Shape/type Description
adata.X [N_i,D] Cell-by-feature matrix used as the gene or marker view and as node features for the spatial view.
adata.obsm["spatial"] [N_i,2] Two-dimensional spatial coordinates aligned row-for-row with adata.X.
adata.uns["patient_code"] String Population-level patient identifier.
adata.uns[<label_name>] Numeric scalar Population metadata selected by --label_name, for example outcome or recurrence. It may be non-finite only when --require_finite_label is not used.

The number and ordering of cells in adata.X and adata.obsm["spatial"] must match. Patient identifiers are also checked against the _c###_ token in each filename. Because the model reads adata.X directly, the H5AD files must already contain the feature representation intended for training.

Spatial runs additionally require an ordered NPZ cache produced by scripts/pheno_gs/cache_space_gm_h5ad.py. It contains:

  • populations: the point clouds used to compute PhenoGS;
  • sample_names: H5AD filenames in lexical order;
  • patients and outcomes: population metadata;
  • feature_scaler_mean and feature_scaler_scale: training-cohort preprocessing statistics;
  • scaler_holdout_patients: the patients excluded when fitting those statistics.

sample_names must match the lexically sorted H5AD filenames exactly.

Population-distance matrix

When the PhenoGS regularizer is enabled, RhinoNet also requires a NumPy .npy matrix

D_pheno: [M, M]

with one row and column per population. It must be finite, symmetric, and ordered exactly like populations for non-spatial data or the sorted H5AD files for spatial data. Entry D_pheno[i,j] is the target distance between populations i and j. The matrix is supplied with --phenoGS_cache; its influence is controlled by --dist_weight.

The included PhenoGS implementation is located at third_party/phenogs/phenoGS_bms.py, with preparation utilities under scripts/pheno_gs/ and portable launchers under slurm/pheno_gs/.

Ordering requirements

The following objects must describe populations in exactly the same order:

  1. the populations array or sorted H5AD filenames;
  2. population-level labels and metadata;
  3. rows and columns of the PhenoGS distance matrix;
  4. rows of latent_representations.npy.

Reordering only one artifact silently changes which metadata and distances belong to a population. Preserve the original order throughout preprocessing, training, latent export, generation, and evaluation.

No datasets are distributed in this repository.

Reproducing representation learning

The reference launchers are:

Dataset configuration Entry point Default representation
Non-spatial population cohort slurm/training/run_pdo_2m.sbatch K=1, graph wavelets, mean pooling, 32-dimensional latent
Spatial cohort A slurm/training/run_charville.sbatch K=2, geometric Hodge Laplacian, 32-dimensional latent
Spatial cohort B slurm/training/run_upmc.sbatch K=2, geometric Hodge Laplacian, 16-dimensional latent

After placing the prepared caches and distance matrices at the paths documented inside a launcher, submit with:

sbatch slurm/training/run_pdo_2m.sbatch
sbatch slurm/training/run_charville.sbatch
sbatch slurm/training/run_upmc.sbatch

The launchers use AdamW for representation learning, a learning rate of 1e-3, weight decay 3e-3, two wavelet scales, global mean pooling, and a PhenoGS weight of 0.1, unless edited explicitly. The best checkpoint is selected by validation loss. Each output directory contains the selected checkpoint, configuration summaries, and latent_representations.npy.

For custom datasets, unsupervised_main.py is the non-spatial entry point and spatial_main.py is the dual-view spatial entry point. Run either script with --help for the complete argument list.

Exporting latents from a checkpoint

Latents can be re-exported without retraining:

python export_rhinonet_latents.py \
  --population_cache data/populations.npz \
  --checkpoint checkpoints/experiment/best_model.pt \
  --output_dir results/latents \
  --latent_dim 32 \
  --K 1 \
  --J 2 \
  --sigma 32 \
  --threshold 0.1

The architectural arguments must match those used to train the checkpoint. The exporter writes:

  • latent_representations.npy
  • embeddings.npy
  • labels.npy
  • latent_export_summary.json

Conditional point-cloud generation

The portable generation launcher trains or loads the stochastic decoder, trains conditional latent flow matching, and generates one population for each held-out condition:

export POPULATION_CACHE="$PWD/data/populations.npz"
export LATENTS="$PWD/results/latents/latent_representations.npy"
export CONDITION_FIELDS="Patient,Treatment"
export OUTPUT_DIR="$PWD/results/conditional_generation"

sbatch slurm/generation/run_generation.sbatch

For recurrence-conditioned cohorts, use a corresponding metadata field, for example:

export CONDITION_FIELDS="Patient,Recurrence"

By default, the launcher trains the decoder for 400 epochs and the latent flow for 30,000 epochs. These values can be changed with DECODER_EPOCHS and FLOW_EPOCHS. The generator preserves the number of cells in each corresponding held-out population and writes:

  • generated_clouds.pkl, containing generated and real held-out point clouds;
  • metrics.json, containing the split, conditioning values, pooled 1-NNA, and condition-stratified 1-NNA;
  • the decoder checkpoint specified by DECODER_PATH.

The conditional split is performed at the population level. At least one training example is retained for every value of every conditioning field so that each categorical embedding is estimable.

Latent graph-path interpolation

interpolate_treatment.py evaluates whether real populations lying between two endpoint populations can be reconstructed from the learned latent geometry. It constructs a weighted k-nearest-neighbor graph over population latents, selects a shortest path from a control endpoint to a second group, removes the path-interior populations from decoder training, and reconstructs those held-out populations from their latent representations.

Example:

python interpolate_treatment.py \
  --population_cache data/populations.npz \
  --latents results/latents/latent_representations.npy \
  --patient 21 \
  --group_field Treatment \
  --control_treatment DMSO \
  --graph_knn 20 \
  --population_knn 20 \
  --decoder_epochs 400 \
  --decoder_path results/interpolation/decoder.pth \
  --output_dir results/interpolation \
  --seed 0 \
  --n_eval_draws 5

This is a graph-path reconstruction experiment, rather than linear interpolation between endpoint vectors. Its outputs include the selected path, held-out intermediate point clouds, Chamfer and EMD reconstruction metrics, and population- and cell-level visualizations.

Evaluation

The repository provides shared metric implementations in utils/metrics.py and utils/population_generation.py. The generation pipeline reports 1-nearest-neighbor accuracy under both Chamfer distance and Earth Mover's Distance. A value near 50% indicates that a nearest-neighbor classifier cannot readily distinguish generated populations from real populations; values far from 50% indicate separability. For graph-path interpolation, Chamfer distance and EMD are computed between each generated intermediate population and its corresponding observed population over independent decoder-noise draws.

Evaluation should use held-out populations and preserve the same feature preprocessing, population identities, and metadata alignment used during training. Reported comparisons should state whether generation is conditional, how populations are matched, and whether cell counts are preserved or sampled.

Reproducibility

Random seeds are exposed by the Python entry points and stored with generated outputs. To reproduce an experiment:

  1. preserve the population ordering in every cache and distance matrix;
  2. use identical preprocessing statistics across training and held-out data;
  3. record the complete launcher and output-directory configuration;
  4. evaluate only the checkpoint selected by validation loss;
  5. repeat stochastic decoder and flow experiments across multiple seeds.

The default reference settings are encoded directly in the Slurm launchers rather than hidden in cluster-specific configuration files.

Repository structure

.
├── models/                         # Graph/simplicial wavelets, autoencoders, decoder, and flow models
├── scripts/
│   ├── pheno_gs/                   # Population-cache and PhenoGS preparation
│   └── preprocessing/              # Dataset preprocessing
├── slurm/
│   ├── generation/                 # Conditional latent-flow generation
│   ├── pheno_gs/                   # PhenoGS distance computation
│   └── training/                   # Reference representation-learning runs
├── third_party/phenogs/            # PhenoGS implementation used by the pipeline
├── utils/                          # Data loading, visualization, memory bank, and metrics
├── unsupervised_main.py            # Non-spatial representation learning
├── spatial_main.py                 # Dual-view spatial representation learning
├── conditional_latent_flow.py      # Conditional latent generation and evaluation
├── interpolate_treatment.py        # Latent graph-path reconstruction
└── export_rhinonet_latents.py      # Checkpoint-based latent export

About

No description, website, or topics provided.

Resources

Stars

3 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages