diff --git a/CHANGELOG.md b/CHANGELOG.md index 05346f0d..411f5882 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -15,6 +15,8 @@ and this project adheres to [Semantic Versioning][]. - {func}`~ehrdata.io.omop.setup_variables` and {func}`~ehrdata.io.omop.setup_interval_variables` with `aggregation_strategy="last"` (or `"first"`) no longer return a missing value when the last (or first) record of an interval has no value in the kept field. The aggregation now considers only the records that do carry a value, so the last (or first) *observed* value is kept. Consequently, records without a value no longer contribute a unit to the unit report, and `is_present` of the long format table marks the intervals in which a value was observed. ([#298](https://github.com/theislab/ehrdata/issues/298)) @eroell ### Fixed + - {func}`~ehrdata.io.write_h5ed` and {func}`~ehrdata.io.write_zarr` now write the columns of `.obs`, `.var` and `.tem` that anndata has no writer for: datetime columns are written as ISO 8601 strings, and `object` columns as a numeric dtype where their values allow it and as strings otherwise, with missing values staying missing. Previously, the output of {func}`~ehrdata.io.omop.setup_obs` could not be written at all, as the OMOP tables bring datetime columns, and `object` columns for the fields a database read finds empty. ([#303](https://github.com/theislab/ehrdata/issues/303)) @eroell + - {func}`~ehrdata.io.write_zarr` with `convert_strings_to_categoricals=True` no longer converts the string columns of the written object's own `.tem` to `categorical` dtype; the conversion now happens on the copy that is written. ([#303](https://github.com/theislab/ehrdata/issues/303)) @eroell - In {func}`~ehrdata.io.omop.setup_variables` and {func}`~ehrdata.io.omop.setup_interval_variables`, the units are now carried along the aggregation: from the very data point that is kept for `"last"`/`"first"`, and shared by the data points that are combined otherwise. Since combining values of different units is not meaningful, an `aggregation_strategy` other than `"last"`, `"first"` and `"count"` now raises a `NotImplementedError` if a feature has more than one `unit_concept_id`. Data tables without a `unit_concept_id` are unaffected, which in `setup_interval_variables` is all of them but `device_exposure` and `dose_era`. ([#300](https://github.com/theislab/ehrdata/issues/300)) @eroell - {func}`~ehrdata.io.omop.setup_variables` and {func}`~ehrdata.io.omop.setup_interval_variables` now mark `is_present` of the long format table for the intervals in which a value was observed with every `aggregation_strategy`, not only with `"last"`/`"first"`. ([#300](https://github.com/theislab/ehrdata/issues/300)) @eroell - {func}`~ehrdata.io.omop.setup_variables` and {func}`~ehrdata.io.omop.setup_interval_variables` with `aggregation_strategy="std"` no longer fail with `Catalog Error: Scalar Function with name std does not exist!`. The strategy is now mapped to DuckDB's `stddev_samp`. ([#300](https://github.com/theislab/ehrdata/issues/300)) @eroell diff --git a/src/ehrdata/io/_array_casting.py b/src/ehrdata/io/_array_casting.py index e8fcd3fa..9620e559 100644 --- a/src/ehrdata/io/_array_casting.py +++ b/src/ehrdata/io/_array_casting.py @@ -4,6 +4,7 @@ from typing import TYPE_CHECKING import numpy as np +import pandas as pd if TYPE_CHECKING: from ehrdata import EHRData @@ -63,3 +64,46 @@ def _cast_arrays_dtype_to_float_or_str_if_nonnumeric_object(edata: EHRData) -> E edata.layers[layer] = array.astype(str) return edata + + +def _cast_dataframe_columns_to_writable_dtype(df: pd.DataFrame, slot: str) -> pd.DataFrame: + """Return a copy of a dataframe of `.obs`, `.var` or `.tem` with the columns anndata cannot write cast. + + anndata has no writer for datetime columns, and h5py rejects an `object` column holding anything + but strings, which is what a database read leaves behind for a column it found empty. + Such columns are cast following the rule that `X` and the layers follow: to a numeric dtype, and + to a string dtype if that fails. Datetimes are written as ISO 8601 strings, from which + :func:`~ehrdata.infer_feature_types` recognizes the column as a date again. Missing values stay missing. + """ + df = df.copy() + cast_to_string = [] + + for column_name in df.columns: + column = df[column_name] + + if pd.api.types.is_datetime64_any_dtype(column): + # element-wise, since a mapped column of ISO 8601 strings is inferred back to datetimes by pandas + iso_8601 = np.array([None if pd.isna(value) else value.isoformat() for value in column], dtype=object) + df[column_name] = pd.Categorical(iso_8601) + cast_to_string.append(column_name) + + elif column.dtype == object: + observed_values = column[column.notna()] + if observed_values.empty: + df[column_name] = np.full(len(column), np.nan) + elif observed_values.map(lambda value: isinstance(value, str)).all(): + # anndata writes an all-string column itself, but not the missing values among them + df[column_name] = pd.Categorical(column) + else: + try: + df[column_name] = pd.to_numeric(column) + except (TypeError, ValueError): + df[column_name] = pd.Categorical(column.astype(str).where(column.notna())) + cast_to_string.append(column_name) + + if cast_to_string: + logger.warning( + f"Columns {cast_to_string} of .{slot} have a dtype that cannot be written, and are written as strings." + ) + + return df diff --git a/src/ehrdata/io/_ondisk.py b/src/ehrdata/io/_ondisk.py index 00ff0dee..84b74365 100644 --- a/src/ehrdata/io/_ondisk.py +++ b/src/ehrdata/io/_ondisk.py @@ -20,6 +20,7 @@ EHRDATA_ONDISK_VERSION, EHRDATA_ONDISK_VERSION_KEY, ) +from ehrdata.io._array_casting import _cast_dataframe_columns_to_writable_dtype from ehrdata.io._coo_codec import is_coo_group, read_coo if TYPE_CHECKING: @@ -70,6 +71,7 @@ def encode_for_disk(edata: EHRData) -> tuple[ad.AnnData, dict[str, sparse.COO]]: Dense 3D arrays are placed into the AnnData's ``.obsm`` under the reserved keys. Sparse 3D arrays, not writeable with AnnData, are pulled out into the returned ``sparse_3d_data`` mapping (keyed by the same reserved ``.obsm`` keys) + Columns of ``.obs`` and ``.var`` of a dtype anndata cannot write are cast, see ``_cast_dataframe_columns_to_writable_dtype``. """ _reject_stray_coo(edata) @@ -99,8 +101,8 @@ def encode_for_disk(edata: EHRData) -> tuple[ad.AnnData, dict[str, sparse.COO]]: adata = ad.AnnData( X=X, - obs=edata.obs.copy(), - var=edata.var.copy(), + obs=_cast_dataframe_columns_to_writable_dtype(edata.obs, "obs"), + var=_cast_dataframe_columns_to_writable_dtype(edata.var, "var"), uns=dict(edata.uns), obsm=obsm, varm=dict(edata.varm), diff --git a/src/ehrdata/io/h5ed.py b/src/ehrdata/io/h5ed.py index 18ca37e7..e5aef115 100644 --- a/src/ehrdata/io/h5ed.py +++ b/src/ehrdata/io/h5ed.py @@ -17,7 +17,11 @@ EHRDATA_ONDISK_VERSION_KEY, ) from ehrdata.core.ehrdata import _silence_anndata_nd_warning -from ehrdata.io._array_casting import _cast_arrays_dtype_to_float_or_str_if_nonnumeric_object, _cast_variables_to_float +from ehrdata.io._array_casting import ( + _cast_arrays_dtype_to_float_or_str_if_nonnumeric_object, + _cast_dataframe_columns_to_writable_dtype, + _cast_variables_to_float, +) from ehrdata.io._coo_codec import write_coo_h5 from ehrdata.io._ondisk import ( _check_020_ehrdata_on_disk_format, @@ -128,6 +132,7 @@ def write_h5ed( `.h5ed` is the ehrdata on-disk format. To write the file, `X` and `layers` cannot be written as `object` dtype. If any of these fields is of `object` dtype, this function will attempt to cast it to a numeric dtype; if this fails, the field will be casted to a string dtype. + The same holds for the columns of `.obs`, `.var` and `.tem`, which additionally cannot be written as a datetime dtype; datetime columns are written as ISO 8601 strings. Args: @@ -159,7 +164,7 @@ def write_h5ed( write_coo_h5( obsm_group.create_group(key), coo, compression=compression, compression_opts=compression_opts ) - ad.io.write_elem(f, "tem", edata.tem) + ad.io.write_elem(f, "tem", _cast_dataframe_columns_to_writable_dtype(edata.tem, "tem")) # Identify the file as ehrdata, namespaced to not clash with anndata's own encoding attrs. f.attrs[EHRDATA_ENCODING_TYPE_KEY] = EHRDATA_ENCODING_TYPE f.attrs[EHRDATA_ONDISK_VERSION_KEY] = str(EHRDATA_ONDISK_VERSION) diff --git a/src/ehrdata/io/zarr.py b/src/ehrdata/io/zarr.py index 545d8209..40ff686f 100644 --- a/src/ehrdata/io/zarr.py +++ b/src/ehrdata/io/zarr.py @@ -15,7 +15,11 @@ EHRDATA_ONDISK_VERSION, EHRDATA_ONDISK_VERSION_KEY, ) -from ehrdata.io._array_casting import _cast_arrays_dtype_to_float_or_str_if_nonnumeric_object, _cast_variables_to_float +from ehrdata.io._array_casting import ( + _cast_arrays_dtype_to_float_or_str_if_nonnumeric_object, + _cast_dataframe_columns_to_writable_dtype, + _cast_variables_to_float, +) from ehrdata.io._coo_codec import write_coo_zarr from ehrdata.io._ondisk import ( _check_020_ehrdata_on_disk_format, @@ -129,6 +133,7 @@ def write_zarr( To write to a `.zarr` store, `X`, and `layers` cannot be written as `object` dtype. If any of these fields is of `object` dtype, this function will attempt to cast it to a numeric dtype; if this fails, the field will be casted to a `str` dtype. + The same holds for the columns of `.obs`, `.var` and `.tem`, which additionally cannot be written as a datetime dtype; datetime columns are written as ISO 8601 strings. Args: @@ -148,11 +153,12 @@ def write_zarr( store = zarr.open_group(filename, mode="a", use_consolidated=False, zarr_format=3) adata, coo_obsm = encode_for_disk(edata) + tem = _cast_dataframe_columns_to_writable_dtype(edata.tem, "tem") if convert_strings_to_categoricals: adata.strings_to_categoricals(adata.obs) adata.strings_to_categoricals(adata.var) - adata.strings_to_categoricals(edata.tem) + adata.strings_to_categoricals(tem) # write_sharded this is a slightly modified version from https://anndata.readthedocs.io/en/stable/tutorials/zarr-v3.html # write_sharded is intended as a future blueprint of implementing better chunking defaults for ehrdata based based on real usecases @@ -194,7 +200,7 @@ def callback( # noqa: PLR0917 # signature is fixed by ad.experimental.write_di for key, coo in coo_obsm.items(): write_coo_zarr(obsm_group.require_group(key), coo) - ad.io.write_elem(store, "tem", edata.tem) + ad.io.write_elem(store, "tem", tem) store.attrs[EHRDATA_ENCODING_TYPE_KEY_ZARR] = EHRDATA_ENCODING_TYPE store.attrs[EHRDATA_ONDISK_VERSION_KEY] = EHRDATA_ONDISK_VERSION diff --git a/tests/conftest.py b/tests/conftest.py index ab9f4b4c..a143b309 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -56,6 +56,40 @@ def _assert_dtype_object_array_with_missing_values_equal(a: np.ndarray, b: np.nd assert np.array_equal(a, b) +def _edata_with_columns_anndata_cannot_write() -> EHRData: + """An EHRData whose `.obs`, `.var` and `.tem` hold the column dtypes anndata cannot write. + + Datetime columns, and `object` columns holding something else than strings: an all-null column as a + database read leaves it behind, and integers next to a missing value. + """ + obs = pd.DataFrame( + { + "birth_datetime": pd.to_datetime(["2000-01-01 01:02:03", None, "2001-02-03 04:05:06"]), + "location_id": np.array([None, None, None], dtype=object), + "care_site_id": np.array([1, None, 3], dtype=object), + }, + index=["0", "1", "2"], + ) + var = pd.DataFrame({"valid_start_date": pd.to_datetime(["1970-01-01", "1970-01-02"])}, index=["0", "1"]) + tem = pd.DataFrame({"timestamp": pd.to_datetime(["2000-01-01", "2000-01-02"])}) + + return EHRData(X=np.zeros((3, 2, 2)), obs=obs, var=var, tem=tem) + + +def _assert_columns_anndata_cannot_write_read_back(edata: EHRData, edata_read: EHRData): + assert edata_read.obs["birth_datetime"].tolist()[::2] == ["2000-01-01T01:02:03", "2001-02-03T04:05:06"] + assert edata_read.obs["birth_datetime"].isna().tolist() == [False, True, False] + assert edata_read.obs["location_id"].isna().all() + assert edata_read.obs["care_site_id"].tolist()[::2] == [1.0, 3.0] + assert edata_read.obs["care_site_id"].isna().tolist() == [False, True, False] + assert edata_read.var["valid_start_date"].tolist() == ["1970-01-01T00:00:00", "1970-01-02T00:00:00"] + assert edata_read.tem["timestamp"].tolist() == ["2000-01-01T00:00:00", "2000-01-02T00:00:00"] + + # the object that was written is left as it was + assert pd.api.types.is_datetime64_any_dtype(edata.obs["birth_datetime"]) + assert edata.obs["location_id"].dtype == object + + @pytest.fixture def csv_basic(): return pd.read_csv("tests/data/toy_csv/csv_basic.csv") diff --git a/tests/io/test_h5ed.py b/tests/io/test_h5ed.py index 398f1670..e7c75c4a 100644 --- a/tests/io/test_h5ed.py +++ b/tests/io/test_h5ed.py @@ -10,9 +10,11 @@ from scipy.sparse import issparse from tests.conftest import ( TEST_DATA_PATH, + _assert_columns_anndata_cannot_write_read_back, _assert_dtype_object_array_with_missing_values_equal, _assert_io_read, _assert_shape_matches, + _edata_with_columns_anndata_cannot_write, ) from ehrdata import EHRData @@ -389,3 +391,14 @@ def test_read_minimal_corpus_h5(): edata_020 = read_h5ed(path_020) _assert_shape_matches(edata_020, (3, 2, 2)) assert np.array_equal(np.asarray(edata_020.layers["tem_data"]), np.arange(3 * 2 * 2, dtype=float).reshape(3, 2, 2)) + + +def test_write_read_h5ed_casts_columns_anndata_cannot_write(tmp_path): + # datetime columns and `object` columns holding non-strings are cast on write instead of failing the write, + # as they did for the `.obs` and `.var` of an OMOP setup (https://github.com/theislab/ehrdata/issues/303) + edata = _edata_with_columns_anndata_cannot_write() + path = tmp_path / "columns_anndata_cannot_write.h5ed" + + write_h5ed(edata, path) + + _assert_columns_anndata_cannot_write_read_back(edata, read_h5ed(path)) diff --git a/tests/io/test_omop.py b/tests/io/test_omop.py index b305cd7c..aa94d2e8 100644 --- a/tests/io/test_omop.py +++ b/tests/io/test_omop.py @@ -2145,3 +2145,32 @@ def test_setup_variables_parquet(omop_connection_vanilla_parquet): [[np.nan, np.nan, np.nan, np.nan], [23.0, np.nan, np.nan, np.nan]], ] assert np.allclose(edata.layers[DEFAULT_TEM_LAYER_NAME], np.array(expected_data), equal_nan=True) + + +@pytest.mark.parametrize("observation_table", VANILLA_PERSONS_WITH_OBSERVATION_TABLE_ENTRY) +def test_write_read_h5ed_of_omop_setup(omop_connection_vanilla, observation_table, tmp_path): + # the datetime columns and the columns the OMOP tables leave empty must not block writing + # what setup_obs and setup_variables built (https://github.com/theislab/ehrdata/issues/303) + con = omop_connection_vanilla + edata = ed.io.omop.setup_obs(backend_handle=con, observation_table=observation_table) + edata = ed.io.omop.setup_variables( + edata, + backend_handle=con, + layer=DEFAULT_TEM_LAYER_NAME, + data_tables=["measurement"], + data_field_to_keep=["value_as_number"], + interval_length_number=1, + interval_length_unit="day", + num_intervals=4, + enrich_var_with_feature_info=True, + enrich_var_with_unit_info=True, + ) + path = tmp_path / f"{observation_table}.h5ed" + + ed.io.write_h5ed(edata, path) + edata_read = ed.io.read_h5ed(path) + + assert edata_read.shape == edata.shape + assert list(edata_read.obs.columns) == list(edata.obs.columns) + assert list(edata_read.var.columns) == list(edata.var.columns) + assert np.allclose(edata_read.layers[DEFAULT_TEM_LAYER_NAME], edata.layers[DEFAULT_TEM_LAYER_NAME], equal_nan=True) diff --git a/tests/io/test_zarr.py b/tests/io/test_zarr.py index 4f460b59..f13c7e8d 100644 --- a/tests/io/test_zarr.py +++ b/tests/io/test_zarr.py @@ -9,10 +9,12 @@ from scipy.sparse import issparse from tests.conftest import ( TEST_DATA_PATH, + _assert_columns_anndata_cannot_write_read_back, _assert_dtype_object_array_with_missing_values_equal, _assert_io_read, _assert_shape_matches, _check_aligned_anndata_parts_equal, + _edata_with_columns_anndata_cannot_write, ) from ehrdata.core.constants import EHRDATA_ONDISK_VERSION @@ -316,3 +318,14 @@ def test_read_minimal_corpus_zarr(): edata_020 = read_zarr(TEST_PATH_ZARR / "edata_minimal_v0_2_0.ehrdata.zarr") _assert_shape_matches(edata_020, (3, 2, 2)) assert np.array_equal(np.asarray(edata_020.layers["tem_data"]), np.arange(3 * 2 * 2, dtype=float).reshape(3, 2, 2)) + + +def test_write_read_zarr_casts_columns_anndata_cannot_write(tmp_path): + # datetime columns and `object` columns holding non-strings are cast on write instead of failing the write, + # as they did for the `.obs` and `.var` of an OMOP setup (https://github.com/theislab/ehrdata/issues/303) + edata = _edata_with_columns_anndata_cannot_write() + path = tmp_path / "columns_anndata_cannot_write.zarr" + + write_zarr(edata, path) + + _assert_columns_anndata_cannot_write_read_back(edata, read_zarr(path))