Skip to content

Commit d442cc1

Browse files
committed
fix: preserve masked array in view
1 parent e0926ca commit d442cc1

2 files changed

Lines changed: 48 additions & 0 deletions

File tree

src/anndata/_core/views.py

Lines changed: 24 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -195,6 +195,25 @@ def toarray(self) -> np.ndarray:
195195
return self.copy()
196196

197197

198+
class MaskedArrayView(_SetItemMixin, np.ma.MaskedArray):
199+
def __new__(
200+
cls,
201+
input_array: Sequence[Any],
202+
view_args: ViewArgs | None = None,
203+
):
204+
arr = np.ma.asarray(input_array).view(cls)
205+
206+
if view_args is not None:
207+
view_args = ElementRef(*view_args)
208+
arr._view_args = view_args
209+
return arr
210+
211+
def __array_finalize__(self, obj: np.ndarray | None):
212+
super().__array_finalize__(obj)
213+
if obj is not None:
214+
self._view_args = getattr(obj, "_view_args", None)
215+
216+
198217
# Extends DaskArray
199218
# Calls parent __new__ constructor since
200219
# even calling astype on a dask array
@@ -341,6 +360,11 @@ def as_view_array(array, view_args):
341360
return ArrayView(array, view_args=view_args)
342361

343362

363+
@as_view.register(np.ma.MaskedArray)
364+
def as_view_masked_array(array, view_args):
365+
return MaskedArrayView(array, view_args=view_args)
366+
367+
344368
@as_view.register(DaskArray)
345369
def as_view_dask_array(array, view_args):
346370
return DaskArrayView(array, view_args=view_args)

tests/test_views.py

Lines changed: 24 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -29,6 +29,7 @@
2929
)
3030
from anndata._core.views import (
3131
ArrayView,
32+
MaskedArrayView,
3233
SparseCSCArrayView,
3334
SparseCSCMatrixView,
3435
SparseCSRArrayView,
@@ -709,6 +710,29 @@ def test_view_retains_ndarray_subclass():
709710
assert view.obsm["foo"].shape == (5, 5)
710711

711712

713+
def test_view_of_masked_array():
714+
mask = np.zeros((10, 5), dtype=bool)
715+
mask[0, 0] = True
716+
data = np.ma.MaskedArray(np.arange(50.0).reshape(10, 5), mask=mask)
717+
718+
adata = ad.AnnData(np.zeros((10, 10)), obsm={"masked": data})
719+
view = adata[:5, :]
720+
721+
masked_view = view.obsm["masked"]
722+
assert isinstance(masked_view, MaskedArrayView)
723+
assert isinstance(masked_view, np.ma.MaskedArray)
724+
assert masked_view.mask[0, 0]
725+
assert not masked_view.mask[1, 0]
726+
assert np.ma.is_masked(masked_view)
727+
728+
with pytest.warns(ImplicitModificationWarning, match=r"obsm"):
729+
masked_view[1, 0] = 1234.0
730+
assert view.obsm["masked"][1, 0] == 1234.0
731+
# Original is untouched: writing to a view actualizes it as a copy.
732+
assert adata.obsm["masked"][1, 0] == 5.0
733+
assert adata.obsm["masked"].mask[0, 0]
734+
735+
712736
def test_modify_uns_in_copy():
713737
# https://github.com/scverse/anndata/issues/571
714738
adata = ad.AnnData(np.ones((5, 5)), uns={"parent": {"key": "value"}})

0 commit comments

Comments
 (0)