Skip to content
Merged
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 docs/releases.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
58 changes: 58 additions & 0 deletions tests/data/_subset_tp5_stereo.py
Original file line number Diff line number Diff line change
@@ -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()
Binary file added tests/data/tp5_stereo_subset.nc
Binary file not shown.
99 changes: 99 additions & 0 deletions tests/test_regrid.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"):
Expand Down
79 changes: 59 additions & 20 deletions xhycom/_regrid.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,6 +30,7 @@
import hashlib
import json
import os
import warnings

import numpy as np
import xarray as xr
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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)
Expand Down
28 changes: 21 additions & 7 deletions xhycom/_transport.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand All @@ -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,
Expand All @@ -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"
)
Expand Down
Loading