diff --git a/src/xregrid/accessors.py b/src/xregrid/accessors.py index 6881725..5604733 100644 --- a/src/xregrid/accessors.py +++ b/src/xregrid/accessors.py @@ -1,6 +1,6 @@ from __future__ import annotations -from typing import Any +from typing import Any, Union import xarray as xr @@ -24,27 +24,75 @@ def __init__(self, xarray_obj: xr.DataArray): """ self._obj = xarray_obj - def to(self, target_grid: xr.Dataset, **kwargs: Any) -> xr.DataArray: + def to( + self, target_grid: Union[xr.Dataset, Regridder], **kwargs: Any + ) -> xr.DataArray: """ - Regrid the DataArray to a target grid. + Regrid the DataArray to a target grid or using a pre-computed Regridder. Parameters ---------- - target_grid : xr.Dataset - The target grid dataset. + target_grid : xr.Dataset or Regridder + The target grid dataset or an existing Regridder instance. **kwargs : Any - Arguments passed to the Regridder constructor. + Arguments passed to the Regridder constructor if target_grid is a Dataset. Returns ------- xr.DataArray The regridded DataArray. """ + if isinstance(target_grid, Regridder): + return target_grid(self._obj) + # Convert DataArray to Dataset to ensure compatibility with Regridder + # if the Regridder needs to inspect the source grid. source_ds = self._obj.to_dataset(name="_tmp_data") regridder = Regridder(source_ds, target_grid, **kwargs) return regridder(self._obj) + def get_regridder(self, target_grid: xr.Dataset, **kwargs: Any) -> Regridder: + """ + Create a Regridder instance for this DataArray. + + Parameters + ---------- + target_grid : xr.Dataset + The target grid dataset. + **kwargs : Any + Arguments passed to the Regridder constructor. + + Returns + ------- + Regridder + The initialized Regridder instance. + """ + source_ds = self._obj.to_dataset(name="_tmp_data") + return Regridder(source_ds, target_grid, **kwargs) + + def plot_diagnostics( + self, target_grid: xr.Dataset, mode: str = "static", **kwargs: Any + ) -> Any: + """ + Visualize regridding diagnostics between this DataArray and a target grid. + + Parameters + ---------- + target_grid : xr.Dataset + The target grid dataset. + mode : str, default 'static' + The plotting mode: 'static' or 'interactive'. + **kwargs : Any + Arguments passed to Regridder.plot_diagnostics. + + Returns + ------- + Any + The plot object. + """ + regridder = self.get_regridder(target_grid, **kwargs) + return regridder.plot_diagnostics(mode=mode, **kwargs) + @xr.register_dataset_accessor("regrid") class RegridDatasetAccessor: @@ -63,21 +111,67 @@ def __init__(self, xarray_obj: xr.Dataset): """ self._obj = xarray_obj - def to(self, target_grid: xr.Dataset, **kwargs: Any) -> xr.Dataset: + def to( + self, target_grid: Union[xr.Dataset, Regridder], **kwargs: Any + ) -> xr.Dataset: """ - Regrid the Dataset to a target grid. + Regrid the Dataset to a target grid or using a pre-computed Regridder. Parameters ---------- - target_grid : xr.Dataset - The target grid dataset. + target_grid : xr.Dataset or Regridder + The target grid dataset or an existing Regridder instance. **kwargs : Any - Arguments passed to the Regridder constructor. + Arguments passed to the Regridder constructor if target_grid is a Dataset. Returns ------- xr.Dataset The regridded Dataset. """ + if isinstance(target_grid, Regridder): + return target_grid(self._obj) + regridder = Regridder(self._obj, target_grid, **kwargs) return regridder(self._obj) + + def get_regridder(self, target_grid: xr.Dataset, **kwargs: Any) -> Regridder: + """ + Create a Regridder instance for this Dataset. + + Parameters + ---------- + target_grid : xr.Dataset + The target grid dataset. + **kwargs : Any + Arguments passed to the Regridder constructor. + + Returns + ------- + Regridder + The initialized Regridder instance. + """ + return Regridder(self._obj, target_grid, **kwargs) + + def plot_diagnostics( + self, target_grid: xr.Dataset, mode: str = "static", **kwargs: Any + ) -> Any: + """ + Visualize regridding diagnostics between this Dataset and a target grid. + + Parameters + ---------- + target_grid : xr.Dataset + The target grid dataset. + mode : str, default 'static' + The plotting mode: 'static' or 'interactive'. + **kwargs : Any + Arguments passed to Regridder.plot_diagnostics. + + Returns + ------- + Any + The plot object. + """ + regridder = self.get_regridder(target_grid, **kwargs) + return regridder.plot_diagnostics(mode=mode, **kwargs) diff --git a/tests/test_accessors.py b/tests/test_accessors.py new file mode 100644 index 0000000..8db94a7 --- /dev/null +++ b/tests/test_accessors.py @@ -0,0 +1,104 @@ +from __future__ import annotations + +import numpy as np +import xarray as xr + +try: + import dask.array as da +except ImportError: + da = None + +from xregrid import Regridder, create_global_grid + + +def test_dataarray_accessor_to_regridder(): + """ + Verify DataArray accessor works with a Regridder instance (Eager & Lazy). + + Follows the Aero Protocol 'Double-Check' Rule. + """ + ds_src = create_global_grid(10, 10) + ds_tgt = create_global_grid(5, 5) + + # 1. Eager (NumPy) + da_src = xr.DataArray( + np.random.rand(18, 36), + dims=("lat", "lon"), + coords={"lat": ds_src.lat, "lon": ds_src.lon}, + name="test_da", + ) + + # Create regridder once + regridder = da_src.regrid.get_regridder(ds_tgt) + assert isinstance(regridder, Regridder) + + # Regrid using target_grid (Dataset) + res_with_ds = da_src.regrid.to(ds_tgt) + + # Regrid using regridder instance (New functionality) + res_with_regridder = da_src.regrid.to(regridder) + + # Verification + xr.testing.assert_allclose(res_with_ds, res_with_regridder) + + # 2. Lazy (Dask) + if da is not None: + da_lazy = da_src.chunk({"lat": 9, "lon": 9}) + res_lazy = da_lazy.regrid.to(regridder) + + assert hasattr(res_lazy.data, "dask") + # Numerical Verification + xr.testing.assert_allclose(res_with_ds, res_lazy.compute()) + + +def test_dataset_accessor_to_regridder(): + """ + Verify Dataset accessor works with a Regridder instance. + """ + ds_src = create_global_grid(10, 10) + ds_tgt = create_global_grid(5, 5) + + ds = xr.Dataset( + {"var1": (("lat", "lon"), np.random.rand(18, 36))}, + coords={"lat": ds_src.lat, "lon": ds_src.lon}, + ) + + # Create regridder once + regridder = ds.regrid.get_regridder(ds_tgt) + assert isinstance(regridder, Regridder) + + # Regrid using target_grid (Dataset) + res_with_ds = ds.regrid.to(ds_tgt) + + # Regrid using regridder instance + res_with_regridder = ds.regrid.to(regridder) + + # Verification + xr.testing.assert_allclose(res_with_ds, res_with_regridder) + + +def test_accessor_plot_diagnostics_smoke(): + """ + Smoke test for accessor plotting methods. + """ + import matplotlib.pyplot as plt + + plt.switch_backend("Agg") + + ds_src = create_global_grid(30, 30) + ds_tgt = create_global_grid(10, 10) + da = xr.DataArray( + np.random.rand(6, 12), + dims=("lat", "lon"), + coords={"lat": ds_src.lat, "lon": ds_src.lon}, + ) + + # Smoke test: simply check if it runs without error + fig = da.regrid.plot_diagnostics(ds_tgt) + assert fig is not None + plt.close(fig) + + ds = da.to_dataset(name="test") + fig2 = ds.regrid.plot_diagnostics(ds_tgt) + assert fig2 is not None + plt.close(fig2)