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.
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/activateTo 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.
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.
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_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;patientsandoutcomes: population metadata;feature_scaler_meanandfeature_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.
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/.
The following objects must describe populations in exactly the same order:
- the
populationsarray or sorted H5AD filenames; - population-level labels and metadata;
- rows and columns of the PhenoGS distance matrix;
- 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.
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.sbatchThe 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.
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.1The architectural arguments must match those used to train the checkpoint. The exporter writes:
latent_representations.npyembeddings.npylabels.npylatent_export_summary.json
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.sbatchFor 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.
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 5This 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.
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.
Random seeds are exposed by the Python entry points and stored with generated outputs. To reproduce an experiment:
- preserve the population ordering in every cache and distance matrix;
- use identical preprocessing statistics across training and held-out data;
- record the complete launcher and output-directory configuration;
- evaluate only the checkpoint selected by validation loss;
- 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.
.
├── 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
