diff --git a/geotiepoints/tests/test_viiinterpolator.py b/geotiepoints/tests/test_viiinterpolator.py index 37eae0c..aa12c91 100644 --- a/geotiepoints/tests/test_viiinterpolator.py +++ b/geotiepoints/tests/test_viiinterpolator.py @@ -6,12 +6,16 @@ """ -import unittest +import dask +import dask.array as da import numpy as np -import xarray as xr import pytest +import xarray as xr + from geotiepoints.viiinterpolator import tie_points_interpolation, tie_points_geo_interpolation +from .utils import CustomScheduler + TEST_N_SCANS = 2 TEST_TIE_POINTS_FACTOR = 2 @@ -19,6 +23,8 @@ TEST_VALID_ALT_TIE_POINTS = TEST_SCAN_ALT_TIE_POINTS * TEST_N_SCANS TEST_INVALID_ALT_TIE_POINTS = TEST_SCAN_ALT_TIE_POINTS * TEST_N_SCANS + 1 TEST_ACT_TIE_POINTS = 4 +TEST_SCAN_ALT_PIXELS = (TEST_SCAN_ALT_TIE_POINTS - 1) * TEST_TIE_POINTS_FACTOR +TEST_ACT_PIXELS = (TEST_ACT_TIE_POINTS - 1) * TEST_TIE_POINTS_FACTOR # Results of latitude/longitude interpolation with simple interpolation on coordinates TEST_LON_1 = np.array( @@ -42,7 +48,7 @@ [18., 18.5, 19., 19.5, 20., 20.5]] ) -# Results of latitude/longitude interpolation on cartesian coordinates (latitude above 60 degrees) +# Results of latitude/longitude interpolation on cartesian coordinates (longitude with a 360 degrees step) TEST_LON_2 = np.array( [[-12., -11.50003808, -11., -10.50011426, -10., -9.50019052], [-10.00243991, -9.5032411, -9.00366173, -8.50454031, -8.00488578, -7.5058423], @@ -65,7 +71,7 @@ [17.98865968, 18.48759253, 18.98798235, 19.48686863, 19.98729684, 20.48613573]] ) -# Results of latitude/longitude interpolation on cartesian coordinates (longitude with a 360 degrees step) +# Results of latitude/longitude interpolation on cartesian coordinates (latitude above 60 degrees) TEST_LON_3 = np.array( [[-12., -11.50444038, -11., -10.50459822, -10., -9.50476197], [-10.07492627, -9.58101155, -9.07759836, -8.5839056, - 8.0803761, -7.58691614], @@ -89,134 +95,105 @@ ) -class TestViiInterpolator(unittest.TestCase): - """Test the vii_utils module.""" +@pytest.fixture(params=[False, True], ids=["numpy", "dask"]) +def use_dask(request): + """Run a test with numpy-backed and with dask-backed tie points.""" + return request.param - def setUp(self): - """Set up the test.""" - # Create the arrays for the interpolation test - # The first has a valid number of n_tie_alt points (multiple of SCAN_ALT_TIE_POINTS) - self.valid_data_for_interpolation = xr.DataArray( - np.arange( - TEST_VALID_ALT_TIE_POINTS * TEST_ACT_TIE_POINTS, - dtype=np.float64, - ).reshape(TEST_VALID_ALT_TIE_POINTS, TEST_ACT_TIE_POINTS), - dims=('num_tie_points_alt', 'num_tie_points_act'), - ) - # The second has an invalid number of n_tie_alt points (not multiple of SCAN_ALT_TIE_POINTS) - self.invalid_data_for_interpolation = xr.DataArray( - np.arange( - TEST_INVALID_ALT_TIE_POINTS * TEST_ACT_TIE_POINTS, - dtype=np.float64, - ).reshape(TEST_INVALID_ALT_TIE_POINTS, TEST_ACT_TIE_POINTS), - dims=('num_tie_points_alt', 'num_tie_points_act'), - ) - # Then two arrays containing valid longitude and latitude data - self.longitude = xr.DataArray( - np.linspace( - -12, - 11, - num=TEST_VALID_ALT_TIE_POINTS * TEST_ACT_TIE_POINTS, - dtype=np.float64, - ).reshape(TEST_VALID_ALT_TIE_POINTS, TEST_ACT_TIE_POINTS), - dims=('num_tie_points_alt', 'num_tie_points_act'), - ) - self.latitude = xr.DataArray( - np.linspace( - 0, - 23, - num=TEST_VALID_ALT_TIE_POINTS * TEST_ACT_TIE_POINTS, - dtype=np.float64, - ).reshape(TEST_VALID_ALT_TIE_POINTS, TEST_ACT_TIE_POINTS), - dims=('num_tie_points_alt', 'num_tie_points_act'), - ) - # Then one containing latitude data above 60 degrees - self.latitude_over60 = xr.DataArray( - np.linspace( - 45, - 68, - num=TEST_VALID_ALT_TIE_POINTS * TEST_ACT_TIE_POINTS, - dtype=np.float64, - ).reshape(TEST_VALID_ALT_TIE_POINTS, TEST_ACT_TIE_POINTS), - dims=('num_tie_points_alt', 'num_tie_points_act'), - ) - # Then one containing longitude data with a 360 degrees step - self.longitude_over360 = xr.DataArray( - np.linspace( - -12, - 11, - num=TEST_VALID_ALT_TIE_POINTS * TEST_ACT_TIE_POINTS, - dtype=np.float64, - ).reshape(TEST_VALID_ALT_TIE_POINTS, TEST_ACT_TIE_POINTS) % 360., - dims=('num_tie_points_alt', 'num_tie_points_act'), - ) - def tearDown(self): - """Tear down the test.""" - # Nothing to do - pass +def _tie_points_data_array(data, use_dask): + """Wrap tie point values in a DataArray with the VII tie point dimensions. - def test_tie_points_interpolation(self): - """# Test the interpolation routine with valid and invalid input.""" - # Test the interpolation routine with valid input - result_valid = tie_points_interpolation( - [self.valid_data_for_interpolation], - TEST_SCAN_ALT_TIE_POINTS, - TEST_TIE_POINTS_FACTOR - )[0] - - act_points_interp = (TEST_ACT_TIE_POINTS - 1) * TEST_TIE_POINTS_FACTOR - num_scans = TEST_VALID_ALT_TIE_POINTS // TEST_SCAN_ALT_TIE_POINTS - scan_alt_points_interp = (TEST_SCAN_ALT_TIE_POINTS - 1) * TEST_TIE_POINTS_FACTOR - - # Across the track - delta_axis_0 = [0., 0.5, 1., 1.5, 2., 2.5] - self.assertTrue(np.allclose(result_valid[0, :], delta_axis_0)) - # Along track - delta_axis_1 = [0., 2., 4., 6., 12., 14., 16., 18] - self.assertTrue(np.allclose(result_valid[:, 0], delta_axis_1)) - - # Test the interpolation routine with invalid input - pytest.raises(ValueError, tie_points_interpolation, - [self.invalid_data_for_interpolation], - TEST_SCAN_ALT_TIE_POINTS, - TEST_TIE_POINTS_FACTOR) - - def test_tie_points_geo_interpolation(self): - """# Test the coordinates interpolation routine with valid and invalid input.""" - # Test the interpolation routine with valid input - lon, lat = tie_points_geo_interpolation( - self.longitude, - self.latitude, - TEST_SCAN_ALT_TIE_POINTS, - TEST_TIE_POINTS_FACTOR - ) - self.assertTrue(np.allclose(lon, TEST_LON_1)) - self.assertTrue(np.allclose(lat, TEST_LAT_1)) + Dask arrays are chunked along the track in whole scans like Satpy's METimage readers do. - lon, lat = tie_points_geo_interpolation( - self.longitude_over360, - self.latitude, - TEST_SCAN_ALT_TIE_POINTS, - TEST_TIE_POINTS_FACTOR - ) - self.assertTrue(np.allclose(lon, TEST_LON_2)) - self.assertTrue(np.allclose(lat, TEST_LAT_2)) + """ + if use_dask: + data = da.from_array(data, chunks=(TEST_SCAN_ALT_TIE_POINTS, -1)) + return xr.DataArray(data, dims=('num_tie_points_alt', 'num_tie_points_act')) + + +def _arange_tie_points(n_tie_alt, use_dask): + """Create tie points counting up from 0 with ``n_tie_alt`` points along the track.""" + data = np.arange(n_tie_alt * TEST_ACT_TIE_POINTS, dtype=np.float64).reshape(n_tie_alt, TEST_ACT_TIE_POINTS) + return _tie_points_data_array(data, use_dask) + + +def _linspace_tie_points(start, stop): + """Create evenly spaced tie point values from ``start`` to ``stop`` for a valid number of scans.""" + data = np.linspace(start, stop, num=TEST_VALID_ALT_TIE_POINTS * TEST_ACT_TIE_POINTS, dtype=np.float64) + return data.reshape(TEST_VALID_ALT_TIE_POINTS, TEST_ACT_TIE_POINTS) + + +def _assert_pixel_array(data_arr, use_dask): + """Check the array type, chunks, and memory layout of interpolated data. + + Satpy's METimage readers expect one chunk of pixel rows per chunk of tie point rows, each + spanning the whole width of the swath. Consumers like pyresample's EWA resampling require + C-contiguous arrays. Dask blocks are checked individually as computing the whole array would + concatenate them into a new C-contiguous array and hide the problem. + + """ + if use_dask: + assert isinstance(data_arr.data, da.Array) + assert data_arr.chunks == ((TEST_SCAN_ALT_PIXELS,) * TEST_N_SCANS, (TEST_ACT_PIXELS,)) + blocks = dask.compute(*data_arr.data.to_delayed().ravel()) + else: + assert isinstance(data_arr.data, np.ndarray) + blocks = [data_arr.data] + assert all(block.flags.c_contiguous for block in blocks) + +def test_tie_points_interpolation(use_dask): + """Test the interpolation routine with valid input.""" + data = _arange_tie_points(TEST_VALID_ALT_TIE_POINTS, use_dask) + with dask.config.set(scheduler=CustomScheduler(max_computes=0)): + result = tie_points_interpolation([data], TEST_SCAN_ALT_TIE_POINTS, TEST_TIE_POINTS_FACTOR)[0] + + _assert_pixel_array(result, use_dask) + # Across the track + np.testing.assert_allclose(result[0, :], [0., 0.5, 1., 1.5, 2., 2.5]) + # Along the track + np.testing.assert_allclose(result[:, 0], [0., 2., 4., 6., 12., 14., 16., 18.]) + + +def test_tie_points_interpolation_invalid_alt_tie_points(use_dask): + """Test that the number of tie points along the track must be a multiple of the tie points per scan.""" + data = _arange_tie_points(TEST_INVALID_ALT_TIE_POINTS, use_dask) + with pytest.raises(ValueError, match="must be a multiple"): + tie_points_interpolation([data], TEST_SCAN_ALT_TIE_POINTS, TEST_TIE_POINTS_FACTOR) + + +@pytest.mark.parametrize( + ("longitude", "latitude", "exp_lon", "exp_lat"), + [ + pytest.param(_linspace_tie_points(-12, 11), _linspace_tie_points(0, 23), TEST_LON_1, TEST_LAT_1, + id="lonlat"), + pytest.param(_linspace_tie_points(-12, 11) % 360., _linspace_tie_points(0, 23), TEST_LON_2, TEST_LAT_2, + id="cartesian_lon_360_step"), + pytest.param(_linspace_tie_points(-12, 11), _linspace_tie_points(45, 68), TEST_LON_3, TEST_LAT_3, + id="cartesian_lat_over_60"), + ], +) +def test_tie_points_geo_interpolation(longitude, latitude, exp_lon, exp_lat, use_dask): + """Test the coordinates interpolation routine in geodetic and cartesian coordinates.""" + # Choosing between geodetic and cartesian interpolation computes the latitude and longitude ranges + with dask.config.set(scheduler=CustomScheduler(max_computes=2)): lon, lat = tie_points_geo_interpolation( - self.longitude, - self.latitude_over60, + _tie_points_data_array(longitude, use_dask), + _tie_points_data_array(latitude, use_dask), TEST_SCAN_ALT_TIE_POINTS, TEST_TIE_POINTS_FACTOR ) - self.assertTrue(np.allclose(lon, TEST_LON_3)) - self.assertTrue(np.allclose(lat, TEST_LAT_3)) - - # Test the interpolation routine with invalid input (different dimensions of the two arrays) - with self.assertRaises(ValueError): - tie_points_geo_interpolation( - self.longitude, - self.invalid_data_for_interpolation, - TEST_SCAN_ALT_TIE_POINTS, - TEST_TIE_POINTS_FACTOR - ) + + _assert_pixel_array(lon, use_dask) + _assert_pixel_array(lat, use_dask) + np.testing.assert_allclose(lon, exp_lon) + np.testing.assert_allclose(lat, exp_lat) + + +def test_tie_points_geo_interpolation_mismatched_shapes(use_dask): + """Test that longitude and latitude must have the same shape.""" + longitude = _tie_points_data_array(_linspace_tie_points(-12, 11), use_dask) + latitude = _arange_tie_points(TEST_INVALID_ALT_TIE_POINTS, use_dask) + with pytest.raises(ValueError, match="don't match"): + tie_points_geo_interpolation(longitude, latitude, TEST_SCAN_ALT_TIE_POINTS, TEST_TIE_POINTS_FACTOR) diff --git a/geotiepoints/viiinterpolator.py b/geotiepoints/viiinterpolator.py index 6b4add8..57267d8 100644 --- a/geotiepoints/viiinterpolator.py +++ b/geotiepoints/viiinterpolator.py @@ -68,18 +68,21 @@ def tie_points_interpolation(data_on_tie_points, scan_alt_tie_points, tie_points if data.shape != (n_tie_alt, n_tie_act) or data.dims != (dim_alt, dim_act): raise ValueError("The dimensions of the arrays are not consistent") - # Interpolate using the xarray interp function twice: first across, then along the scan + # Interpolate using the xarray interp function twice: first along, then across the scan # (much faster than interpolating directly in the two dimensions) + # Interpolating across the scan last keeps dask chunks of whole pixel rows, one per chunk of + # tie point rows, but it produces Fortran-ordered arrays. Make them C-contiguous as required + # by consumers like pyresample's EWA resampling. data = data.assign_coords({dim_alt: tie_grid_alt, dim_act: tie_grid_act}) data_pixel = data.interp({dim_alt: pixel_grid_alt}, assume_sorted=True) \ .interp({dim_act: pixel_grid_act}, assume_sorted=True).drop_vars([dim_alt, dim_act]) + data_pixel = _as_c_contiguous(data_pixel) data_on_pixel_points.append(data_pixel) return data_on_pixel_points - def tie_points_geo_interpolation(longitude, latitude, scan_alt_tie_points, tie_points_factor, lat_threshold_use_cartesian=60., @@ -176,3 +179,21 @@ def _xyz2lonlat(x_coords, y_coords, z_coords, z_threshold_use_xy=0.8): np.sign(z_coords) * (90. - np.rad2deg(np.arcsin(r / MEAN_EARTH_RADIUS))) ) return lons, lats + + +def _as_c_contiguous(data_arr): + """Get a copy of the DataArray with a C-contiguous numpy array or C-contiguous dask blocks. + + Args: + data_arr: xarray DataArray backed by a numpy or dask array. + + Returns: + xarray DataArray with the same dimensions, coordinates and attributes. + + """ + data = data_arr.data + if isinstance(data, da.Array): + data = data.map_blocks(np.ascontiguousarray, dtype=data.dtype, meta=np.array((), dtype=data.dtype)) + else: + data = np.ascontiguousarray(data) + return data_arr.copy(data=data)