From 55781c22da9b759c28274b02b32d07cc5be92bac Mon Sep 17 00:00:00 2001 From: ilan-gold Date: Fri, 30 May 2025 11:33:28 +0200 Subject: [PATCH 1/3] (chore): add thorough test cases --- properties/test_pandas_roundtrip.py | 20 ++++++++++++++++---- xarray/core/dataset.py | 6 ++++-- xarray/core/extension_array.py | 3 --- xarray/core/variable.py | 9 ++------- 4 files changed, 22 insertions(+), 16 deletions(-) diff --git a/properties/test_pandas_roundtrip.py b/properties/test_pandas_roundtrip.py index 04babad1a23..fea14baafe3 100644 --- a/properties/test_pandas_roundtrip.py +++ b/properties/test_pandas_roundtrip.py @@ -138,22 +138,34 @@ def test_roundtrip_pandas_dataframe_datetime(df) -> None: "extension_array", [ pd.Categorical(["a", "b", "c"]), - pd.array([1, 2, 3], dtype="int64"), + pd.array([1, 2, 3], dtype="int64[pyarrow]"), pd.array(["a", "b", "c"], dtype="string"), pd.arrays.IntervalArray( [pd.Interval(0, 1), pd.Interval(1, 5), pd.Interval(2, 6)] ), + pd.arrays.TimedeltaArray._from_sequence(pd.TimedeltaIndex(["1h", "2h", "3h"])), + pd.arrays.DatetimeArray._from_sequence( + pd.DatetimeIndex(["2023-01-01", "2023-01-02", "2023-01-03"], freq="D") + ), np.array([1, 2, 3], dtype="int64"), ], + ids=["cat", "pyarrow", "string", "interval", "timedelta", "datetime", "numpy"], ) -def test_roundtrip_1d_pandas_extension_array(extension_array) -> None: +@pytest.mark.parametrize("is_index", [True, False]) +def test_roundtrip_1d_pandas_extension_array(extension_array, is_index) -> None: df = pd.DataFrame({"arr": extension_array}) + if is_index: + df = df.set_index("arr") arr = xr.Dataset.from_dataframe(df)["arr"] roundtripped = arr.to_pandas() - assert (df["arr"] == roundtripped).all() + df_arr_to_test = df.index if is_index else df["arr"] + assert (df_arr_to_test == roundtripped).all() # `NumpyExtensionArray` types are not roundtripped, including `StringArray` which subtypes. if isinstance(extension_array, pd.arrays.NumpyExtensionArray): assert isinstance(arr.data, np.ndarray) else: - assert df["arr"].dtype == roundtripped.dtype + assert ( + df_arr_to_test.dtype + == (roundtripped.index if is_index else roundtripped).dtype + ) xr.testing.assert_identical(arr, roundtripped.to_xarray()) diff --git a/xarray/core/dataset.py b/xarray/core/dataset.py index 76422a09af4..319b1e01f94 100644 --- a/xarray/core/dataset.py +++ b/xarray/core/dataset.py @@ -99,7 +99,6 @@ parse_dims_as_set, ) from xarray.core.variable import ( - UNSUPPORTED_EXTENSION_ARRAY_TYPES, IndexVariable, Variable, as_variable, @@ -7272,7 +7271,10 @@ def from_dataframe(cls, dataframe: pd.DataFrame, sparse: bool = False) -> Self: extension_arrays = [] for k, v in dataframe.items(): if not is_extension_array_dtype(v) or isinstance( - v.array, UNSUPPORTED_EXTENSION_ARRAY_TYPES + v.array, + pd.arrays.DatetimeArray + | pd.arrays.TimedeltaArray + | pd.arrays.NumpyExtensionArray, ): arrays.append((k, np.asarray(v))) else: diff --git a/xarray/core/extension_array.py b/xarray/core/extension_array.py index 9377b442aab..0c0312e0e23 100644 --- a/xarray/core/extension_array.py +++ b/xarray/core/extension_array.py @@ -82,9 +82,6 @@ class PandasExtensionArray(Generic[T_ExtensionArray], NDArrayMixin): def __post_init__(self): if not isinstance(self.array, pd.api.extensions.ExtensionArray): raise TypeError(f"{self.array} is not an pandas ExtensionArray.") - # This does not use the UNSUPPORTED_EXTENSION_ARRAY_TYPES whitelist because - # we do support extension arrays from datetime, for example, that need - # duck array support internally via this class. if isinstance(self.array, pd.arrays.NumpyExtensionArray): raise TypeError( "`NumpyExtensionArray` should be converted to a numpy array in `xarray` internally." diff --git a/xarray/core/variable.py b/xarray/core/variable.py index 32fe55e2ac8..6b952ae7d3b 100644 --- a/xarray/core/variable.py +++ b/xarray/core/variable.py @@ -63,11 +63,6 @@ ) # https://github.com/python/mypy/issues/224 BASIC_INDEXING_TYPES = integer_types + (slice,) -UNSUPPORTED_EXTENSION_ARRAY_TYPES = ( - pd.arrays.DatetimeArray, - pd.arrays.TimedeltaArray, - pd.arrays.NumpyExtensionArray, -) if TYPE_CHECKING: from xarray.core.types import ( @@ -196,7 +191,7 @@ def _maybe_wrap_data(data): """ if isinstance(data, pd.Index): return PandasIndexingAdapter(data) - if isinstance(data, UNSUPPORTED_EXTENSION_ARRAY_TYPES): + if isinstance(data, pd.arrays.NumpyExtensionArray): return data.to_numpy() if isinstance(data, pd.api.extensions.ExtensionArray): return PandasExtensionArray(data) @@ -262,7 +257,7 @@ def convert_non_numpy_type(data): if ( isinstance(data, pd.Series) and pd.api.types.is_extension_array_dtype(data) - and not isinstance(data.array, UNSUPPORTED_EXTENSION_ARRAY_TYPES) + and not isinstance(data.array, pd.arrays.NumpyExtensionArray) ): pandas_data = data.array else: From 5bdd8a7220a5eb89eda49761dadff3ecfe566a1a Mon Sep 17 00:00:00 2001 From: ilan-gold Date: Fri, 30 May 2025 11:34:16 +0200 Subject: [PATCH 2/3] (fix): cleanup test --- xarray/tests/test_variable.py | 30 +++++++++++++----------------- 1 file changed, 13 insertions(+), 17 deletions(-) diff --git a/xarray/tests/test_variable.py b/xarray/tests/test_variable.py index 1e7c32dec1e..d9e094bd98e 100644 --- a/xarray/tests/test_variable.py +++ b/xarray/tests/test_variable.py @@ -15,6 +15,7 @@ from xarray import DataArray, Dataset, IndexVariable, Variable, set_options from xarray.core import dtypes, duck_array_ops, indexing from xarray.core.common import full_like, ones_like, zeros_like +from xarray.core.extension_array import PandasExtensionArray from xarray.core.indexing import ( BasicIndexer, CopyOnWriteArray, @@ -2757,15 +2758,15 @@ def test_tz_datetime(self) -> None: warnings.simplefilter("ignore") actual: T_DuckArray = as_compatible_data(times_s) assert actual.array == times_s - assert actual.array.dtype == pd.DatetimeTZDtype("s", tz) # type: ignore[arg-type] + assert actual.array.dtype == times_s.dtype # type: ignore[arg-type] series = pd.Series(times_s) with warnings.catch_warnings(): warnings.simplefilter("ignore") actual2: T_DuckArray = as_compatible_data(series) - np.testing.assert_array_equal(actual2, np.asarray(series.values)) - assert actual2.dtype == np.dtype("datetime64[s]") + np.testing.assert_array_equal(actual2, np.asarray(series.array)) + assert actual2.dtype == times_s.dtype def test_full_like(self) -> None: # For more thorough tests, see test_variable.py @@ -3096,8 +3097,13 @@ def test_datetime_conversion(values, unit) -> None: else: # The only case where a non-datetime64 dtype can occur currently is in # the case that the variable is backed by a timezone-aware - # DatetimeIndex, and thus is hidden within the PandasIndexingAdapter class. - assert isinstance(var._data, PandasIndexingAdapter) + # DatetimeIndex/DateTimeArray, and thus is hidden within the PandasIndexingAdapter/PandasExtensionArray class. + assert isinstance( + var._data, + PandasIndexingAdapter + if isinstance(values, pd.DatetimeIndex) + else PandasExtensionArray, + ) assert var._data.array.dtype == pd.DatetimeTZDtype( "ns", pytz.timezone("America/New_York") ) @@ -3132,19 +3138,9 @@ def test_pandas_two_only_datetime_conversion_warnings( ) -> None: # todo: check for redundancy (suggested per review) var = Variable(["time"], data.astype(dtype)) # type: ignore[arg-type] - - # we internally convert series to numpy representations to avoid too much nastiness with extension arrays - # when calling data.array e.g., with NumpyExtensionArrays - if isinstance(data, pd.Series): - assert var.dtype == np.dtype("datetime64[s]") - elif var.dtype.kind == "M": - assert var.dtype == dtype - else: - # The only case where a non-datetime64 dtype can occur currently is in - # the case that the variable is backed by a timezone-aware - # DatetimeIndex, and thus is hidden within the PandasIndexingAdapter class. + assert var.dtype == dtype + if isinstance(data, pd.DatetimeIndex): assert isinstance(var._data, PandasIndexingAdapter) - assert var._data.array.dtype == pd.DatetimeTZDtype("s", tz_ny) @pytest.mark.parametrize( From 6da431ebeb6c2271f2798e4df26f0dcff67d99b2 Mon Sep 17 00:00:00 2001 From: ilan-gold Date: Fri, 30 May 2025 14:21:47 +0200 Subject: [PATCH 3/3] (fix): remove check on `from_dataframe` --- xarray/core/dataset.py | 5 +---- 1 file changed, 1 insertion(+), 4 deletions(-) diff --git a/xarray/core/dataset.py b/xarray/core/dataset.py index 319b1e01f94..e9f7d66d7b3 100644 --- a/xarray/core/dataset.py +++ b/xarray/core/dataset.py @@ -7271,10 +7271,7 @@ def from_dataframe(cls, dataframe: pd.DataFrame, sparse: bool = False) -> Self: extension_arrays = [] for k, v in dataframe.items(): if not is_extension_array_dtype(v) or isinstance( - v.array, - pd.arrays.DatetimeArray - | pd.arrays.TimedeltaArray - | pd.arrays.NumpyExtensionArray, + v.array, pd.arrays.NumpyExtensionArray ): arrays.append((k, np.asarray(v))) else: