Varix-Dim48-Varix

A Varix variational autoencoder with a 48-dimensional latent space, trained on single-cell RNA-seq data. Unlike Ontix, the latent dimensions are not constrained by gene ontologies, providing a compact general-purpose embedding of the transcriptome.

Part of the autoencodix pretrained model collection.

Usage

Install the autoencodix package:

pip install autoencodix

Download the model and use it to generate embeddings for your own scRNA-seq data (AnnData with genes as Ensembl IDs, log1p-normalized counts):

from huggingface_hub import snapshot_download
import autoencodix as acx
import anndata as ad
import pandas as pd
import numpy as np
import scanpy
from scipy import sparse
from autoencodix.data._numeric_dataset import NumericDataset
from autoencodix.data._datasetcontainer import DatasetContainer
from autoencodix.configs.varix_config import VarixConfig

# Download and load the pretrained model
model_name = "Varix-Dim48-Varix"
repo_id = f"autoencodix/{model_name}"
model_file = snapshot_download(repo_id=repo_id)
model_file = model_file + "/large_varix_final_model_Dim48_Varix.pkl"

loaded_varix = acx.Varix.load(file_path=model_file)
loaded_varix._trainer._config.device = "cpu"  # Switch device if you want to run on CPU or GPU

# Load your data (AnnData with adata.var.index as Ensembl gene IDs)
adata = ad.read_h5ad("path/to/your_data.h5ad")

# Match the input gene space of the pretrained model, zero-padding any missing genes
anndata_template = ad.AnnData(
    X=sparse.csr_matrix(np.zeros((1, len(loaded_varix.result.model.feature_order)))),
    var=pd.DataFrame(index=loaded_varix.result.model.feature_order),
)
adata = ad.concat([anndata_template, adata], axis=0, join="outer").copy()
adata = adata[1:, loaded_varix.result.model.feature_order].copy()

# The model expects log1p normalized counts
scanpy.pp.log1p(adata, copy=False)

test_dataset = NumericDataset(
    data=adata.X,
    config=VarixConfig(),
    sample_ids=adata.obs.index,
    metadata=adata.obs.loc[adata.obs.index, :],
    split_indices=None,
    feature_ids=adata.var.index,
)
acx_container = DatasetContainer(train=None, valid=None, test=test_dataset)

# Generate embeddings
result = loaded_varix.predict(data=acx_container)
df_latent = result.get_latent_df(split="test", epoch=-1)
df_latent

df_latent is a pandas.DataFrame with one row per cell and 48 columns, one per latent dimension.

Further usage

For a full walkthrough including latent space visualization, marker gene explanation with xAI/LLMs, and synthetic data generation, see the tutorial notebook.

Downloads last month

-

Downloads are not tracked for this model. How to track
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support

Collection including autoencodix/Varix-Dim48-Varix