diff --git a/glymur/jp2k.py b/glymur/jp2k.py index 87028e1..ea8e955 100644 --- a/glymur/jp2k.py +++ b/glymur/jp2k.py @@ -763,10 +763,10 @@ class Jp2k(Jp2kBox): """ Slicing protocol. """ - if isinstance(index, slice) and ( - index.start == None and + if ((isinstance(index, slice) and + (index.start == None and index.stop == None and - index.step == None): + index.step == None)) or (index == ...)): # Case of jp2[:] = data, i.e. write the entire image. # # Should have a slice object where start = stop = step = None @@ -787,8 +787,8 @@ class Jp2k(Jp2kBox): area = (row, 0, row + 1, codestream.segment[1].xsiz) return self.read(area=area).squeeze() - if isinstance(pargs, slice): - # Case of jp2[:], i.e. retrieve the entire image. + if isinstance(pargs, slice) or pargs is ...: + # Case of jp2[:] or jp2[...], i.e. retrieve the entire image. # # Should have a slice object where start = stop = step = None return self.read() @@ -806,6 +806,20 @@ class Jp2k(Jp2kBox): elif len(pargs) == 3: return pixel[pargs[2]] + if isinstance(pargs, tuple) and any(x is ... for x in pargs): + + nrows = codestream.segment[1].ysiz + ncols = codestream.segment[1].xsiz + nbands = codestream.segment[1].Csiz + # Reformulate without the ellipsis. + if pargs[0] is ...: + newindex = (slice(0, nrows), slice(0, ncols), pargs[1]) + elif pargs[1] is ...: + # Assume we have something like (r,...) where r is a scalar. + newindex = pargs[0] + + return self.__getitem__(newindex) + # Assuming pargs is a tuple of slices from now on. rows = pargs[0] cols = pargs[1] diff --git a/glymur/test/test_jp2k.py b/glymur/test/test_jp2k.py index 2868e84..ef86744 100644 --- a/glymur/test/test_jp2k.py +++ b/glymur/test/test_jp2k.py @@ -71,6 +71,16 @@ class SliceProtocolBase(unittest.TestCase): @unittest.skipIf(os.name == "nt", "NamedTemporaryFile issue on windows") class TestSliceProtocolBaseWrite(SliceProtocolBase): + def test_write_ellipsis(self): + expected = self.j2k_data + + with tempfile.NamedTemporaryFile(suffix='.j2k') as tfile: + j = Jp2k(tfile.name, 'wb') + j[...] = self.j2k_data + actual = j.read() + + np.testing.assert_array_equal(actual, expected) + def test_basic_write(self): expected = self.j2k_data @@ -236,6 +246,21 @@ class TestSliceProtocolRead(SliceProtocolBase): expected = self.jp2.read(area=(0, 0, 202, 202), rlevel=1) np.testing.assert_array_equal(actual, expected) + def test_ellipsis_full_read(self): + actual = self.j2k[...] + expected = self.j2k_data + np.testing.assert_array_equal(actual, expected) + + def test_ellipsis_band_select(self): + actual = self.j2k[..., 0] + expected = self.j2k_data[..., 0] + np.testing.assert_array_equal(actual, expected) + + def test_ellipsis_row_select(self): + actual = self.j2k[0, ...] + expected = self.j2k_data[0, ...] + np.testing.assert_array_equal(actual, expected) + def test_slice_protocol_2d_reduce_resolution(self): d = self.j2k[:] self.assertEqual(d.shape, (800, 480, 3))