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()