Skip to content

Commit 83e329d

Browse files
authored
Merge pull request #68 from fish-pace/copilot/check-spatial-extent-xarray-dataset
Pre-slice grid dataset to minimal spatial extent before xoak k-d tree indexing
2 parents 5a751d0 + 3701014 commit 83e329d

2 files changed

Lines changed: 262 additions & 0 deletions

File tree

‎src/point_collocation/core/engine.py‎

Lines changed: 89 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -527,6 +527,15 @@ def _execute_plan(
527527
"Use plan.show_variables() to inspect the dataset."
528528
)
529529

530+
# For grid+xoak, pre-slice the dataset to the spatial extent
531+
# of the query points before building the k-d tree. A global
532+
# granule with only a few scattered points would otherwise cause
533+
# xoak to index the entire global grid, which is very slow.
534+
if spatial_method == "xoak" and geometry == "grid":
535+
pt_lats = [float(plan.points.loc[idx]["lat"]) for idx in pt_indices]
536+
pt_lons = [float(plan.points.loc[idx]["lon"]) for idx in pt_indices]
537+
ds = _slice_grid_to_points(ds, pt_lats, pt_lons, lat_name, lon_name)
538+
530539
for pt_idx in pt_indices:
531540
row = plan.points.loc[pt_idx].to_dict()
532541
row["granule_id"] = gm.granule_id
@@ -597,6 +606,86 @@ def _execute_plan(
597606
return df
598607

599608

609+
def _slice_grid_to_points(
610+
ds: xr.Dataset,
611+
lats: list[float],
612+
lons: list[float],
613+
lat_name: str,
614+
lon_name: str,
615+
buffer_deg: float = 1.0,
616+
) -> xr.Dataset:
617+
"""Slice a regular-grid dataset to the smallest region covering *lats*/*lons*.
618+
619+
When ``geometry='grid'`` and ``spatial_method='xoak'``, building a k-d tree
620+
over an entire global granule is very slow if only a few points need to be
621+
matched. This function slices the dataset to a padded bounding box around
622+
the query points so xoak indexes the minimum required region.
623+
624+
Only applies to datasets with 1-D coordinate arrays (regular grids). Returns
625+
*ds* unchanged for 2-D coordinates or if the resulting slice would be empty.
626+
627+
Parameters
628+
----------
629+
ds:
630+
The dataset to slice.
631+
lats, lons:
632+
Latitudes and longitudes of the query points.
633+
lat_name, lon_name:
634+
Coordinate names detected by :func:`_find_geoloc_pair`.
635+
buffer_deg:
636+
Extra degrees to pad the bounding box on each side (default 1°).
637+
Ensures at least one grid cell surrounds each query point.
638+
639+
Returns
640+
-------
641+
xr.Dataset
642+
A lazy slice of *ds* covering the padded bounding box, or *ds* unchanged
643+
if the coordinates are not 1-D or the slice would be empty.
644+
"""
645+
lat_coord = ds.coords.get(lat_name) if lat_name in ds.coords else ds.get(lat_name)
646+
lon_coord = ds.coords.get(lon_name) if lon_name in ds.coords else ds.get(lon_name)
647+
648+
if lat_coord is None or lon_coord is None:
649+
return ds
650+
if lat_coord.ndim != 1 or lon_coord.ndim != 1:
651+
return ds
652+
653+
lat_min_data = float(lat_coord.min())
654+
lat_max_data = float(lat_coord.max())
655+
lon_min_data = float(lon_coord.min())
656+
lon_max_data = float(lon_coord.max())
657+
658+
min_lat = max(min(lats) - buffer_deg, lat_min_data)
659+
max_lat = min(max(lats) + buffer_deg, lat_max_data)
660+
min_lon = max(min(lons) - buffer_deg, lon_min_data)
661+
max_lon = min(max(lons) + buffer_deg, lon_max_data)
662+
663+
# xarray slice() is order-aware: if the coordinate is stored in descending
664+
# order (e.g. 90→-90), the larger bound must come first.
665+
lat_vals = lat_coord.values
666+
if len(lat_vals) > 1 and lat_vals[0] > lat_vals[-1]:
667+
lat_slice = slice(max_lat, min_lat)
668+
else:
669+
lat_slice = slice(min_lat, max_lat)
670+
671+
lon_vals = lon_coord.values
672+
if len(lon_vals) > 1 and lon_vals[0] > lon_vals[-1]:
673+
lon_slice = slice(max_lon, min_lon)
674+
else:
675+
lon_slice = slice(min_lon, max_lon)
676+
677+
sliced = ds.sel({lat_name: lat_slice, lon_name: lon_slice})
678+
679+
# Guard against an empty slice (e.g., all query points fall outside the
680+
# coordinate range, or the grid is coarser than the buffer).
681+
lat_dim = lat_coord.dims[0]
682+
lon_dim = lon_coord.dims[0]
683+
if sliced.sizes.get(lat_dim, 0) == 0 or sliced.sizes.get(lon_dim, 0) == 0:
684+
return ds
685+
686+
return sliced
687+
688+
600689
def _extract_nearest(
601690
ds: xr.Dataset,
602691
row: dict,

‎tests/test_plan.py‎

Lines changed: 173 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -3396,6 +3396,179 @@ def test_grid_matchup_with_xoak_returns_nearest_value(
33963396
assert len(result) == 1
33973397
assert not math.isnan(result.loc[0, "sst"])
33983398

3399+
def test_grid_matchup_xoak_global_granule_returns_nearest_value(
3400+
self, tmp_path: pathlib.Path, monkeypatch: pytest.MonkeyPatch
3401+
) -> None:
3402+
"""geometry='grid' + xoak on a global granule slices correctly and returns a value."""
3403+
pytest.importorskip("xoak") # skip if xoak not installed
3404+
3405+
# Large global grid (181 lats × 361 lons) with the query point near centre.
3406+
lats = list(range(-90, 91)) # integers -90, -89, …, 90
3407+
lons = list(range(-180, 181)) # integers -180, -179, …, 180
3408+
nc_path = str(tmp_path / "global_grid.nc")
3409+
_make_l3_dataset(lats, lons, seed=99).to_netcdf(nc_path, engine="netcdf4")
3410+
3411+
mock_ea = MagicMock()
3412+
mock_ea.open.return_value = [nc_path]
3413+
monkeypatch.setitem(__import__("sys").modules, "earthaccess", mock_ea)
3414+
3415+
# Query a single point at (lat=10, lon=20).
3416+
pts = pd.DataFrame(
3417+
{
3418+
"lat": [10.0],
3419+
"lon": [20.0],
3420+
"time": pd.to_datetime(["2023-06-01T12:00:00"]),
3421+
}
3422+
)
3423+
gm = GranuleMeta(
3424+
granule_id="https://example.com/global_grid.nc",
3425+
begin=pd.Timestamp("2023-06-01T00:00:00Z"),
3426+
end=pd.Timestamp("2023-06-01T23:59:59Z"),
3427+
bbox=(-180.0, -90.0, 180.0, 90.0),
3428+
result_index=0,
3429+
)
3430+
p = Plan(
3431+
points=pts,
3432+
results=[object()],
3433+
granules=[gm],
3434+
point_granule_map={0: [0]},
3435+
source_kwargs={"short_name": "TEST"},
3436+
time_buffer=pd.Timedelta(0),
3437+
)
3438+
3439+
result = pc.matchup(
3440+
p,
3441+
geometry="grid",
3442+
variables=["sst"],
3443+
spatial_method="xoak",
3444+
open_dataset_kwargs={"engine": "netcdf4"},
3445+
)
3446+
3447+
assert "sst" in result.columns
3448+
assert len(result) == 1
3449+
assert not math.isnan(result.loc[0, "sst"])
3450+
3451+
3452+
class TestSliceGridToPoints:
3453+
"""Unit tests for the _slice_grid_to_points helper."""
3454+
3455+
def test_slices_ascending_coords(self) -> None:
3456+
"""Dataset with ascending lat/lon is sliced to the point bounding box + buffer."""
3457+
from point_collocation.core.engine import _slice_grid_to_points
3458+
3459+
lats = list(range(-90, 91))
3460+
lons = list(range(-180, 181))
3461+
ds = xr.Dataset(
3462+
{"sst": (["lat", "lon"], np.zeros((len(lats), len(lons))))},
3463+
coords={"lat": lats, "lon": lons},
3464+
)
3465+
3466+
sliced = _slice_grid_to_points(ds, [10.0], [20.0], "lat", "lon", buffer_deg=2.0)
3467+
3468+
# The slice should cover [8, 12] lat and [18, 22] lon (within 2° buffer).
3469+
assert float(sliced["lat"].min()) >= 8.0
3470+
assert float(sliced["lat"].max()) <= 12.0
3471+
assert float(sliced["lon"].min()) >= 18.0
3472+
assert float(sliced["lon"].max()) <= 22.0
3473+
# Original dataset should be much larger.
3474+
assert sliced.sizes["lat"] < ds.sizes["lat"]
3475+
assert sliced.sizes["lon"] < ds.sizes["lon"]
3476+
3477+
def test_slices_descending_lat_coords(self) -> None:
3478+
"""Dataset with descending lat (90→-90) is sliced correctly."""
3479+
from point_collocation.core.engine import _slice_grid_to_points
3480+
3481+
lats = list(range(90, -91, -1)) # integers 90, 89, …, -90 (descending)
3482+
lons = list(range(-180, 181))
3483+
ds = xr.Dataset(
3484+
{"sst": (["lat", "lon"], np.zeros((len(lats), len(lons))))},
3485+
coords={"lat": lats, "lon": lons},
3486+
)
3487+
3488+
sliced = _slice_grid_to_points(ds, [5.0], [0.0], "lat", "lon", buffer_deg=1.0)
3489+
3490+
assert sliced.sizes["lat"] > 0
3491+
assert sliced.sizes["lon"] > 0
3492+
assert sliced.sizes["lat"] < ds.sizes["lat"]
3493+
3494+
def test_single_point_uses_buffer(self) -> None:
3495+
"""A single query point still produces a non-empty slice thanks to the buffer."""
3496+
from point_collocation.core.engine import _slice_grid_to_points
3497+
3498+
lats = list(range(-90, 91))
3499+
lons = list(range(-180, 181))
3500+
ds = xr.Dataset(
3501+
{"sst": (["lat", "lon"], np.zeros((len(lats), len(lons))))},
3502+
coords={"lat": lats, "lon": lons},
3503+
)
3504+
3505+
sliced = _slice_grid_to_points(ds, [0.0], [0.0], "lat", "lon", buffer_deg=1.0)
3506+
3507+
# 1° buffer each side → at least 3 lat values and 3 lon values.
3508+
assert sliced.sizes["lat"] >= 3
3509+
assert sliced.sizes["lon"] >= 3
3510+
3511+
def test_empty_slice_falls_back_to_full_dataset(self) -> None:
3512+
"""If the buffered box is outside the grid, the full dataset is returned."""
3513+
from point_collocation.core.engine import _slice_grid_to_points
3514+
3515+
lats = [0.0, 1.0, 2.0]
3516+
lons = [0.0, 1.0, 2.0]
3517+
ds = xr.Dataset(
3518+
{"sst": (["lat", "lon"], np.zeros((3, 3)))},
3519+
coords={"lat": lats, "lon": lons},
3520+
)
3521+
3522+
# Query point far outside the dataset range.
3523+
sliced = _slice_grid_to_points(ds, [50.0], [50.0], "lat", "lon", buffer_deg=0.5)
3524+
3525+
# Should fall back to the full dataset unchanged.
3526+
assert sliced.sizes["lat"] == ds.sizes["lat"]
3527+
assert sliced.sizes["lon"] == ds.sizes["lon"]
3528+
3529+
def test_2d_coords_returns_unchanged(self) -> None:
3530+
"""2-D (swath-style) coordinates are not sliced."""
3531+
from point_collocation.core.engine import _slice_grid_to_points
3532+
3533+
lat_2d = np.array([[0.0, 1.0], [2.0, 3.0]])
3534+
lon_2d = np.array([[10.0, 11.0], [12.0, 13.0]])
3535+
ds = xr.Dataset(
3536+
{"sst": (["nrows", "ncols"], np.zeros((2, 2)))},
3537+
coords={
3538+
"lat": (["nrows", "ncols"], lat_2d),
3539+
"lon": (["nrows", "ncols"], lon_2d),
3540+
},
3541+
)
3542+
3543+
sliced = _slice_grid_to_points(ds, [1.0], [11.0], "lat", "lon")
3544+
3545+
# 2-D coords → no slicing; sizes must be unchanged.
3546+
assert sliced.sizes == ds.sizes
3547+
3548+
def test_multiple_points_uses_union_bbox(self) -> None:
3549+
"""Multiple query points: slice covers the union bounding box."""
3550+
from point_collocation.core.engine import _slice_grid_to_points
3551+
3552+
lats = list(range(-90, 91))
3553+
lons = list(range(-180, 181))
3554+
ds = xr.Dataset(
3555+
{"sst": (["lat", "lon"], np.zeros((len(lats), len(lons))))},
3556+
coords={"lat": lats, "lon": lons},
3557+
)
3558+
3559+
# Two points that are far apart; the slice must cover both.
3560+
sliced = _slice_grid_to_points(
3561+
ds, [-30.0, 30.0], [-60.0, 60.0], "lat", "lon", buffer_deg=1.0
3562+
)
3563+
3564+
assert float(sliced["lat"].min()) <= -30.0
3565+
assert float(sliced["lat"].max()) >= 30.0
3566+
assert float(sliced["lon"].min()) <= -60.0
3567+
assert float(sliced["lon"].max()) >= 60.0
3568+
# Still smaller than the full global grid.
3569+
assert sliced.sizes["lat"] < ds.sizes["lat"]
3570+
assert sliced.sizes["lon"] < ds.sizes["lon"]
3571+
33993572

34003573
class TestShowVariablesLayout:
34013574
"""Tests for plan.show_variables(geometry=...) with both open methods."""

0 commit comments

Comments
 (0)