Skip to content

Commit 3735151

Browse files
committed
update: preserve global index name
1 parent f4a6b50 commit 3735151

3 files changed

Lines changed: 54 additions & 4 deletions

File tree

CHANGELOG.md

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -8,6 +8,12 @@ and this project adheres to [Semantic Versioning][].
88
[keep a changelog]: https://keepachangelog.com/en/1.1.0/
99
[semantic versioning]: https://semver.org/spec/v2.0.0.html
1010

11+
## [0.4.1]
12+
13+
### Fixed
14+
15+
- `update()` now preserves the index name of the global `obs_names` and `var_names`.
16+
1117
## [0.4.0]
1218

1319
### Added
@@ -202,6 +208,7 @@ To copy the annotations explicitly, you will need to use `pull_obs()` and/or `pu
202208

203209
Initial `mudata` release with `MuData`, previously a part of the `muon` framework.
204210

211+
[0.4.1]: https://github.com/scverse/mudata/releases/tag/v0.4.1
205212
[0.4.0]: https://github.com/scverse/mudata/releases/tag/v0.4.0
206213
[0.3.10]: https://github.com/scverse/mudata/releases/tag/v0.3.10
207214
[0.3.9]: https://github.com/scverse/mudata/releases/tag/v0.3.9

src/mudata/_core/mudata.py

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -541,6 +541,7 @@ def _update_attr(
541541

542542
data_global = getattr(self, attr)
543543
prev_index = data_global.index
544+
prev_index_name = prev_index.name
544545

545546
attr_duplicated = not data_global.index.is_unique or self._check_duplicated_attr_names(attr)
546547
attr_intersecting = self._check_intersecting_attr_names(attr)
@@ -604,6 +605,7 @@ def calc_attrm_update():
604605
if not attr_duplicated:
605606
# Shared axis
606607
data_mod = pd.concat(dfs, join="outer", axis=1 if axis == self._axis or self._axis == -1 else 0, sort=False)
608+
data_mod.index.name = prev_index_name
607609
for mod in self._mod.keys():
608610
fix_attrmap_col(data_mod, mod, rowcol)
609611

@@ -671,8 +673,7 @@ def calc_attrm_update():
671673

672674
data_mod.reset_index(level=list(range(1, data_mod.index.nlevels)), inplace=True)
673675
data_global.reset_index(level=list(range(1, data_global.index.nlevels)), inplace=True)
674-
data_mod.index.set_names(None, inplace=True)
675-
data_global.index.set_names(None, inplace=True)
676+
data_mod.index.set_names(prev_index_name, inplace=True)
676677

677678
# get adata positions and remove columns from the data frame
678679
mdict = {}

tests/test_update.py

Lines changed: 44 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -127,6 +127,8 @@ def add_mdata_global_columns(md: MuData, rng: np.random.Generator) -> MuData:
127127
@pytest.fixture
128128
def mdata(rng: np.random.Generator, modalities: Mapping[str, AnnData], axis: Axis):
129129
md = MuData(modalities, axis=axis)
130+
md.obs.index.name = "obs_idx"
131+
md.var.index.name = "var_idx"
130132

131133
return add_mdata_global_columns(md, rng)
132134

@@ -188,12 +190,21 @@ def test_update_simple(mdata: MuData, axis: Axis):
188190
getattr(mdata, f"{attr}_names")[: mdata["mod1"].shape[axis]] == getattr(mdata["mod1"], f"{attr}_names")
189191
).all()
190192

193+
mdata.update()
194+
assert getattr(mdata, attr).index.name == f"{attr}_idx"
195+
assert getattr(mdata, oattr).index.name == f"{oattr}_idx"
196+
191197

192198
def test_update_simple_empty_modalities(modalities: Mapping[str, AnnData], axis: Axis):
199+
attr = "obs" if axis == 0 else "var"
200+
oattr = "var" if axis == 0 else "obs"
201+
193202
for mod in modalities.values():
194203
mod.obs = pd.DataFrame(index=mod.obs_names)
195204
mod.var = pd.DataFrame(index=mod.var_names)
196205
mdata = MuData(modalities)
206+
getattr(mdata, attr).index.name = f"{attr}_idx"
207+
getattr(mdata, oattr).index.name = f"{oattr}_idx"
197208

198209
old_obsnames = mdata.obs_names
199210
old_varnames = mdata.var_names
@@ -205,17 +216,21 @@ def test_update_simple_empty_modalities(modalities: Mapping[str, AnnData], axis:
205216

206217
assert (old_obsnames == mdata.obs_names).all()
207218
assert (old_varnames == mdata.var_names).all()
219+
assert getattr(mdata, attr).index.name == f"{attr}_idx"
220+
assert getattr(mdata, oattr).index.name == f"{oattr}_idx"
208221

209222

210223
def test_update_add_modality(rng: np.random.Generator, modalities: Mapping[str, AnnData], axis: Axis):
211224
modnames = list(modalities.keys())
212225
mdata = add_mdata_global_columns(
213226
MuData({modname: modalities[modname] for modname in modnames[:-2]}, axis=axis), rng
214227
)
215-
216228
attr = "obs" if axis == 0 else "var"
217229
oattr = "var" if axis == 0 else "obs"
218230

231+
getattr(mdata, attr).index.name = f"{attr}_idx"
232+
getattr(mdata, oattr).index.name = f"{oattr}_idx"
233+
219234
for i in (-2, -1):
220235
old_attrnames = getattr(mdata, f"{attr}_names")
221236
old_oattrnames = getattr(mdata, f"{oattr}_names")
@@ -226,6 +241,8 @@ def test_update_add_modality(rng: np.random.Generator, modalities: Mapping[str,
226241

227242
mdata.mod[modnames[i]] = modalities[modnames[i]]
228243
mdata.update()
244+
assert getattr(mdata, attr).index.name == f"{attr}_idx"
245+
assert getattr(mdata, oattr).index.name == f"{oattr}_idx"
229246

230247
for mod in mdata.mod.keys():
231248
assert mdata.obsmap[mod].dtype.kind == "u"
@@ -272,6 +289,8 @@ def test_update_delete_modality(mdata: MuData, axis: Axis):
272289

273290
del mdata.mod[modnames[0]]
274291
mdata.update()
292+
assert getattr(mdata, attr).index.name == f"{attr}_idx"
293+
assert getattr(mdata, oattr).index.name == f"{oattr}_idx"
275294

276295
for mod in mdata.mod.keys():
277296
assert mdata.obsmap[mod].dtype.kind == "u"
@@ -295,6 +314,8 @@ def test_update_delete_modality(mdata: MuData, axis: Axis):
295314

296315
del mdata.mod[modnames[2]]
297316
mdata.update()
317+
assert getattr(mdata, attr).index.name == f"{attr}_idx"
318+
assert getattr(mdata, oattr).index.name == f"{oattr}_idx"
298319

299320
assert mdata.shape[1 - axis] == sum(mod.shape[1 - axis] for mod in mdata.mod.values())
300321
assert (getattr(mdata, oattr)["batch"] == fullobatch[keptomask]).all()
@@ -340,14 +361,23 @@ def test_update_intersecting(rng: np.random.Generator, modalities: Mapping[str,
340361
assert mdata.shape[axis] == axisnames.shape[0]
341362
assert (getattr(mdata, f"{attr}_names") == axisnames).all()
342363

364+
getattr(mdata, attr).index.name = f"{attr}_idx"
365+
getattr(mdata, oattr).index.name = f"{oattr}_idx"
366+
mdata.update()
367+
assert getattr(mdata, attr).index.name == f"{attr}_idx"
368+
assert getattr(mdata, oattr).index.name == f"{oattr}_idx"
369+
343370

344371
def test_update_intersecting_after_filtering(mdata):
372+
attr = "obs" if mdata.axis == 0 else "var"
345373
oattr = "var" if mdata.axis == 0 else "obs"
346374
orig_shape = mdata.shape
347375

348376
for mod in mdata.mod.values():
349377
setattr(mod, f"{oattr}_names", [f"{oattr}{j}" for j in range(mod.shape[1 - mdata.axis])])
350378
mdata.update()
379+
assert getattr(mdata, attr).index.name == f"{attr}_idx"
380+
assert getattr(mdata, oattr).index.name == f"{oattr}_idx"
351381

352382
subset = (slice(None), slice(5))
353383
subset = subset[mdata.axis], subset[1 - mdata.axis]
@@ -359,6 +389,8 @@ def test_update_intersecting_after_filtering(mdata):
359389

360390
assert mdata["mod1"].shape[1 - mdata.axis] == 5
361391
mdata.update()
392+
assert getattr(mdata, attr).index.name == f"{attr}_idx"
393+
assert getattr(mdata, oattr).index.name == f"{oattr}_idx"
362394
getattr(mdata, f"pull_{oattr}")(prefix_unique=False, join_nonunique=True)
363395
assert mdata.shape[mdata.axis] == orig_shape[mdata.axis]
364396
assert mdata.shape[1 - mdata.axis] == sum(mod.shape[1 - mdata.axis] for mod in mdata.mod.values())
@@ -368,13 +400,16 @@ def test_update_intersecting_after_filtering(mdata):
368400
).sum()
369401

370402

371-
def test_update_after_filter_obs_adata(mdata: MuData, axis: Axis):
403+
def test_update_after_filter_obs_adata(mdata: MuData):
372404
"""
373405
Check for https://github.com/scverse/muon/issues/44
374406
"""
375407
# Replicate in-place filtering in muon:
376408
# mu.pp.filter_obs(mdata['mod1'], 'min_count', lambda x: (x < -2))
377409

410+
attr = "obs" if mdata.axis == 0 else "var"
411+
oattr = "var" if mdata.axis == 0 else "obs"
412+
378413
old_obsnames = mdata.obs_names
379414
old_varnames = mdata.var_names
380415

@@ -388,6 +423,8 @@ def test_update_after_filter_obs_adata(mdata: MuData, axis: Axis):
388423

389424
mdata.mod["mod3"] = mdata["mod3"][mdata["mod3"].obs["min_count"] < -2].copy()
390425
mdata.update()
426+
assert getattr(mdata, attr).index.name == f"{attr}_idx"
427+
assert getattr(mdata, oattr).index.name == f"{oattr}_idx"
391428

392429
for mod in mdata.mod.keys():
393430
assert mdata.obsmap[mod].dtype.kind == "u"
@@ -411,12 +448,17 @@ def test_update_after_obs_reordered(mdata: MuData):
411448
"""
412449
Update should work if obs are reordered.
413450
"""
451+
attr = "obs" if mdata.axis == 0 else "var"
452+
oattr = "var" if mdata.axis == 0 else "obs"
453+
414454
some_obs_names = mdata.obs_names.values[:2]
415455

416456
true_obsm_values = get_attrm_values(mdata, "obs", "test", some_obs_names)
417457

418458
mdata.mod["mod1"] = mdata["mod1"][::-1].copy()
419459
mdata.update()
460+
assert getattr(mdata, attr).index.name == f"{attr}_idx"
461+
assert getattr(mdata, oattr).index.name == f"{oattr}_idx"
420462

421463
for mod in mdata.mod.keys():
422464
assert mdata.obsmap[mod].dtype.kind == "u"

0 commit comments

Comments
 (0)