diff --git a/changelog/980.bugfix.1.rst b/changelog/980.bugfix.1.rst new file mode 100644 index 000000000..b62ba24a8 --- /dev/null +++ b/changelog/980.bugfix.1.rst @@ -0,0 +1 @@ +~ndcube.NDCube.axis_world_coords and ~ndcube.NDCube.axis_world_coords_values with ``wcs=cube.extra_coords`` no longer transpose extra coordinates that span more than one array axis, such as a 2-D ~astropy.coordinates.SkyCoord lookup table on array axes (0, 1), or a WCS-backed extra coordinate whose mapping reorders correlated pixel axes. A multi-dimensional lookup table is now matched to its array axes in the order given. diff --git a/changelog/980.bugfix.rst b/changelog/980.bugfix.rst new file mode 100644 index 000000000..b995377d4 --- /dev/null +++ b/changelog/980.bugfix.rst @@ -0,0 +1 @@ +Slicing an ~ndcube.NDCube whose extra_coords include a lookup table spanning more than one array axis now works. Integer indexing no longer gives wrong or NaN values, and tables whose axes were given as a list, including every such table loaded from ASDF, no longer raise a TypeError. diff --git a/ndcube/asdf/converters/tests/test_ndcube_converter.py b/ndcube/asdf/converters/tests/test_ndcube_converter.py index b8497c165..d5c297aeb 100644 --- a/ndcube/asdf/converters/tests/test_ndcube_converter.py +++ b/ndcube/asdf/converters/tests/test_ndcube_converter.py @@ -6,8 +6,11 @@ from packaging.version import Version import asdf +import astropy.units as u import astropy.wcs +from astropy.coordinates import SkyCoord +from ndcube import NDCube from ndcube.tests.helpers import assert_cubes_equal @@ -47,3 +50,14 @@ def test_serialization_sliced_ndcube(expected_cube, tmp_path): with asdf.open(file_path) as af: assert_cubes_equal(af["ndcube_gwcs"], sndc, rtol=1e-12) + + +def test_serialization_multi_axis_extra_coords_can_be_sliced(tmp_path): + cube = NDCube(np.zeros((3, 4)), wcs=astropy.wcs.WCS(naxis=2)) + sky = SkyCoord(np.arange(12).reshape(3, 4) * u.deg, np.ones((3, 4)) * u.deg) + cube.extra_coords.add(("lon", "lat"), (0, 1), sky, mesh=False) + file_path = tmp_path / "test.asdf" + with asdf.AsdfFile({"ndcube": cube}) as af: + af.write_to(file_path) + with asdf.open(file_path) as af: + assert af["ndcube"][1:].extra_coords.keys() == cube[1:].extra_coords.keys() diff --git a/ndcube/extra_coords/extra_coords.py b/ndcube/extra_coords/extra_coords.py index bcbf6724d..8985c1dc9 100644 --- a/ndcube/extra_coords/extra_coords.py +++ b/ndcube/extra_coords/extra_coords.py @@ -359,8 +359,8 @@ def _getitem_lookup_tables(self, item): item = list(item) + [slice(None)] * (ndims - len(item)) n_dropped_dims = np.cumsum([isinstance(i, Integral) for i in item]) for lut_axis, lut in self._lookup_tables: - lut_axes = (lut_axis,) if not isinstance(lut_axis, tuple) else lut_axis - new_lut_axes = tuple(ax - n_dropped_dims[ax] for ax in lut_axes) + lut_axes = (lut_axis,) if isinstance(lut_axis, Integral) else tuple(lut_axis) + new_lut_axes = tuple(ax - n_dropped_dims[ax] for ax in lut_axes if not isinstance(item[ax], Integral)) lut_slice = tuple(item[i] for i in lut_axes) if isinstance(lut_slice, tuple) and len(lut_slice) == 1: lut_slice = lut_slice[0] diff --git a/ndcube/extra_coords/tests/test_extra_coords.py b/ndcube/extra_coords/tests/test_extra_coords.py index 577e7e32c..edff34023 100644 --- a/ndcube/extra_coords/tests/test_extra_coords.py +++ b/ndcube/extra_coords/tests/test_extra_coords.py @@ -560,3 +560,14 @@ def test_length1_extra_coord(wave_lut): sec = ec[item] assert (sec.wcs.pixel_to_world(0) == wave_lut[item]).all() assert (sec.wcs.world_to_pixel(wave_lut[item])[0] == [0]).all() + + +@pytest.mark.parametrize("axes", [(0, 1), [0, 1], (1, 0)]) +@pytest.mark.parametrize("item", [np.s_[:], np.s_[1:], np.s_[1]]) +def test_slice_multi_axis_lookup_table(axes, item): + cube = NDCube(np.zeros((3, 4, 5)), wcs=WCS(naxis=3)) + lon = np.arange(12).reshape(3, 4) + sky = SkyCoord(lon * u.deg, np.ones((3, 4)) * u.deg) + cube.extra_coords.add(("lon", "lat"), axes, sky if axes[0] == 0 else sky.T, mesh=False) + sliced = cube[item] + np.testing.assert_allclose(sliced.axis_world_coords(wcs=sliced.extra_coords)[0].ra.deg, lon[item]) diff --git a/ndcube/ndcube.py b/ndcube/ndcube.py index 87ece2c04..9eb845425 100644 --- a/ndcube/ndcube.py +++ b/ndcube/ndcube.py @@ -515,8 +515,10 @@ def _generate_world_coords(self, pixel_corners, wcs, *, needed_axes, units=None) ranges = [np.arange(i) - 0.5 for i in pixel_shape] else: ranges = [np.arange(i) for i in pixel_shape] + pixel_axes = np.arange(len(ranges)) # Limit the pixel dimensions to the ones present in the ExtraCoords if isinstance(wcs, ExtraCoords): + pixel_axes = np.asarray(wcs.mapping) ranges = [ranges[i] for i in wcs.mapping] wcs = wcs.wcs if wcs is None: @@ -556,8 +558,8 @@ def _generate_world_coords(self, pixel_corners, wcs, *, needed_axes, units=None) for idx in world_axes_indices: array_slice = np.zeros((wcs.pixel_n_dim,), dtype=object) array_slice[wcs.axis_correlation_matrix[idx]] = slice(None) - tmp_world = world[idx][tuple(array_slice)].T - world_coords[idx] = tmp_world + order = np.argsort(pixel_axes[wcs.axis_correlation_matrix[idx]])[::-1] + world_coords[idx] = world[idx][tuple(array_slice)].transpose(order) if units: for i, (coord, unit) in enumerate(zip(world_coords, wcs.world_axis_units)): world_coords[i] = coord << u.Unit(unit)