Skip to content
Open
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
19 changes: 0 additions & 19 deletions src/anndata/_core/aligned_mapping.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,7 +22,6 @@
axis_len,
convert_to_dict,
deprecation_msg,
raise_value_error_if_multiindex_columns,
warn,
warn_once,
)
Expand Down Expand Up @@ -351,24 +350,6 @@ def to_df(self) -> pd.DataFrame:
df[f"{key}{icolumn + 1}"] = column
return df

def _validate_value(self, val: _AlignedAny, key: str) -> V:
if isinstance(val, pd.DataFrame):
raise_value_error_if_multiindex_columns(val, f"{self.attrname}[{key!r}]")
if not val.index.equals(self.dim_names):
# Could probably also re-order index if it’s contained
try:
pd.testing.assert_index_equal(val.index, self.dim_names)
except AssertionError as e:
msg = f"value.index does not match parent’s {self.dim} names:\n{e}"
raise ValueError(msg) from None
else:
msg = "Index.equals and pd.testing.assert_index_equal disagree"
raise AssertionError(msg)
val.index.name = (
self.dim_names.name
) # this is consistent with AnnData.obsm.setter and AnnData.varm.setter
return super()._validate_value(val, key)

@property
def dim_names(self) -> pd.Index:
return (self.parent.obs_names, self.parent.var_names)[self._axis]
Expand Down
40 changes: 19 additions & 21 deletions tests/test_obsmvarm.py
Original file line number Diff line number Diff line change
@@ -1,18 +1,24 @@
from __future__ import annotations

from functools import partial
from typing import TYPE_CHECKING

import joblib
import numpy as np
import pandas as pd
import pytest
from scipy import sparse

import anndata as ad
from anndata import AnnData
from anndata.compat import CupyArray
from anndata.tests.helpers import as_cupy, get_multiindex_columns_df, jnp
from anndata.tests.helpers import as_cupy, assert_equal, get_multiindex_columns_df, jnp
from anndata.utils import asarray

if TYPE_CHECKING:
from pathlib import Path
from typing import Literal

M, N = (100, 100)


Expand Down Expand Up @@ -81,26 +87,6 @@ def test_setting_ndarray(adata: AnnData):
assert h == joblib.hash(adata)


def test_setting_dataframe(adata: AnnData):
obsm_df = pd.DataFrame(dict(b_1=np.ones(M), b_2=["a"] * M), index=adata.obs_names)
varm_df = pd.DataFrame(dict(b_1=np.ones(N), b_2=["a"] * N), index=adata.var_names)

adata.obsm["b"] = obsm_df
assert np.all(adata.obsm["b"] == obsm_df)
adata.varm["b"] = varm_df
assert np.all(adata.varm["b"] == varm_df)

bad_obsm_df = obsm_df.copy()
bad_obsm_df.reset_index(inplace=True)
with pytest.raises(ValueError, match=r"index does not match.*obs names"):
adata.obsm["c"] = bad_obsm_df

bad_varm_df = varm_df.copy()
bad_varm_df.reset_index(inplace=True)
with pytest.raises(ValueError, match=r"index does not match.*var names"):
adata.varm["c"] = bad_varm_df


def test_setting_sparse(adata: AnnData):
obsm_sparse = sparse.random(M, 100, format="csr")
assert isinstance(obsm_sparse, sparse.csr_matrix)
Expand Down Expand Up @@ -182,3 +168,15 @@ def test_1d_declaration(array_type):
def test_1d_set(adata, array_type):
adata.varm["1d-array"] = array_type(np.ones(adata.shape[1]))
assert adata.varm["1d-array"].shape == (adata.shape[1], 1)


@pytest.mark.parametrize("axis", ["obs", "var"])
def test_roundtrips_df_with_different_index(
tmp_path: Path, adata: AnnData, axis: Literal["obs", "var"]
):
getattr(adata, f"{axis}m")["df"] = pd.DataFrame(
index=[f"not_{i}" for i in range(getattr(adata, axis).shape[0])]
)
adata.write_h5ad(tmp_path / "foo.h5ad")
roundtripped = ad.read_h5ad(tmp_path / "foo.h5ad")
assert_equal(adata, roundtripped)
Loading