From c83e36ec7cca98d85cc1c3ade7ba056b8030bdc0 Mon Sep 17 00:00:00 2001 From: Lukas Heumos Date: Fri, 18 Sep 2026 13:02:35 +0200 Subject: [PATCH 1/3] Keep scCODA and tascCODA working under numpyro 0.22 numpyro 0.22 turns distribution validation on by default. scCODA replaces zero counts with a pseudocount of 0.5, so `counts` and the derived `n_total` fall outside the discrete DirichletMultinomial support: sampling now raises and NUTS initialization fails with "Cannot find valid initial parameters". The log-prob itself is well defined for fractional counts, so the likelihood site opts out of validation, and the prior predictive rounds the number of trials, which is the only well-defined choice when actually drawing counts. Verified against numpyro 0.22.0 in the pertpy conda env: all 11 tests in tests/tools/_coda pass, where run_nuts and make_arviz failed before. --- src/pertpy/tools/_coda/_base_coda.py | 3 ++- src/pertpy/tools/_coda/_sccoda.py | 5 ++++- src/pertpy/tools/_coda/_tasccoda.py | 5 ++++- 3 files changed, 10 insertions(+), 3 deletions(-) diff --git a/src/pertpy/tools/_coda/_base_coda.py b/src/pertpy/tools/_coda/_base_coda.py index e32bb9f5..e5ed6086 100644 --- a/src/pertpy/tools/_coda/_base_coda.py +++ b/src/pertpy/tools/_coda/_base_coda.py @@ -137,7 +137,8 @@ def _build_arviz_from_adata( predict_kwargs = { "counts": None, "covariates": jnp.array(sample_adata.obsm["covariate_matrix"], dtype="float64"), - "n_total": jnp.array(sample_adata.obsm["sample_counts"], dtype="float64"), + # Sampling needs a whole number of trials, unlike the log-prob, which tolerates the 0.5 pseudocounts. + "n_total": jnp.array(np.rint(cast_dense(sample_adata.obsm["sample_counts"])), dtype="float64"), "ref_index": jnp.array(sample_adata.uns["scCODA_params"]["reference_index"]), "sample_adata": sample_adata, } diff --git a/src/pertpy/tools/_coda/_sccoda.py b/src/pertpy/tools/_coda/_sccoda.py index def6a063..6663012e 100644 --- a/src/pertpy/tools/_coda/_sccoda.py +++ b/src/pertpy/tools/_coda/_sccoda.py @@ -283,7 +283,10 @@ def model( ) # Calculate DM-distributed counts - predictions = npy.sample("counts", npd.DirichletMultinomial(concentrations, n_total), obs=counts) + # Pseudocounts of 0.5 leave `counts` outside the discrete support, which the well-defined log-prob tolerates but validation does not. + predictions = npy.sample( + "counts", npd.DirichletMultinomial(concentrations, n_total, validate_args=False), obs=counts + ) return predictions diff --git a/src/pertpy/tools/_coda/_tasccoda.py b/src/pertpy/tools/_coda/_tasccoda.py index 5abafcbb..616e34c4 100644 --- a/src/pertpy/tools/_coda/_tasccoda.py +++ b/src/pertpy/tools/_coda/_tasccoda.py @@ -454,7 +454,10 @@ def model( ) # Calculate DM-distributed counts - predictions = npy.sample("counts", npd.DirichletMultinomial(concentrations, n_total), obs=counts) + # Pseudocounts of 0.5 leave `counts` outside the discrete support, which the well-defined log-prob tolerates but validation does not. + predictions = npy.sample( + "counts", npd.DirichletMultinomial(concentrations, n_total, validate_args=False), obs=counts + ) return predictions From 34a3d6c29c5a89fc95a65a16b233949d0b55ab3a Mon Sep 17 00:00:00 2001 From: Lukas Heumos Date: Fri, 18 Sep 2026 13:06:07 +0200 Subject: [PATCH 2/3] Drop the comments from the numpyro 0.22 fix --- src/pertpy/tools/_coda/_base_coda.py | 1 - src/pertpy/tools/_coda/_sccoda.py | 1 - src/pertpy/tools/_coda/_tasccoda.py | 1 - 3 files changed, 3 deletions(-) diff --git a/src/pertpy/tools/_coda/_base_coda.py b/src/pertpy/tools/_coda/_base_coda.py index e5ed6086..9d291ff5 100644 --- a/src/pertpy/tools/_coda/_base_coda.py +++ b/src/pertpy/tools/_coda/_base_coda.py @@ -137,7 +137,6 @@ def _build_arviz_from_adata( predict_kwargs = { "counts": None, "covariates": jnp.array(sample_adata.obsm["covariate_matrix"], dtype="float64"), - # Sampling needs a whole number of trials, unlike the log-prob, which tolerates the 0.5 pseudocounts. "n_total": jnp.array(np.rint(cast_dense(sample_adata.obsm["sample_counts"])), dtype="float64"), "ref_index": jnp.array(sample_adata.uns["scCODA_params"]["reference_index"]), "sample_adata": sample_adata, diff --git a/src/pertpy/tools/_coda/_sccoda.py b/src/pertpy/tools/_coda/_sccoda.py index 6663012e..b2035fa0 100644 --- a/src/pertpy/tools/_coda/_sccoda.py +++ b/src/pertpy/tools/_coda/_sccoda.py @@ -283,7 +283,6 @@ def model( ) # Calculate DM-distributed counts - # Pseudocounts of 0.5 leave `counts` outside the discrete support, which the well-defined log-prob tolerates but validation does not. predictions = npy.sample( "counts", npd.DirichletMultinomial(concentrations, n_total, validate_args=False), obs=counts ) diff --git a/src/pertpy/tools/_coda/_tasccoda.py b/src/pertpy/tools/_coda/_tasccoda.py index 616e34c4..81b26077 100644 --- a/src/pertpy/tools/_coda/_tasccoda.py +++ b/src/pertpy/tools/_coda/_tasccoda.py @@ -454,7 +454,6 @@ def model( ) # Calculate DM-distributed counts - # Pseudocounts of 0.5 leave `counts` outside the discrete support, which the well-defined log-prob tolerates but validation does not. predictions = npy.sample( "counts", npd.DirichletMultinomial(concentrations, n_total, validate_args=False), obs=counts ) From d2aaab6872f0d308c5267ebf8a98907a6cec578b Mon Sep 17 00:00:00 2001 From: Lukas Heumos Date: Fri, 18 Sep 2026 13:10:24 +0200 Subject: [PATCH 3/3] Round the predictive trial counts without a cast --- src/pertpy/tools/_coda/_base_coda.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/pertpy/tools/_coda/_base_coda.py b/src/pertpy/tools/_coda/_base_coda.py index 9d291ff5..250209fa 100644 --- a/src/pertpy/tools/_coda/_base_coda.py +++ b/src/pertpy/tools/_coda/_base_coda.py @@ -137,7 +137,7 @@ def _build_arviz_from_adata( predict_kwargs = { "counts": None, "covariates": jnp.array(sample_adata.obsm["covariate_matrix"], dtype="float64"), - "n_total": jnp.array(np.rint(cast_dense(sample_adata.obsm["sample_counts"])), dtype="float64"), + "n_total": jnp.rint(jnp.array(sample_adata.obsm["sample_counts"], dtype="float64")), "ref_index": jnp.array(sample_adata.uns["scCODA_params"]["reference_index"]), "sample_adata": sample_adata, }