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
1 change: 1 addition & 0 deletions changelog/980.bugfix.1.rst
Original file line number Diff line number Diff line change
@@ -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.
1 change: 1 addition & 0 deletions changelog/980.bugfix.rst
Original file line number Diff line number Diff line change
@@ -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.
14 changes: 14 additions & 0 deletions ndcube/asdf/converters/tests/test_ndcube_converter.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand Down Expand Up @@ -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()
4 changes: 2 additions & 2 deletions ndcube/extra_coords/extra_coords.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]
Expand Down
11 changes: 11 additions & 0 deletions ndcube/extra_coords/tests/test_extra_coords.py
Original file line number Diff line number Diff line change
Expand Up @@ -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])
6 changes: 4 additions & 2 deletions ndcube/ndcube.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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)
Expand Down
Loading