diff --git a/pyproject.toml b/pyproject.toml index 03181eadb..eebd22b06 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -244,3 +244,38 @@ memray-flame = "memray flamegraph --temporal" [tool.pixi.environments] profiling = { features = ["profiling"], solve-group = "default" } + +[tool.hatch.envs.test] +dependency-groups = ["test"] + +[tool.hatch.envs.test-anndata-pandas] +template = "test" +extra-dependencies = ["zarr>=3"] +scripts.test = ["pip list|grep anndata && pip list|grep pandas && pytest {args}"] +scripts.test-readwrite = ["pip list|grep anndata && pip list|grep pandas && pytest tests/io/test_readwrite.py"] +scripts.test-all = ["pip list|grep anndata && pip list|grep pandas && pytest ."] + +[[tool.hatch.envs.test-anndata-pandas.matrix]] +anndata-pandas = [ + "0.13-2", + "0.13-3", + "0.12-2" + # support for pandas>=3 is available only in anndata>=0.13 +] + +[tool.hatch.envs.test-anndata-pandas.overrides] +matrix.anndata-pandas.extra-dependencies = [ + # every option where the if-condition is True gets included + + # anndata 0.13 + {value="anndata~=0.13", if = ["0.13-2", "0.13-3"]}, + + # anndata 0.12 + {value="anndata>=0.12,<0.13", if = ["0.12-2"]}, + + # pandas 2 + {value="pandas>=2.3,<3", if = ["0.13-2", "0.12-2"]}, + + # pandas 3 + {value="pandas~=3.0", if = ["0.13-3"]}, +] diff --git a/src/spatialdata/_core/spatialdata.py b/src/spatialdata/_core/spatialdata.py index fb55ab086..1baa1c33d 100644 --- a/src/spatialdata/_core/spatialdata.py +++ b/src/spatialdata/_core/spatialdata.py @@ -1114,6 +1114,7 @@ def write( sdata_formats: SpatialDataFormatType | list[SpatialDataFormatType] | None = None, shapes_geometry_encoding: Literal["WKB", "geoarrow"] | None = None, raster_compressor: dict[Literal["lz4", "zstd"], int] | None = None, + convert_table_strings_to_categoricals: bool = False, ) -> None: """ Write the `SpatialData` object to a Zarr store. @@ -1166,6 +1167,9 @@ def write( compression level which should be inclusive between 0 and 9. For compression, `lz4` and `zstd` are supported. If not specified, the compression will be `lz4` with compression level 5. Bytes are automatically ordered for more efficient compression. + convert_table_strings_to_categoricals + If True, convert string columns of all tables to categoricals before writing. + Note that this will have a side effect of modifying string columns into categoricals in place. """ from spatialdata._io._utils import _resolve_zarr_store, _validate_compressor_args from spatialdata._io.format import _parse_formats @@ -1194,6 +1198,7 @@ def write( parsed_formats=parsed, shapes_geometry_encoding=shapes_geometry_encoding, raster_compressor=raster_compressor, + convert_table_strings_to_categoricals=convert_table_strings_to_categoricals, ) if self.path != file_path and update_sdata_path: @@ -1212,6 +1217,7 @@ def _write_element( parsed_formats: dict[str, SpatialDataFormatType] | None = None, shapes_geometry_encoding: Literal["WKB", "geoarrow"] | None = None, raster_compressor: dict[Literal["lz4", "zstd"], int] | None = None, + convert_table_strings_to_categoricals: bool = False, ) -> None: from spatialdata._io.io_zarr import _get_groups_for_element @@ -1279,6 +1285,7 @@ def _write_element( group=element_type_group, name=element_name, element_format=parsed_formats["tables"], + convert_strings_to_categoricals=convert_table_strings_to_categoricals, ) else: raise ValueError(f"Unknown element type: {element_type}") @@ -1290,6 +1297,7 @@ def write_element( sdata_formats: SpatialDataFormatType | list[SpatialDataFormatType] | None = None, shapes_geometry_encoding: Literal["WKB", "geoarrow"] | None = None, raster_compressor: dict[Literal["lz4", "zstd"], int] | None = None, + convert_table_strings_to_categoricals: bool = False, ) -> None: """ Write a single element, or a list of elements, to the Zarr store used for backing. @@ -1308,11 +1316,14 @@ def write_element( shapes_geometry_encoding Whether to use the WKB or geoarrow encoding for GeoParquet. See :meth:`geopandas.GeoDataFrame.to_parquet` for details. If None, uses the value from :attr:`spatialdata.settings.shapes_geometry_encoding`. - raster_compressor + raster_compressor A lenght-1 dictionary with as key the type of compression to use for images and labels and as value the compression level which should be inclusive between 0 and 9. For compression, `lz4` and `zstd` are supported. If not specified, the compression will be `lz4` with compression level 5. Bytes are automatically ordered for more efficient compression. + convert_table_strings_to_categoricals + If True, and if element to be written is a table, convert string columns to categoricals before writing. + Note that this will have a side effect of modifying string columns into categoricals in place. Notes ----- @@ -1332,6 +1343,7 @@ def write_element( sdata_formats=sdata_formats, shapes_geometry_encoding=shapes_geometry_encoding, raster_compressor=raster_compressor, + convert_table_strings_to_categoricals=convert_table_strings_to_categoricals, ) return @@ -1368,6 +1380,7 @@ def write_element( parsed_formats=parsed_formats, shapes_geometry_encoding=shapes_geometry_encoding, raster_compressor=raster_compressor, + convert_table_strings_to_categoricals=convert_table_strings_to_categoricals, ) # After every write, metadata should be consolidated, otherwise this can lead to IO problems like when deleting. if self.has_consolidated_metadata(): diff --git a/src/spatialdata/_io/exceptions.py b/src/spatialdata/_io/exceptions.py new file mode 100644 index 000000000..66f5802b7 --- /dev/null +++ b/src/spatialdata/_io/exceptions.py @@ -0,0 +1,26 @@ +from __future__ import annotations + +from ome_zarr.format import Format + + +class FormatVersionUnknownError(ValueError): + """Exception raised when an unknown element format is encountered.""" + + def __init__(self, element_type: str, version_encountered: Format): + self.element_type = element_type + self.version_encountered = version_encountered + self.message = ( + f"Encountered unknown element format version " + f"`{self.version_encountered}` for element of type `{self.element_type}`" + ) + super().__init__(self.message) + + +class WritingToZarrV2DeprecationWarning(DeprecationWarning): + """Warning raised when writing to zarr v2 format.""" + + message = ( + "Writing to zarr v2 format is currently deprecated in spatialdata " + "and will be removed in a future version. " + "Please consider writing to zarr v3." + ) diff --git a/src/spatialdata/_io/io_points.py b/src/spatialdata/_io/io_points.py index 03ef33389..bb203cad2 100644 --- a/src/spatialdata/_io/io_points.py +++ b/src/spatialdata/_io/io_points.py @@ -1,5 +1,6 @@ from __future__ import annotations +import warnings from pathlib import Path import zarr @@ -12,6 +13,7 @@ _write_metadata, overwrite_coordinate_transformations_non_raster, ) +from spatialdata._io.exceptions import WritingToZarrV2DeprecationWarning from spatialdata._io.format import CurrentPointsFormat, PointsFormats, _parse_version from spatialdata.models import get_axes_names from spatialdata.transformations._utils import ( @@ -65,6 +67,10 @@ def write_points( element_format The format of the points element used to store it. """ + if element_format.zarr_format == 2: + warnings.warn( + message=WritingToZarrV2DeprecationWarning.message, category=WritingToZarrV2DeprecationWarning, stacklevel=2 + ) axes = get_axes_names(points) transformations = _get_transformations(points) assert transformations is not None # mypy: validate_element() in _write_element guarantees this diff --git a/src/spatialdata/_io/io_raster.py b/src/spatialdata/_io/io_raster.py index 276f016bd..b9a2964f0 100644 --- a/src/spatialdata/_io/io_raster.py +++ b/src/spatialdata/_io/io_raster.py @@ -1,5 +1,6 @@ from __future__ import annotations +import warnings from collections.abc import Sequence from pathlib import Path from typing import Any, Literal, TypeGuard, cast @@ -23,6 +24,7 @@ overwrite_channel_names, overwrite_coordinate_transformations_raster, ) +from spatialdata._io.exceptions import WritingToZarrV2DeprecationWarning from spatialdata._io.format import ( CurrentRasterFormat, RasterFormatType, @@ -581,6 +583,11 @@ def write_image( raster_compressor: dict[Literal["lz4", "zstd"], int] | None = None, **metadata: str | JSONDict | list[JSONDict], ) -> None: + if element_format.zarr_format == 2: + warnings.warn( + message=WritingToZarrV2DeprecationWarning.message, category=WritingToZarrV2DeprecationWarning, stacklevel=2 + ) + _write_raster( raster_type="image", raster_data=image, @@ -603,6 +610,11 @@ def write_labels( raster_compressor: dict[Literal["lz4", "zstd"], int] | None = None, **metadata: JSONDict, ) -> None: + if element_format.zarr_format == 2: + warnings.warn( + message=WritingToZarrV2DeprecationWarning.message, category=WritingToZarrV2DeprecationWarning, stacklevel=2 + ) + _write_raster( raster_type="labels", raster_data=labels, diff --git a/src/spatialdata/_io/io_shapes.py b/src/spatialdata/_io/io_shapes.py index 3b6e18e39..f8528868d 100644 --- a/src/spatialdata/_io/io_shapes.py +++ b/src/spatialdata/_io/io_shapes.py @@ -1,5 +1,6 @@ from __future__ import annotations +import warnings from pathlib import Path from typing import Any, Literal @@ -15,6 +16,7 @@ _write_metadata, overwrite_coordinate_transformations_non_raster, ) +from spatialdata._io.exceptions import WritingToZarrV2DeprecationWarning from spatialdata._io.format import ( CurrentShapesFormat, ShapesFormats, @@ -93,6 +95,11 @@ def write_shapes( Whether to use the WKB or geoarrow encoding for GeoParquet. See :meth:`geopandas.GeoDataFrame.to_parquet` for details. If None, uses the value from :attr:`spatialdata.settings.shapes_geometry_encoding`. """ + if element_format.zarr_format == 2: + warnings.warn( + message=WritingToZarrV2DeprecationWarning.message, category=WritingToZarrV2DeprecationWarning, stacklevel=2 + ) + from spatialdata.config import settings if geometry_encoding is None: diff --git a/src/spatialdata/_io/io_table.py b/src/spatialdata/_io/io_table.py index 3eb4b0927..931572389 100644 --- a/src/spatialdata/_io/io_table.py +++ b/src/spatialdata/_io/io_table.py @@ -1,5 +1,7 @@ from __future__ import annotations +import warnings +from importlib.metadata import version from pathlib import Path import numpy as np @@ -8,7 +10,10 @@ from anndata import read_zarr as read_anndata_zarr from anndata._io.specs import write_elem as write_adata from ome_zarr.format import Format +from packaging.version import Version +from spatialdata._io._utils import _resolve_zarr_store +from spatialdata._io.exceptions import FormatVersionUnknownError, WritingToZarrV2DeprecationWarning from spatialdata._io.format import ( CurrentTablesFormat, TablesFormats, @@ -55,17 +60,67 @@ def write_table( name: str, group_type: str = "ngff:regions_table", element_format: Format = CurrentTablesFormat(), + convert_strings_to_categoricals: bool = False, ) -> None: + """ + Write a table to a Zarr store. + + Parameters + ---------- + table + The table to write. + group + The table will be written into a subgroup of this group + name + The name of the subgroup of `group` to which table is to be written. + group_type + The type of the group. + element_format + The format to use for writing the table. + convert_strings_to_categoricals + If True, convert string columns to categoricals before writing. + Note that this will have a side effect of modifying dtypes of the input table in place. + """ + if element_format.zarr_format == 2: + warnings.warn( + message=WritingToZarrV2DeprecationWarning.message, category=WritingToZarrV2DeprecationWarning, stacklevel=2 + ) + if TableModel.ATTRS_KEY in table.uns: region, region_key, instance_key = get_table_keys(table) TableModel.validate(table) else: region, region_key, instance_key = (None, None, None) - write_adata(group, name, table) - tables_group = group[name] - tables_group.attrs["spatialdata-encoding-type"] = group_type - tables_group.attrs["region"] = region - tables_group.attrs["region_key"] = region_key - tables_group.attrs["instance_key"] = instance_key - tables_group.attrs["version"] = element_format.spatialdata_format_version + # Ensure the table group exists + table_group = group.require_group(name=name) + + if element_format not in TablesFormats.values(): + raise FormatVersionUnknownError(element_type="table", version_encountered=element_format) + + if element_format.zarr_format == 3 and Version(version("anndata")) >= Version("0.13"): + # `write_zarr` in anndata v0.13 and above can only write to zarr v3 + # solution of passing resolved store directly roughly based on: + # https://github.com/scverse/anndata/issues/1548#issuecomment-2199801855 + + # resolve the store from the group + resolved_store = _resolve_zarr_store(table_group) + + # Write the table to the path of the table group + table.write_zarr( + store=resolved_store, + consolidate_metadata=False, + convert_strings_to_categoricals=convert_strings_to_categoricals, + ) + + else: + if convert_strings_to_categoricals: + table.strings_to_categoricals() + + write_adata(group, name, table) + + table_group.attrs["spatialdata-encoding-type"] = group_type + table_group.attrs["region"] = region + table_group.attrs["region_key"] = region_key + table_group.attrs["instance_key"] = instance_key + table_group.attrs["version"] = element_format.spatialdata_format_version diff --git a/tests/io/test_readwrite.py b/tests/io/test_readwrite.py index 034c01d37..609bdc12f 100644 --- a/tests/io/test_readwrite.py +++ b/tests/io/test_readwrite.py @@ -4,6 +4,7 @@ import os import tempfile from collections.abc import Callable +from importlib.metadata import version from pathlib import Path from typing import Any, Literal @@ -17,6 +18,7 @@ from anndata import AnnData from numpy.random import default_rng from packaging.version import Version +from pandas.testing import assert_series_equal from shapely import MultiPolygon, Polygon from upath import UPath from xarray import DataArray @@ -1289,13 +1291,13 @@ def test_read_sdata(tmp_path: Path, points: SpatialData) -> None: assert_spatial_data_objects_are_identical(sdata_from_path, sdata_from_zarr_group) -def test_sdata_with_nan_in_obs(tmp_path: Path) -> None: +@pytest.mark.parametrize("convert_strings_to_categoricals", (True, False)) +def test_sdata_with_nan_in_obs(tmp_path: Path, convert_strings_to_categoricals: bool) -> None: """Test writing SpatialData with mixed string/NaN values in obs works correctly. Regression test for https://github.com/scverse/spatialdata/issues/399 Previously this raised TypeError: expected unicode string, found nan. - Now the write succeeds, though NaN values in object-dtype columns are - converted to the string "nan" after round-trip. + Now the write succeeds, and NaN values are preserved round trip """ from spatialdata.models import TableModel @@ -1319,8 +1321,17 @@ def test_sdata_with_nan_in_obs(tmp_path: Path) -> None: assert sdata["table"].obs["column_only_region1"].iloc[1] is np.nan assert np.isnan(sdata["table"].obs["column_only_region2"].iloc[0]) + dtypes_before_writing = sdata["table"].obs.dtypes.copy() + path = tmp_path / "data.zarr" - sdata.write(path) + sdata.write(path, convert_table_strings_to_categoricals=convert_strings_to_categoricals) + + if convert_strings_to_categoricals: + expected_dtypes = dtypes_before_writing.copy() + expected_dtypes["column_only_region1"] = "category" + assert_series_equal(sdata["table"].obs.dtypes, expected_dtypes) + else: + assert_series_equal(sdata["table"].obs.dtypes, dtypes_before_writing) sdata2 = SpatialData.read(path) assert "column_only_region1" in sdata2["table"].obs.columns @@ -1329,8 +1340,12 @@ def test_sdata_with_nan_in_obs(tmp_path: Path) -> None: assert r1.iloc[0] == "string" assert r2.iloc[1] == 3 - if Version(pd.__version__) >= Version("3"): - assert pd.isna(r1.iloc[1]) - else: # After round-trip, NaN in object-dtype column becomes string "nan" on pandas 2 - assert r1.iloc[1] == "nan" assert np.isnan(r2.iloc[0]) + + if Version(version("pandas")) >= Version("3"): + assert pd.isna(r1.iloc[1]) + else: # After round-trip, NaN in object-dtype column becomes string + if convert_strings_to_categoricals: + assert pd.isna(r1.iloc[1]) + else: + assert r1.iloc[1] == "nan"