diff --git a/src/pertpy/tools/_coda/_base_coda.py b/src/pertpy/tools/_coda/_base_coda.py index e32bb9f5..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(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, } diff --git a/src/pertpy/tools/_coda/_sccoda.py b/src/pertpy/tools/_coda/_sccoda.py index def6a063..b2035fa0 100644 --- a/src/pertpy/tools/_coda/_sccoda.py +++ b/src/pertpy/tools/_coda/_sccoda.py @@ -283,7 +283,9 @@ def model( ) # Calculate DM-distributed counts - predictions = npy.sample("counts", npd.DirichletMultinomial(concentrations, n_total), obs=counts) + 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..81b26077 100644 --- a/src/pertpy/tools/_coda/_tasccoda.py +++ b/src/pertpy/tools/_coda/_tasccoda.py @@ -454,7 +454,9 @@ def model( ) # Calculate DM-distributed counts - predictions = npy.sample("counts", npd.DirichletMultinomial(concentrations, n_total), obs=counts) + predictions = npy.sample( + "counts", npd.DirichletMultinomial(concentrations, n_total, validate_args=False), obs=counts + ) return predictions