Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
120 changes: 101 additions & 19 deletions src/pertpy/tools/_milo.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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


Expand Down Expand Up @@ -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."""

Expand Down Expand Up @@ -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`:
Expand Down Expand Up @@ -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(
Expand All @@ -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"]
Expand Down
9 changes: 7 additions & 2 deletions src/pertpy/tools/_milo_glmm.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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:
Expand Down
60 changes: 59 additions & 1 deletion tests/tools/test_milo.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"])
Expand Down Expand Up @@ -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()
Expand Down
Loading