acx-pretrained-models
Collection
Collection of larger autoencoder models pretrained on scRNASeq data using Autoencodix package • 15 items • Updated
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.
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.
For a full walkthrough including latent space visualization, marker gene explanation with xAI/LLMs, and synthetic data generation, see the tutorial notebook.