Skip to content

Commit d8c3f10

Browse files
committed
covariates: accept string columns and treat them as categoricals
1 parent ea5455d commit d8c3f10

4 files changed

Lines changed: 26 additions & 10 deletions

File tree

src/mofaflex/_core/datasets/misc.py

Lines changed: 13 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,5 @@
11
from collections.abc import Mapping, Sequence
2+
from functools import reduce
23

34
import numpy as np
45
import pandas as pd
@@ -54,19 +55,26 @@ def merge_covariates(covariates: Mapping[str, Mapping[str, pd.DataFrame]]):
5455
for group_covars in covariates.values():
5556
for view_covars in group_covars.values():
5657
dtypes = view_covars.dtypes
58+
cats = None
5759
if dtypes.nunique() > 1:
5860
raise ValueError("Mixed dtypes for a covariate are not supported.")
5961
if dtypes.iloc[0] == "category":
60-
categories = (
61-
view_covars.iloc[0].cat.categories
62-
if categories is None
63-
else categories.union(view_covars.iloc[0].cat.categories)
62+
cats = view_covars.iloc[0].cat.categories
63+
elif pd.api.types.is_string_dtype(dtypes.iloc[0]):
64+
cats = reduce(
65+
lambda x, y: x.union(y), (pd.Index(col.dropna().unique()) for _, col in view_covars.items())
6466
)
67+
if cats is not None:
68+
categories = pd.Index(cats) if categories is None else categories.union(cats)
6569
for group_covars in covariates.values():
6670
for view_covars in group_covars.values():
67-
if view_covars.dtypes.iloc[0] == "category":
71+
dtypes = view_covars.dtypes
72+
if dtypes.iloc[0] == "category":
6873
for col in view_covars.columns:
6974
view_covars[col] = view_covars[col].cat.set_categories(categories)
75+
elif pd.api.types.is_string_dtype(dtypes.iloc[0]):
76+
for col in view_covars.columns:
77+
view_covars[col] = pd.Categorical(view_covars[col], categories=categories)
7078

7179
# ensure the covariate value is consistent across views (nanmean or first)
7280
merged_covariates = {}

tests/conftest.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -68,6 +68,7 @@ def _adata(likelihood, nobs, nvar, var_names=None, obs_names=None):
6868
adata.obs["gvar_normal"] = rng.random(size=(nobs))
6969
adata.obs["gvar_bernoulli"] = rng.binomial(1, 0.5, size=(nobs))
7070
adata.obs["gvar_categorical"] = pd.Categorical(rng.choice(["A", "B", "C"], size=(nobs)))
71+
adata.obs["gvar_string"] = rng.choice(["A", "B", "C"], size=(nobs))
7172
adata.varm["annot_df"] = pd.DataFrame(
7273
rng.choice([False, True], size=(nvar, 10)), columns=[f"annot_{i}" for i in range(10)], index=adata.var_names
7374
)

tests/test_datasets_common.py

Lines changed: 5 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -4,14 +4,16 @@
44
from mofaflex._core.datasets import MofaFlexDataset, merge_covariates
55

66

7-
@pytest.mark.parametrize("axis", (0, 1))
8-
def test_merge_covariates_preserves_axis_order(axis, random_adata):
7+
@pytest.mark.parametrize(
8+
("axis", "key", "mkey"), ((0, None, "covar_array"), (1, None, "covar_array"), (0, "gvar_string", None))
9+
)
10+
def test_merge_covariates_preserves_axis_order(axis, key, mkey, random_adata):
911
"""merge_covariates must keep rows in the dataset's sample/feature order, not lexically sorted."""
1012

1113
adata = random_adata("Normal", 500, 500)
1214
dataset = MofaFlexDataset(adata)
1315

14-
covars = dataset.get_covariates(axis, mkey="covar_array")
16+
covars = dataset.get_covariates(axis, key=key, mkey=mkey)
1517
merged = next(iter(merge_covariates(covars).values()))
1618
canonical = next(iter(dataset.get_names(axis).values()))
1719

tests/test_mofaflex_integration.py

Lines changed: 7 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -100,7 +100,11 @@ def model_api_untrained_only():
100100
"argfor,argname,argval",
101101
[
102102
("likelihood_normal", "scale_per_group", False),
103-
("term_mofaflex", "guiding_vars_obs_keys", ["gvar_normal", "gvar_bernoulli", "gvar_categorical"]),
103+
(
104+
"term_mofaflex",
105+
"guiding_vars_obs_keys",
106+
["gvar_normal", "gvar_bernoulli", "gvar_categorical", "gvar_string"],
107+
),
104108
("term_mofaflex", "weight_prior", "Normal"),
105109
("term_mofaflex", "weight_prior", "Laplace"),
106110
("term_mofaflex", "weight_prior", "Horseshoe"),
@@ -174,6 +178,7 @@ def test_integration(
174178
"gvar_normal": "Normal",
175179
"gvar_bernoulli": "Bernoulli",
176180
"gvar_categorical": "Categorical",
181+
"gvar_string": "Categorical",
177182
},
178183
**termargs,
179184
)
@@ -211,7 +216,7 @@ def test_integration(
211216
assert "get_weight_annotations" in dir(model)
212217
assert all(isinstance(annot, pd.DataFrame) for annot in model.get_weight_annotations().values())
213218
elif argname == "guiding_vars_obs_keys":
214-
assert model.n_guided_factors == model.terms["_"].n_guided_factors == 3
219+
assert model.n_guided_factors == model.terms["_"].n_guided_factors == 4
215220
else:
216221
assert (
217222
model.n_factors

0 commit comments

Comments
 (0)