diff --git a/docs/releases.md b/docs/releases.md index b38a0cb..e15b834 100644 --- a/docs/releases.md +++ b/docs/releases.md @@ -10,6 +10,8 @@ ### 🐛 Bug Fixes +* Fix transport and regrid for non-lat/lon grids (e.g. stereographic) ([#21](https://github.com/nansencenter/xhycom/pull/21)) + ### 📚 Documentation ### 🔧 Internal diff --git a/tests/data/_subset_tp5_stereo.py b/tests/data/_subset_tp5_stereo.py new file mode 100644 index 0000000..88915dc --- /dev/null +++ b/tests/data/_subset_tp5_stereo.py @@ -0,0 +1,58 @@ +"""Generate the bundled stereographic (TOPAZ5) target-grid test fixture. + +The full TP5 stereographic output is 1137×1185 cells with multiple 3-D fields +(~1 GB). For tests we need only a small curvilinear target covering the TP0 +model domain (Nordic Seas, ~lat 60–80 N, lon -20 to 20 E), spatially coarsened +to a few dozen cells. + +Run from a machine that can see the source file:: + + python tests/data/_subset_tp5_stereo.py + +Source: /cluster/projects/nn2993k/nlo043/TP5a0.06/staged/Hy2.2/archm_1993_01.nc +Re-run only if the fixture needs regenerating; the product is committed. +""" +import os + +import numpy as np +import xarray as xr + +SRC = ("/cluster/projects/nn2993k/nlo043/TP5a0.06/staged/Hy2.2/archm_1993_01.nc") +DST = os.path.join(os.path.dirname(os.path.abspath(__file__)), + "tp5_stereo_subset.nc") + +# Region overlapping the TP0 test domain (lat 60–80 N, lon -20 to 20 E) in +# the stereographic projected grid (y/x in degrees of the projection). +# Indices derived by masking on latitude/longitude of the full grid. +Y_SLICE = slice(139, 639) # ~500 rows covering approx lat 60–80 N +X_SLICE = slice(626, 1126) # ~500 cols covering approx lon -20 to 20 E +STRIDE = 20 # → ~25×25 cells +DEPTH_N = 3 # keep only the shallowest depth levels + + +def main() -> None: + ds = xr.open_dataset(SRC) + sub = ds.isel(y=Y_SLICE, x=X_SLICE).isel( + y=slice(None, None, STRIDE), + x=slice(None, None, STRIDE), + depth=slice(None, DEPTH_N), + time=0, + ) + # Keep geographic coords, depth levels, bathymetry, and one tracer for + # mask derivation (thetao NaN at land and below seafloor). + keep_vars = ["model_depth", "thetao"] + sub = sub[keep_vars] + # Ensure longitude/latitude 2-D coords are included (they are non-index + # coords on the y/x dims — xarray carries them through isel). + encoding = {v: {"zlib": True, "complevel": 4} for v in sub.data_vars} + sub.to_netcdf(DST, encoding=encoding) + print(f"wrote {DST} ({os.path.getsize(DST) / 1e3:.0f} KB)") + print("dims:", dict(sub.sizes)) + print("coords:", list(sub.coords)) + # Sanity: lon/lat should span the TP0 domain. + print(f"lon range: {float(sub.longitude.min()):.1f} .. {float(sub.longitude.max()):.1f}") + print(f"lat range: {float(sub.latitude.min()):.1f} .. {float(sub.latitude.max()):.1f}") + + +if __name__ == "__main__": + main() diff --git a/tests/data/tp5_stereo_subset.nc b/tests/data/tp5_stereo_subset.nc new file mode 100644 index 0000000..b3dac1d Binary files /dev/null and b/tests/data/tp5_stereo_subset.nc differ diff --git a/tests/test_regrid.py b/tests/test_regrid.py index 8215cd6..4a8a7a4 100644 --- a/tests/test_regrid.py +++ b/tests/test_regrid.py @@ -646,6 +646,105 @@ def test_apply_target_mask_surface_only_derived(): ) +def _curvilinear_target(ny: int = 8, nx: int = 8): + """A small synthetic curvilinear (2-D lon/lat) target grid. + + Mimics the structure of a TOPAZ5 stereographic output: dims are ``y`` and + ``x`` (1-D projected coordinates), and ``longitude`` / ``latitude`` are + 2-D non-index coordinates carrying the geographic position of each cell. + """ + x1d = np.linspace(2.0, 8.0, nx) + y1d = np.linspace(42.0, 48.0, ny) + lon2d, lat2d = np.meshgrid(x1d, y1d) + depth = np.array([5.0, 50.0, 200.0]) + return xr.Dataset( + {"thetao": (("depth", "y", "x"), np.ones((len(depth), ny, nx)))}, + coords={ + "y": y1d, + "x": x1d, + "latitude": (("y", "x"), lat2d), + "longitude": (("y", "x"), lon2d), + "depth": depth, + }, + ) + + +@pytest.fixture +def tp5_stereo(tmp_path): + """Small stereographic fixture from tests/data/tp5_stereo_subset.nc.""" + import os + + path = os.path.join(os.path.dirname(__file__), "data", "tp5_stereo_subset.nc") + if not os.path.exists(path): + pytest.skip( + "tp5_stereo_subset.nc not found — regenerate with _subset_tp5_stereo.py" + ) + return xr.open_dataset(path) + + +# --------------------------------------------------------------------------- +# Curvilinear target tests +# --------------------------------------------------------------------------- +def test_subset_target_skips_curvilinear(): + """_subset_target returns the target unchanged for a 2-D curvilinear grid.""" + from xhycom._regrid import _subset_target + + ds = _curvilinear_ds() # source lon 0..10, lat 40..50 + tgt = _curvilinear_target() + result = _subset_target(tgt, ds) + assert result is tgt + + +def test_regrid_horizontal_curvilinear_target_recovers_field(): + """Bilinear regrid to a curvilinear (2-D lon/lat) target recovers a constant field.""" + pytest.importorskip("xesmf") + ds = _curvilinear_ds() # source lon 0..10, lat 40..50, temp = lat + tgt = _curvilinear_target() # target lon 2..8, lat 42..48 — fully inside source + out = xhycom.regrid_horizontal(ds, target=tgt, method="bilinear") + # Dims must be y/x (from the target), not lat/lon. + assert "y" in out.dims and "x" in out.dims + assert "lat" not in out.dims and "lon" not in out.dims + # Geographic coords must be attached with the target's names. + assert "latitude" in out.coords and "longitude" in out.coords + assert out["latitude"].dims == ("y", "x") + # temp == lat everywhere; bilinear must recover latitude at interior points. + np.testing.assert_allclose( + out["temp"].isel(k=0).values, + out["latitude"].values, + atol=5e-3, + ) + + +def test_regrid_horizontal_curvilinear_target_yx_coords_attached(): + """y/x index coordinates from the target are preserved on the output.""" + pytest.importorskip("xesmf") + ds = _curvilinear_ds() + tgt = _curvilinear_target() + out = xhycom.regrid_horizontal(ds, target=tgt, method="bilinear") + assert "y" in out.coords and "x" in out.coords + np.testing.assert_array_equal(out["y"].values, tgt["y"].values) + np.testing.assert_array_equal(out["x"].values, tgt["x"].values) + + +def test_regrid_horizontal_curvilinear_conservative_warns(): + """Requesting conservative for a curvilinear target warns and falls back to bilinear.""" + pytest.importorskip("xesmf") + ds = _curvilinear_ds() + tgt = _curvilinear_target() + with pytest.warns(UserWarning, match="curvilinear"): + out = xhycom.regrid_horizontal(ds, target=tgt, method="conservative") + assert "y" in out.dims # bilinear fallback produced output + + +def test_regrid_horizontal_tp5_stereo_fixture(tp5_stereo): + """Regridding to the bundled TP5 stereographic fixture does not raise.""" + pytest.importorskip("xesmf") + ds = _curvilinear_ds() + out = xhycom.regrid_horizontal(ds, target=tp5_stereo, method="bilinear") + assert "y" in out.dims and "x" in out.dims + assert "latitude" in out.coords and "longitude" in out.coords + + def test_regrid_invalid_order(): ds = _curvilinear_ds() with pytest.raises(ValueError, match="order"): diff --git a/xhycom/_regrid.py b/xhycom/_regrid.py index 497bb7e..7f3863b 100644 --- a/xhycom/_regrid.py +++ b/xhycom/_regrid.py @@ -30,6 +30,7 @@ import hashlib import json import os +import warnings import numpy as np import xarray as xr @@ -77,7 +78,7 @@ def _load_grid(grid: xr.Dataset | str | None) -> xr.Dataset | None: def _target_lonlat(tgt: xr.Dataset) -> tuple[np.ndarray, np.ndarray]: - """1-D target longitudes / latitudes from a grid Dataset.""" + """Target longitudes / latitudes from a grid Dataset (1-D or 2-D).""" lon = tgt["longitude"] if "longitude" in tgt.variables else tgt["lon"] lat = tgt["latitude"] if "latitude" in tgt.variables else tgt["lat"] return np.asarray(lon.values), np.asarray(lat.values) @@ -119,6 +120,10 @@ def _subset_target(tgt: xr.Dataset, ds: xr.Dataset, pad: float = 1.0) -> xr.Data longitude box rather than risk dropping covered cells. """ lon_name, lat_name = _lonlat_names(tgt) + # Curvilinear targets (2-D lon/lat) are already regional; subsetting + # along 1-D dims using a 2-D boolean mask is ill-defined, so skip it. + if tgt[lon_name].ndim > 1: + return tgt tlon = np.asarray(tgt[lon_name].values) tlat = np.asarray(tgt[lat_name].values) slat = np.asarray(ds["lat"].values) @@ -605,21 +610,39 @@ def regrid_horizontal( lon = np.asarray(lon) lat = np.asarray(lat) - target_ds = xr.Dataset({"lat": (["lat"], lat), "lon": (["lon"], lon)}) - - conservative = method.startswith("conservative") - - # Conservative remapping needs cell corner bounds on both grids. - if conservative: - src = _add_source_bounds(src, grid) - # Latitude edges must stay within [-90, 90]: midpoint extrapolation of a - # target row sitting on the pole (e.g. GLORYS' top row at exactly 90 N) - # otherwise lands a cell corner past the pole — an invalid spherical - # coordinate. Clamping caps that cell at the pole instead. - target_ds = target_ds.assign( - lon_b=("lon_b", _edges_1d(lon)), - lat_b=("lat_b", np.clip(_edges_1d(lat), -90.0, 90.0)), - ) + curvilinear = lon.ndim == 2 + + if curvilinear: + # Curvilinear targets (e.g. a stereographic Arctic grid): 2-D lon/lat + # require a 2-D xESMF target dataset. Conservative remapping is not + # supported for curvilinear targets (cell-corner bounds in geographic + # coordinates are non-trivial for projected grids), so fall back to + # bilinear with a warning. + if method.startswith("conservative"): + warnings.warn( + "Conservative regridding is not supported for curvilinear " + "targets (2-D lon/lat); falling back to 'bilinear'.", + UserWarning, + stacklevel=3, + ) + method = "bilinear" + target_ds = xr.Dataset({"lat": (["y", "x"], lat), "lon": (["y", "x"], lon)}) + conservative = False + else: + target_ds = xr.Dataset({"lat": (["lat"], lat), "lon": (["lon"], lon)}) + conservative = method.startswith("conservative") + # Conservative remapping needs cell corner bounds on both grids. + if conservative: + src = _add_source_bounds(src, grid) + # Latitude edges must stay within [-90, 90]: midpoint extrapolation + # of a target row sitting on the pole (e.g. GLORYS' top row at + # exactly 90 N) otherwise lands a cell corner past the pole — an + # invalid spherical coordinate. Clamping caps that cell at the + # pole instead. + target_ds = target_ds.assign( + lon_b=("lon_b", _edges_1d(lon)), + lat_b=("lat_b", np.clip(_edges_1d(lat), -90.0, 90.0)), + ) # Thickness-weight layered fields for conservative remapping so the # volume-integrated content (field * layer thickness) is conserved: remap @@ -664,10 +687,26 @@ def regrid_horizontal( out[v] = out[v] / denom out[v].attrs = layer_attrs[v] - out["lon"].attrs.setdefault("standard_name", "longitude") - out["lon"].attrs.setdefault("units", "degrees_east") - out["lat"].attrs.setdefault("standard_name", "latitude") - out["lat"].attrs.setdefault("units", "degrees_north") + if curvilinear: + # xESMF does not carry non-index geographic coords to the output; + # re-attach them using the target's coordinate names. + lon_cname, lat_cname = ( + _lonlat_names(tgt) if tgt is not None else ("longitude", "latitude") + ) + extra: dict = { + lon_cname: (("y", "x"), lon), + lat_cname: (("y", "x"), lat), + } + if tgt is not None: + for d in ("y", "x"): + if d in tgt.coords and d in out.dims: + extra[d] = tgt[d] + out = out.assign_coords(extra) + else: + out["lon"].attrs.setdefault("standard_name", "longitude") + out["lon"].attrs.setdefault("units", "degrees_east") + out["lat"].attrs.setdefault("standard_name", "latitude") + out["lat"].attrs.setdefault("units", "degrees_north") if tgt is not None and apply_target_mask: out = _apply_target_mask(out, tgt, surface_only=True) diff --git a/xhycom/_transport.py b/xhycom/_transport.py index d9c1eb1..7e0a61b 100644 --- a/xhycom/_transport.py +++ b/xhycom/_transport.py @@ -206,27 +206,41 @@ def transport( if op not in _OPS: raise ValueError(f"Unknown operator {op!r}. Use: {sorted(_OPS)}") - # For regular (1-D coordinate) grids use coordinate-value selection so the + # For regular (1-D lat/lon coordinate) grids use coordinate-value selection so the # result is correct even when ds is a spatial subset of the grid that was # originally passed to resolve() (e.g. HYCOM regridded to a clipped GLORYS). - if resolved.y_dim in ds.coords and ds[resolved.y_dim].ndim == 1: + # Guard: only treat the y dimension as a lat coordinate if cell_lat values + # actually fall within its range — otherwise (e.g. TOPAZ's stereographic y + # coordinate runs -55..55, not -90..90) fall back to integer isel(). + _y_is_latlon = ( + resolved.y_dim in ds.coords + and ds[resolved.y_dim].ndim == 1 + and float(ds[resolved.y_dim].min()) <= float(np.min(resolved.cell_lat)) + and float(np.max(resolved.cell_lat)) <= float(ds[resolved.y_dim].max()) + ) + if _y_is_latlon: _lat_idx = xr.DataArray(resolved.cell_lat, dims="section") _lon_idx = xr.DataArray(resolved.cell_lon, dims="section") sel = {resolved.y_dim: _lat_idx, resolved.x_dim: _lon_idx} _sel_kw: dict = {"method": "nearest"} + _use_isel = False else: j_da = xr.DataArray(resolved.j, dims="section") i_da = xr.DataArray(resolved.i, dims="section") sel = {resolved.y_dim: j_da, resolved.x_dim: i_da} _sel_kw = {} + _use_isel = True + + def _select(da: xr.DataArray) -> xr.DataArray: + return da.isel(**sel) if _use_isel else da.sel(**sel, **_sel_kw) theta = np.radians(resolved.bearing_deg) cos_t = xr.DataArray(np.cos(theta), dims="section") sin_t = xr.DataArray(np.sin(theta), dims="section") w = xr.DataArray(resolved.cell_width_km * 1e3, dims="section") - u = ds[u_var].sel(**sel, **_sel_kw) - v = ds[v_var].sel(**sel, **_sel_kw) + u = _select(ds[u_var]) + v = _select(ds[v_var]) v_normal = u * cos_t - v * sin_t # positive = rightward z_vals = ds[z_dim].values.astype(float) @@ -239,7 +253,7 @@ def transport( if constraints: cmask: xr.DataArray | None = None for cvar, (op, threshold) in constraints.items(): - val = ds[cvar].sel(**sel, **_sel_kw) + val = _select(ds[cvar]) cond: xr.DataArray = { "lt": val < threshold, "le": val <= threshold, @@ -260,14 +274,14 @@ def _tp_integrate(da: xr.DataArray) -> xr.DataArray: _tp_integrate(v_normal) * 1e-6, "volume transport", "Sv" ) if compute_heat: - t = ds[t_var].sel(**sel, **_sel_kw) + t = _select(ds[t_var]) out_vars["heat"] = _attach_attrs( _tp_integrate((t - t_ref) * v_normal) * rho0 * cp * 1e-12, "heat transport", "TW", ) if compute_salt: - s = ds[s_var].sel(**sel, **_sel_kw) + s = _select(ds[s_var]) out_vars["salt"] = _attach_attrs( _tp_integrate(s * v_normal) * rho0 / 1000.0, "salt transport", "kg s-1" )