From 370d1cd45a981adc3cb1613f154459d289ac1b19 Mon Sep 17 00:00:00 2001 From: Lukas Heumos Date: Wed, 23 Sep 2026 14:10:48 +0200 Subject: [PATCH] Make the pydeseq2 Milo solver reproduce R Milo The pydeseq2 solver tested with the Wald test of DESeq2, whose fitted means are floored at 0.5. Both cost almost all power in neighbourhoods that one group barely populates, which are common in Milo, so on the Stephenson data the solver found none of the 144 Asymptomatic neighbourhoods miloR calls. The solver now normalises like R Milo with library size times TMM factors, ported from calcNormFactors of edgeR and matching it to machine precision, keeps the MAP dispersions of pydeseq2 and tests with a likelihood ratio test fitted without the floor. The reported log fold change comes from a refit with the prior count of edgeR, so it stays finite when a group has no cells. Cook's outlier handling is dropped because edgeR has none. logCPM now is a log CPM instead of the baseMean of DESeq2, shared with the mixed model. On identical neighbourhood counts of the Stephenson data the overlap of significant neighbourhoods with miloR rises from 0.79-0.81 to 0.84-0.88 (Jaccard) and to 71 of miloR's 144 Asymptomatic calls. On a benchmark built from the same neighbourhoods its power matches miloR (recall 0.48 against 0.50) while it makes fewer false calls under null data (3.8 against 14.8 with permuted labels, 3.1 against 5.9 with ten against ten healthy donors). The remaining gap is the dispersion estimate, which would need edgeR itself. test_da_nhoods_default_contrast compares p-values instead of SpatialFDR, which is constant on its six sample fixture with the new test. --- src/pertpy/tools/_milo.py | 120 +++++++++++++++++++++++++++------ src/pertpy/tools/_milo_glmm.py | 9 ++- tests/tools/test_milo.py | 60 ++++++++++++++++- 3 files changed, 167 insertions(+), 22 deletions(-) diff --git a/src/pertpy/tools/_milo.py b/src/pertpy/tools/_milo.py index 7647989d..30d03a70 100644 --- a/src/pertpy/tools/_milo.py +++ b/src/pertpy/tools/_milo.py @@ -20,7 +20,7 @@ from pertpy._doc import _doc_params, doc_common_plot_args from pertpy._logger import logger from pertpy._types import CSBase, cast_frame, cast_matrix -from pertpy.tools._milo_glmm import fit_nb_glmm_nhoods, parse_random_effects, random_effect_matrices +from pertpy.tools._milo_glmm import fit_nb_glmm_nhoods, log_cpm, parse_random_effects, random_effect_matrices if TYPE_CHECKING: from collections.abc import Collection, Sequence @@ -30,7 +30,9 @@ from matplotlib.figure import Figure from numpy.typing import ArrayLike +from scipy.linalg import null_space from scipy.sparse import coo_matrix, csr_matrix, issparse, spmatrix +from scipy.stats import chi2, false_discovery_control, rankdata from sklearn.metrics.pairwise import euclidean_distances @@ -96,6 +98,76 @@ def _weighted_bh(pvalues: np.ndarray, weights: np.ndarray) -> np.ndarray: return out +def _tmm_factors( + counts: np.ndarray, lib_size: np.ndarray, *, log_ratio_trim: float = 0.3, sum_trim: float = 0.05 +) -> np.ndarray: + """TMM normalisation factors of a features x samples count matrix, ported from ``calcNormFactors`` of edgeR.""" + counts = np.asarray(counts, dtype=float) + counts = counts[(counts > 0).any(axis=1)] + lib_size = np.asarray(lib_size, dtype=float) + n_samples = counts.shape[1] + if counts.shape[0] == 0 or n_samples == 1: + return np.ones(n_samples) + + upper_quartiles = np.quantile(counts, 0.75, axis=0) / lib_size + if np.median(upper_quartiles) < 1e-20: + ref = int(np.argmax(np.sqrt(counts).sum(axis=0))) + else: + ref = int(np.argmin(np.abs(upper_quartiles - upper_quartiles.mean()))) + + factors = np.ones(n_samples) + for i in range(n_samples): + with np.errstate(divide="ignore", invalid="ignore"): + obs, reference = counts[:, i] / lib_size[i], counts[:, ref] / lib_size[ref] + log_ratio = np.log2(obs / reference) + abs_expr = (np.log2(obs) + np.log2(reference)) / 2 + variance = (lib_size[i] - counts[:, i]) / lib_size[i] / counts[:, i] + variance += (lib_size[ref] - counts[:, ref]) / lib_size[ref] / counts[:, ref] + finite = np.isfinite(log_ratio) & np.isfinite(abs_expr) + log_ratio, abs_expr, variance = log_ratio[finite], abs_expr[finite], variance[finite] + if log_ratio.size == 0 or np.max(np.abs(log_ratio)) < 1e-6: + continue + n = log_ratio.size + lo_ratio, lo_expr = np.floor(n * log_ratio_trim) + 1, np.floor(n * sum_trim) + 1 + rank_ratio, rank_expr = rankdata(log_ratio), rankdata(abs_expr) + keep = (rank_ratio >= lo_ratio) & (rank_ratio <= n + 1 - lo_ratio) + keep &= (rank_expr >= lo_expr) & (rank_expr <= n + 1 - lo_expr) + with np.errstate(invalid="ignore"): + log_factor = np.sum(log_ratio[keep] / variance[keep]) / np.sum(1 / variance[keep]) + factors[i] = 2 ** np.nan_to_num(log_factor) + return factors / np.exp(np.mean(np.log(factors))) + + +def _nb_lrt( + counts: np.ndarray, + lib_size: np.ndarray, + design: np.ndarray, + contrast: np.ndarray, + dispersions: np.ndarray, + *, + prior_count: float = 0.125, +) -> tuple[np.ndarray, np.ndarray]: + """Log2 fold change of ``contrast`` and its negative binomial likelihood ratio test p-value for every neighbourhood. + + Like the test of edgeR in R Milo and unlike a Wald test, a likelihood ratio test keeps its power in neighbourhoods that a group barely populates. + The fits therefore drop the floor of 0.5 that pydeseq2 puts on fitted means, which would cap the fold change in exactly those neighbourhoods. + As in edgeR, the fold change comes from a refit with ``prior_count`` added in proportion to the library sizes, which keeps it finite when a group has no cells. + """ + from pydeseq2.utils import irls_solver, nb_nll + + reduced = design @ null_space(contrast[None, :]) + prior = prior_count * lib_size / lib_size.mean() + logfc = np.empty(len(counts)) + statistic = np.empty(len(counts)) + for i, (y, dispersion) in enumerate(zip(counts, dispersions, strict=True)): + _, mu_full, *_ = irls_solver(y, lib_size, design, dispersion, min_mu=1e-6) + _, mu_reduced, *_ = irls_solver(y, lib_size, reduced, dispersion, min_mu=1e-6) + statistic[i] = 2 * (nb_nll(y, mu_reduced, dispersion) - nb_nll(y, mu_full, dispersion)) + shrunk, *_ = irls_solver(y + prior, lib_size + 2 * prior, design, dispersion, min_mu=1e-6) + logfc[i] = contrast @ shrunk / np.log(2) + return logfc, chi2.sf(np.maximum(statistic, 0), df=1) + + class Milo: """Python implementation of Milo.""" @@ -356,8 +428,8 @@ def da_nhoods( tol: Convergence tolerance of a mixed model fit. solver: The solver to fit the model to, ignored for a mixed model. The "edger" solver requires R, rpy2 and edgeR to be installed and reproduces the R implementation. - The "pydeseq2" requires pydeseq2 to be installed. - It uses the Wald test of DESeq2 instead of the quasi-likelihood F-test of edgeR, so its results are close to but not identical with those of R Milo. + The "pydeseq2" solver requires pydeseq2 but not R. + It normalises like R Milo, estimates the dispersions with pydeseq2 and tests with a likelihood ratio test, which comes close to the quasi-likelihood F-test of edgeR without reproducing it exactly. Returns: None, modifies `milo_mdata['milo']` in place, adding the results of the DA test to `.var`: @@ -560,14 +632,10 @@ def da_nhoods( if find_spec("pydeseq2") is None: raise ImportError("pydeseq2 is required but not installed. Install with: pip install pydeseq2") - import warnings - from pydeseq2.dds import DeseqDataSet - from pydeseq2.ds import DeseqStats - - warnings.filterwarnings("always", message=".*(alpha).*") counts_filtered = count_mat[np.ix_(keep_nhoods, keep_smp)] + lib_size_filtered = lib_size[keep_smp] design_df_filtered = design_df.copy() design_df_filtered = design_df_filtered.astype( @@ -578,28 +646,42 @@ def da_nhoods( counts=pd.DataFrame(counts_filtered.T, index=design_df_filtered.index), metadata=design_df_filtered, design=fixed if fixed.startswith("~") else f"~{fixed}", - refit_cooks=True, - size_factors_fit_type="poscounts", ) design_matrix = cast_frame(dds.obsm["design_matrix"]) _check_residual_df(design_matrix, design) - dds.deseq2() + + effective_lib_size = lib_size_filtered * _tmm_factors(counts_filtered, lib_size_filtered) + size_factors = effective_lib_size / np.exp(np.mean(np.log(effective_lib_size))) + normed_counts = counts_filtered.T / size_factors[:, None] + dds.obs["size_factors"] = size_factors + dds.layers["normed_counts"] = normed_counts + dds.var["_normed_means"] = normed_counts.mean(axis=0) + dds.fit_genewise_dispersions() + dds.fit_dispersion_trend() + dds.fit_dispersion_prior() + dds.fit_MAP_dispersions() contrast = ( _contrast_vector(list(design_matrix.columns), model_contrasts, reference_levels) if model_contrasts is not None else np.eye(design_matrix.shape[1])[-1] ) - stat_res = DeseqStats(dds, contrast=contrast) - stat_res.summary() - res = stat_res.results_df - - res = res.rename( - columns={"baseMean": "logCPM", "log2FoldChange": "logFC", "pvalue": "PValue", "padj": "FDR"} + logfc, pvalues = _nb_lrt( + counts_filtered, + effective_lib_size, + design_matrix.to_numpy(dtype=float), + contrast, + cast_frame(dds.var)["dispersions"].to_numpy(), + ) + res = pd.DataFrame( + { + "logFC": logfc, + "logCPM": log_cpm(counts_filtered), + "PValue": pvalues, + "FDR": false_discovery_control(pvalues), + } ) - - res = res[["logCPM", "logFC", "PValue", "FDR"]] res.index = sample_adata.var_names[keep_nhoods] written = [*res.columns, "SpatialFDR"] diff --git a/src/pertpy/tools/_milo_glmm.py b/src/pertpy/tools/_milo_glmm.py index 91f595d8..d581ec12 100644 --- a/src/pertpy/tools/_milo_glmm.py +++ b/src/pertpy/tools/_milo_glmm.py @@ -357,6 +357,12 @@ def pseudo_likelihood( ) +def log_cpm(counts: np.ndarray) -> np.ndarray: + """Log2 of the mean counts per million of every neighbourhood across samples.""" + library_size = counts.sum(axis=0) + return np.log2(np.mean(counts / np.where(library_size > 0, library_size, 1), axis=1) * 1e6 + 1e-12) + + def fit_nb_glmm_nhoods( counts: np.ndarray, X: np.ndarray, @@ -373,8 +379,7 @@ def fit_nb_glmm_nhoods( The reported log fold change is ``contrast`` applied to the fixed effects, defaulting to the last column of the model matrix, which is the coefficient the edgeR solver tests. Its p-value comes from a t-test whose degrees of freedom follow :func:`between_within_df`. """ - library_size = counts.sum(axis=0) - logcpm = np.log2(np.mean(counts / np.where(library_size > 0, library_size, 1), axis=1) * 1e6 + 1e-12) + logcpm = log_cpm(counts) weights = np.zeros(X.shape[1]) if contrast is None else np.asarray(contrast, dtype=float) if contrast is None: diff --git a/tests/tools/test_milo.py b/tests/tools/test_milo.py index c5c04d25..5c18c9bf 100644 --- a/tests/tools/test_milo.py +++ b/tests/tools/test_milo.py @@ -9,6 +9,7 @@ from mudata import MuData import pertpy as pt +from pertpy.tools._milo import _nb_lrt, _tmm_factors @pytest.fixture(params=["edger", "pydeseq2"]) @@ -234,10 +235,67 @@ def test_da_nhoods_default_contrast(da_nhoods_mdata, milo, solver): milo.da_nhoods(mdata, design="~condition", model_contrasts="conditionConditionB-conditionConditionA", solver=solver) contr_results = mdata["milo"].var.copy() - assert np.corrcoef(contr_results["SpatialFDR"], default_results["SpatialFDR"])[0, 1] > 0.99 + assert np.corrcoef(contr_results["PValue"], default_results["PValue"])[0, 1] > 0.99 assert np.corrcoef(contr_results["logFC"], default_results["logFC"])[0, 1] > 0.99 +def test_da_nhoods_pydeseq2_reproduces_edger(da_nhoods_mdata, milo): + pytest.importorskip("rpy2") + try: + from rpy2.robjects.packages import importr + + importr("edgeR") + except Exception: # noqa: BLE001 + pytest.skip("Required R package 'edgeR' not available") + mdata = da_nhoods_mdata.copy() + milo.da_nhoods(mdata, design="~condition", solver="edger") + edger = mdata["milo"].var[["logFC", "PValue"]].copy() + milo.da_nhoods(mdata, design="~condition", solver="pydeseq2") + pydeseq2 = mdata["milo"].var[["logFC", "PValue"]] + + assert np.corrcoef(edger["logFC"], pydeseq2["logFC"])[0, 1] > 0.99 + assert edger["PValue"].corr(pydeseq2["PValue"], method="spearman") > 0.9 + + +def test_tmm_factors_match_edger(): + counts = np.array( + [ + [10, 12, 30, 0], + [0, 3, 5, 8], + [25, 20, 18, 40], + [7, 0, 2, 9], + [100, 80, 120, 95], + [3, 6, 0, 1], + [45, 50, 38, 60], + [0, 0, 4, 2], + [15, 22, 17, 13], + [60, 30, 75, 55], + [8, 11, 9, 0], + [33, 27, 41, 36], + [0, 0, 0, 0], + ] + ) + # calcNormFactors(counts, lib.size=..., method="TMM") of edgeR 4.8.2 + np.testing.assert_allclose( + _tmm_factors(counts, counts.sum(0)), + [0.977957934816150, 1.193511265666823, 0.941900259756086, 0.909595673308097], + ) + np.testing.assert_allclose( + _tmm_factors(counts, counts.sum(0) * np.array([1, 2, 1, 3])), + [1.588354120588916, 0.759830740837740, 1.529790909716994, 0.541631270583240], + ) + + +def test_nb_lrt_keeps_power_when_a_group_has_no_cells(): + """A Wald test loses its power when a group has no cells, the likelihood ratio test that R Milo relies on does not.""" + counts = np.array([[0, 0, 0, 0, 12, 15, 9, 14]], dtype=float) + design = np.column_stack([np.ones(8), np.repeat([0.0, 1.0], 4)]) + logfc, pvalues = _nb_lrt(counts, np.full(8, 1000.0), design, np.array([0.0, 1.0]), np.array([0.1])) + + assert pvalues[0] < 1e-3 + assert 3 < logfc[0] < 10 + + @pytest.fixture def three_condition_mdata(adata, milo): adata = adata.copy()