Skip to content

Quick start

from prism import PRISMConfig, PRISMTrainer, generate_prism_data

data, truth = generate_prism_data(n_subjects=50, n_genes=60, n_cells_per_subject=200,
                                  rho=0.6, seed=42, device="cpu")

cfg = PRISMConfig(n_genes=data.n_genes, n_covars=data.n_covars, n_context=data.n_context,
                  ot_project_q=True, auto_q_prior=False,  # PRISM-OT, the configuration used in the paper
                  max_em_iter=30, wandb_enabled=False, device="cpu")
res = PRISMTrainer(cfg).fit(data)

print("general DE (FDR < 0.05):", int((res.q_values_de < 0.05).sum()))
print("context DE (FDR < 0.05):", int((res.q_values_context.min(dim=1).values < 0.05).sum()))

The fit takes about 4 minutes on a 4-core CPU. wandb_enabled=False turns off Weights & Biases logging, which is on by default whenever wandb is installed.

Your data

PrismData.from_anndata reads raw counts from adata.X and per-cell columns from adata.obs. Include an intercept column among the covariates, and standardise continuous covariates and context variables.

from prism import PrismData

adata.obs["intercept"] = 1.0
data = PrismData.from_anndata(adata, covar_cols=["intercept", "age", "sex"], context_cols=["ctx"],
                              condition_col="disease", subject_col="donor", device="cuda")

Then fit it with the configuration above, using device="cuda".

Outputs

field shape meaning
res.q_hat (cells,) posterior probability that each cell is in the disease state (0 for control donors)
res.rho_hat (disease donors,) fraction of each disease donor's cells in the disease state
res.alpha_de_hat (genes,) constant disease effect (natural-log fold change)
res.theta_hat (genes, context dims) context modulation of the disease effect
res.q_values_de (genes,) BH-adjusted p-values for H0: alpha = 0 (general DE)
res.q_values_context (genes, context dims) BH-adjusted p-values for H0: theta = 0 (context DE)

res.de_summary(gene_names=...) returns the per-gene estimates and tests as a pandas DataFrame.