Skip to content
Open
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
9 changes: 4 additions & 5 deletions src/nuthatch/backends/terracotta.py
Original file line number Diff line number Diff line change
Expand Up @@ -92,7 +92,6 @@ def lon_base_change(ds, to_base="base180", lon_dim='lon'):
ds = ds.sortby('lon')
return ds


@register_backend
class TerracottaBackend(DatabaseBackend, FileBackend):
"""
Expand Down Expand Up @@ -188,10 +187,11 @@ def write(self, ds, upsert=False, primary_keys=None):
else:
ds = ds.transpose('y', 'x')

# Adapt the CRS
# Label the dataset as WGS84.
# TODO: If the dataset already has a different CRS, we should reproject it to WGS84.
ds.rio.write_crs("epsg:4326", inplace=True)
ds = ds.rio.reproject('EPSG:3857', resampling=self.resampling, nodata=np.nan)
ds.rio.write_crs("epsg:3857", inplace=True)
# Write rows north-to-south, as expected by terracotta
ds = ds.sortby("y", ascending=False)

# Insert the parameters.
with self.driver.connect(verify=False):
Expand Down Expand Up @@ -268,4 +268,3 @@ def delete(self):
for dataset in datasets:
logger.info(f"Deleting datasets {datasets} from terracotta.")
self.driver.delete({'key': dataset})

157 changes: 157 additions & 0 deletions tests/backends/test_terracotta_alignment.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,157 @@
import os
from pathlib import Path

import fsspec
import numpy as np
import pytest
import rasterio
import xarray as xr
from rasterio.transform import rowcol, xy
from shapely.geometry import box, mapping
from terracotta import get_driver
from rasterio.enums import Resampling

from nuthatch.backends.terracotta import TerracottaBackend

# Test grids are regular north-up reference grids.
GRID_SPECS = {
"global1_5": {"grid_size": 1.5},
"global0_25": {"grid_size": 0.25},
"global0_1_avoidpoles": {"grid_size": 0.1, "lat_offset": 5.0},
}
def _reference_grid_coordinates(grid_name):
grid_spec = GRID_SPECS[grid_name]
grid_size = grid_spec["grid_size"]
lat_offset = grid_spec.get("lat_offset", 0.0) + grid_size/2.0
lons = np.arange(-180.0, 180.0, grid_size)
lats = np.arange(-90.0 + lat_offset, 90.0 - lat_offset, grid_size)
return lons, lats


def _reference_grid_dataset(grid_name):
lons, lats = _reference_grid_coordinates(grid_name)
values = np.arange(len(lats) * len(lons), dtype=float).reshape(len(lats), len(lons))
return xr.Dataset(
{"precip": (("lat", "lon"), values)},
coords={"lat": lats, "lon": lons},
)

AFRICA_BOUNDS = (-19.5, -40.5, 55.5, 40.5)
def _clip_source_to_pseudo_africa(ds):
return (
ds.rio.write_crs("EPSG:4326")
.rio.set_spatial_dims("lon", "lat")
.rio.clip(
[mapping(box(*AFRICA_BOUNDS))],
"EPSG:4326",
drop=True,
)
)


def _sqlite_database_path(path: Path) -> str:
return os.path.relpath(path, Path.cwd())


def _make_test_backend(tmp_path: Path, cache_key: str):
backend = TerracottaBackend.__new__(TerracottaBackend)
backend.lat_dim = "lat"
backend.lon_dim = "lon"
backend.time_dim = "time"
backend.resampling = Resampling.nearest
backend.cache_key = cache_key
backend.path = str(tmp_path / f"{cache_key}.terracotta")
backend.override_path = backend.path
backend.fs = fsspec.filesystem("file")
Path(backend.path).mkdir(parents=True, exist_ok=True)
backend.driver = get_driver(
_sqlite_database_path(tmp_path / "terracotta_scope.sqlite"),
provider="sqlite",
)
try:
backend.driver.get_keys()
except Exception:
backend.driver.create(["key"])
return backend


def _reference_points_from_dataset(ds):
lon_indices = [ds.sizes["lon"] // 6, ds.sizes["lon"] // 2, (5 * ds.sizes["lon"]) // 6]
lat_indices = [ds.sizes["lat"] // 6, ds.sizes["lat"] // 2, (5 * ds.sizes["lat"]) // 6]
return [
(float(ds.lon.values[lon_idx]), float(ds.lat.values[lat_idx]))
for lon_idx, lat_idx in zip(lon_indices, lat_indices, strict=False)
]


@pytest.mark.parametrize("grid_name", GRID_SPECS)
def test_terracotta_write_preserves_latlon_resolution_for_crops(tmp_path, grid_name):
source = _reference_grid_dataset(grid_name)
clipped = _clip_source_to_pseudo_africa(source)

global_backend = _make_test_backend(tmp_path, "global")
africa_backend = _make_test_backend(tmp_path, "africa")
global_backend.write(source)
africa_backend.write(clipped)

global_path = tmp_path / "global.terracotta" / "_.tif"
africa_path = tmp_path / "africa.terracotta" / "_.tif"

with (
rasterio.open(global_path) as global_tif,
rasterio.open(africa_path) as africa_tif,
):
assert global_tif.crs == africa_tif.crs
assert global_tif.res == pytest.approx(africa_tif.res)

@pytest.mark.parametrize("grid_name", GRID_SPECS)
def test_terracotta_returns_same_tile_values_for_global_and_africa_scopes(tmp_path, grid_name):

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I think an easier way to test this would be just to verify that the reprojected xarrays are the same and their x/y coordinates are subsets of one another

source = _reference_grid_dataset(grid_name)
clipped = _clip_source_to_pseudo_africa(source)

global_backend = _make_test_backend(tmp_path, "global")
africa_backend = _make_test_backend(tmp_path, "africa")
global_backend.write(source)
africa_backend.write(clipped)
driver = global_backend.driver

tile_bounds = driver.get_metadata(("africa",))["bounds"]
global_tile = driver.get_raster_tile(
{"key": "global"},
tile_bounds=tile_bounds,
tile_size=(256, 256),
preserve_values=True,
)
africa_tile = driver.get_raster_tile(
{"key": "africa"},
tile_bounds=tile_bounds,
tile_size=(256, 256),
preserve_values=True,
)

assert np.array_equal(global_tile.filled(-9999), africa_tile.filled(-9999))


@pytest.mark.parametrize("grid_name", GRID_SPECS)
@pytest.mark.parametrize("scope", ["global", "africa"])
def test_terracotta_write_preserves_reference_points_in_native_crs(
tmp_path, grid_name, scope
):
source = _reference_grid_dataset(grid_name)
ds_to_write = source if scope == "global" else _clip_source_to_pseudo_africa(source)
backend = _make_test_backend(tmp_path, scope)
backend.write(ds_to_write)
reference_points = _reference_points_from_dataset(ds_to_write)
tif_path = tmp_path / f"{scope}.terracotta" / "_.tif"

with rasterio.open(tif_path) as tif:
assert tif.crs.to_epsg() == 4326

for lon, lat in reference_points:
row, col = rowcol(tif.transform, lon, lat)
assert 0 <= row < tif.height
assert 0 <= col < tif.width

observed_x, observed_y = xy(tif.transform, row, col, offset="center")
assert observed_x == pytest.approx(lon, abs=tif.res[0] / 2.0)
assert observed_y == pytest.approx(lat, abs=abs(tif.res[1]) / 2.0)
Loading