@@ -127,6 +127,8 @@ def add_mdata_global_columns(md: MuData, rng: np.random.Generator) -> MuData:
127127@pytest .fixture
128128def 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
192198def 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
210223def 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
344371def 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