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
80 changes: 28 additions & 52 deletions src/pertpy/tools/_milo.py
Original file line number Diff line number Diff line change
Expand Up @@ -34,20 +34,22 @@
from sklearn.metrics.pairwise import euclidean_distances


def _contrast_vector(columns: list[str], model_contrasts: str) -> np.ndarray:
def _contrast_vector(columns: list[str], model_contrasts: str, reference_levels: Collection[str] = ()) -> np.ndarray:
"""Turn an R style contrast such as ``conditionB-conditionA`` into weights over formulaic design columns.

Formulaic names a coefficient ``condition[T.B]`` where R names it ``conditionB``, so the columns are matched on their R spelling.
A term naming a reference level in ``reference_levels`` has no coefficient under treatment coding and weighs zero.
"""
r_names = {column.replace("[T.", "").replace("[", "").replace("]", ""): column for column in columns}
weights = pd.Series(0.0, index=columns)
for sign, term in re.findall(r"([+-]?)\s*([^+-]+)", model_contrasts):
name = term.strip()
if name not in r_names:
if name in r_names:
weights[r_names[name]] += -1.0 if sign == "-" else 1.0
elif name not in reference_levels:
raise ValueError(
f"Contrast term {name!r} does not match any coefficient of the design. Available: {sorted(r_names)}."
)
weights[r_names[name]] += -1.0 if sign == "-" else 1.0
return weights.to_numpy()


