Skip to content
Draft
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
2 changes: 2 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
44 changes: 44 additions & 0 deletions src/ehrdata/io/_array_casting.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@
from typing import TYPE_CHECKING

import numpy as np
import pandas as pd

if TYPE_CHECKING:
from ehrdata import EHRData
Expand Down Expand Up @@ -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
6 changes: 4 additions & 2 deletions src/ehrdata/io/_ondisk.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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)

Expand Down Expand Up @@ -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),
Expand Down
9 changes: 7 additions & 2 deletions src/ehrdata/io/h5ed.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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)
Expand Down
12 changes: 9 additions & 3 deletions src/ehrdata/io/zarr.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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:
Expand All @@ -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
Expand Down Expand Up @@ -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
34 changes: 34 additions & 0 deletions tests/conftest.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Expand Down
13 changes: 13 additions & 0 deletions tests/io/test_h5ed.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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))
29 changes: 29 additions & 0 deletions tests/io/test_omop.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
13 changes: 13 additions & 0 deletions tests/io/test_zarr.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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))
Loading