Expand Down Expand Up @@ -153,7 +155,8 @@ def make_nhoods(
Otherwise:

nhoods: :class:`scipy.sparse.csr_matrix` in `adata.obsm['nhoods']`.
A binary matrix of cell to neighbourhood assignments. Neighbourhoods in the columns are ordered by the order of the index cell in adata.obs_names
A binary matrix of cell to neighbourhood assignments, in which every index cell belongs to its own neighbourhood.
Neighbourhoods in the columns are ordered by the order of the index cell in adata.obs_names

nhood_ixs_refined: pandas.Series in `adata.obs['nhood_ixs_refined']`.
A boolean indicating whether a cell is an index for a neighbourhood
Expand Down Expand Up @@ -223,6 +226,7 @@ def make_nhoods(
refined_vertices = np.unique(refined_vertices)
refined_vertices.sort()

knn_graph.setdiag(1) # type: ignore[union-attr]
nhoods = knn_graph[:, refined_vertices]
adata.obsm["nhoods"] = nhoods

Expand Down Expand Up @@ -351,9 +355,9 @@ def da_nhoods(
max_iter: Maximum number of iterations of a mixed model fit.
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 is the closest to the R implementation.
The "edger" solver requires R, rpy2 and edgeR to be installed and reproduces the R implementation.
The "pydeseq2" requires pydeseq2 to be installed.
It is still very comparable to the "edger" solver but might be a bit slower.
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.

Returns:
None, modifies `milo_mdata['milo']` in place, adding the results of the DA test to `.var`:
Expand Down Expand Up @@ -439,14 +443,18 @@ def da_nhoods(
if isinstance(design_df[column].dtype, pd.CategoricalDtype):
design_df[column] = design_df[column].cat.remove_unused_categories()

fixed = fixed_design if add_intercept and model_contrasts is None else fixed_design + " + 0"
reference_levels = {
f"{column}{design_df[column].astype('category').cat.categories[0]}"
for column in design_df.select_dtypes(exclude="number").columns
}
if random_effects:
if find_spec("formulaic_contrasts") is None:
raise ImportError(
"formulaic-contrasts is required for mixed models. Install with: pip install pertpy[de]"
)
from formulaic_contrasts import FormulaicContrasts

fixed = fixed_design if add_intercept and model_contrasts is None else fixed_design + " + 0"
design_matrix = FormulaicContrasts(design_df, fixed).design_matrix
_check_residual_df(design_matrix, design)
counts_filtered = count_mat[np.ix_(keep_nhoods, keep_smp)]
Expand All @@ -457,7 +465,7 @@ def da_nhoods(
np.asarray(design_matrix, dtype=float),
random_effect_matrices(design_df, random_effects),
np.log(lib_size_filtered),
contrast=_contrast_vector(list(design_matrix.columns), model_contrasts)
contrast=_contrast_vector(list(design_matrix.columns), model_contrasts, reference_levels)
if model_contrasts is not None
else None,
reml=reml,
Expand Down Expand Up @@ -485,12 +493,10 @@ def da_nhoods(
from rpy2.robjects.vectors import FloatVector

# Define model matrix
if not add_intercept or model_contrasts is not None:
design = design + " + 0"
design_df = design_df.astype(dict.fromkeys(design_df.select_dtypes(exclude=["number"]).columns, "category"))
with localconverter(ro.default_converter + pandas2ri.converter):
design_r = pandas2ri.py2rpy(design_df)
formula_r = stats.formula(design)
formula_r = stats.formula(fixed)
model = stats.model_matrix(object=formula_r, data=design_r)
model_np = np.array(model)
_check_residual_df(model_np, design)
Expand All @@ -504,7 +510,7 @@ def da_nhoods(
dge = edgeR.DGEList(counts=count_mat_r, lib_size=lib_size_r)
dge = edgeR.calcNormFactors(dge, method="TMM")
dge = edgeR.estimateDisp(dge, model)
fit = edgeR.glmQLFit(dge, model, robust=True)
fit = edgeR.glmQLFit(dge, model, robust=True, legacy=True)
# Test
n_coef = model_np.shape[1]
if model_contrasts is not None:
Expand All @@ -518,7 +524,7 @@ def da_nhoods(

get_model_cols = STAP(r_str, "get_model_cols")
with localconverter(ro.default_converter + numpy2ri.converter + pandas2ri.converter):
model_mat_cols = get_model_cols.get_model_cols(design_df, design)
model_mat_cols = get_model_cols.get_model_cols(design_df, fixed)
with localconverter(ro.default_converter + pandas2ri.converter + numpy2ri.converter):
model_df = pandas2ri.rpy2py(model)
model_df = pd.DataFrame(model_df)
Expand Down Expand Up @@ -568,54 +574,24 @@ def da_nhoods(
dict.fromkeys(design_df_filtered.select_dtypes(exclude=["number"]).columns, "category")
)

design_clean = design if design.startswith("~") else f"~{design}"

dds = DeseqDataSet(
counts=pd.DataFrame(counts_filtered.T, index=design_df_filtered.index),
metadata=design_df_filtered,
design=design_clean,
design=fixed if fixed.startswith("~") else f"~{fixed}",
refit_cooks=True,
size_factors_fit_type="poscounts",
)

_check_residual_df(dds.obsm["design_matrix"], design) # type: ignore[arg-type]
design_matrix = cast_frame(dds.obsm["design_matrix"])
_check_residual_df(design_matrix, design)
dds.deseq2()

if model_contrasts is not None and "-" in model_contrasts:
if "(" in model_contrasts or "+" in model_contrasts.split("-")[1]:
raise ValueError(
f"Complex contrasts like '{model_contrasts}' are not supported by pydeseq2. "
"Use simple pairwise contrasts (e.g., 'GroupA-GroupB') or switch to solver='edger'."
)

parts = model_contrasts.split("-")
factor_name = design_clean.replace("~", "").split("+")[-1].strip()
group1 = parts[0].replace(factor_name, "").strip()
group2 = parts[1].replace(factor_name, "").strip()
if factor_name not in design_df_filtered.columns:
raise ValueError(
f"Contrast factor {factor_name!r} is not a column of the design dataframe. "
f"Available columns: {list(design_df_filtered.columns)}."
)
if not isinstance(design_df_filtered[factor_name].dtype, pd.CategoricalDtype):
design_df_filtered[factor_name] = design_df_filtered[factor_name].astype("category")
available_levels = list(design_df_filtered[factor_name].cat.categories)
missing = [g for g in (group1, group2) if g not in available_levels]
if missing:
raise ValueError(
f"Contrast levels {missing!r} not found in factor {factor_name!r}. "
f"Available levels: {available_levels}. "
f"Contrasts must follow the form '{factor_name}<level_a>-{factor_name}<level_b>' "
"with both levels present in the data."
)
stat_res = DeseqStats(dds, contrast=[factor_name, group1, group2])
else:
factor_name = design_clean.replace("~", "").split("+")[-1].strip()
if not isinstance(design_df_filtered[factor_name], pd.CategoricalDtype):
design_df_filtered[factor_name] = design_df_filtered[factor_name].astype("category")
categories = design_df_filtered[factor_name].cat.categories
stat_res = DeseqStats(dds, contrast=[factor_name, categories[-1], categories[0]])

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

Expand Down
64 changes: 64 additions & 0 deletions tests/tools/test_milo.py
Original file line number Diff line number Diff line change
Expand Up @@ -70,6 +70,13 @@ def test_make_nhoods_sizes(adata, milo):
assert knn_graph.sum(0).min() <= adata.obsm["nhoods"].sum(0).min()


def test_make_nhoods_contains_index_cells(adata, milo):
adata = adata.copy()
milo.make_nhoods(adata)
index_cells = np.flatnonzero(adata.obs["nhood_ixs_refined"] == 1)
assert np.all(np.asarray(adata.obsm["nhoods"][index_cells, np.arange(len(index_cells))]) == 1)


def test_make_nhoods_neighbors_key(adata, milo):
adata = adata.copy()
k = adata.uns["neighbors"]["params"]["n_neighbors"]
Expand Down Expand Up @@ -231,6 +238,63 @@ def test_da_nhoods_default_contrast(da_nhoods_mdata, milo, solver):
assert np.corrcoef(contr_results["logFC"], default_results["logFC"])[0, 1] > 0.99


@pytest.fixture
def three_condition_mdata(adata, milo):
adata = adata.copy()
milo.make_nhoods(adata)
rng = np.random.default_rng(seed=42)
conditions = ["ConditionA", "ConditionB", "ConditionC"]
adata.obs["condition"] = rng.choice(conditions, size=adata.n_obs)
da_cells = adata.obs["louvain"] == "1"
adata.obs.loc[da_cells, "condition"] = rng.choice(conditions, size=da_cells.sum(), p=[0.1, 0.8, 0.1])
adata.obs["replicate"] = rng.choice(["R1", "R2", "R3"], size=adata.n_obs)
adata.obs["sample"] = adata.obs["replicate"] + adata.obs["condition"]
return milo.count_nhoods(adata, sample_col="sample")


def test_da_nhoods_contrast_of_single_coefficient(three_condition_mdata, milo, solver):
"""A contrast naming a single coefficient tests that coefficient instead of the last level of the design."""
mdata = three_condition_mdata
index_cells = mdata["milo"].var["index_cell"]
enriched = (mdata["rna"].obs.loc[index_cells, "louvain"] == "1").to_numpy()

milo.da_nhoods(mdata, design="~replicate+condition", model_contrasts="conditionConditionB", solver=solver)
b_vs_a = mdata["milo"].var["logFC"].to_numpy()
milo.da_nhoods(mdata, design="~replicate+condition", model_contrasts="conditionConditionC", solver=solver)
c_vs_a = mdata["milo"].var["logFC"].to_numpy()

assert np.nanmean(b_vs_a[enriched]) > 1
assert np.nanmean(b_vs_a[enriched]) > np.nanmean(c_vs_a[enriched]) + 1


def test_da_nhoods_contrast_against_reference_level(three_condition_mdata, milo):
"""The reference level has no coefficient, so subtracting it leaves the contrast unchanged."""
mdata = three_condition_mdata
milo.da_nhoods(mdata, design="~replicate+condition", model_contrasts="conditionConditionB", solver="pydeseq2")
coefficient = mdata["milo"].var["logFC"].to_numpy()
milo.da_nhoods(
mdata,
design="~replicate+condition",
model_contrasts="conditionConditionB-conditionConditionA",
solver="pydeseq2",
)
np.testing.assert_allclose(mdata["milo"].var["logFC"].to_numpy(), coefficient)


def test_da_nhoods_continuous_covariate_per_unit(da_nhoods_mdata, milo, solver):
"""The log fold change of a continuous covariate is per unit, so rescaling the covariate rescales it."""
mdata = da_nhoods_mdata.copy()
obs = mdata["rna"].obs
obs["dose"] = (obs["condition"] == "ConditionB") + obs["replicate"].str[1].astype(float) / 10

milo.da_nhoods(mdata, design="~dose", solver=solver)
per_unit = mdata["milo"].var["logFC"].to_numpy()
obs["dose"] *= 10
milo.da_nhoods(mdata, design="~dose", solver=solver)

np.testing.assert_allclose(mdata["milo"].var["logFC"].to_numpy() * 10, per_unit, rtol=1e-3, atol=1e-3)


@pytest.mark.skipif(find_spec("formulaic_contrasts") is None, reason="formulaic-contrasts not available")
def test_da_nhoods_glmm(da_nhoods_mdata, milo):
mdata = da_nhoods_mdata.copy()
Expand Down
Loading