diff --git a/opendm/orthophoto.py b/opendm/orthophoto.py index 27f42b81a..679d6c432 100644 --- a/opendm/orthophoto.py +++ b/opendm/orthophoto.py @@ -265,6 +265,28 @@ def feather_raster(input_raster, output_raster, blend_distance=20): return output_raster +def _read_window_gated(ds, src_window, dst_shape, dtype): + """Read ``src_window`` into a ``dst_shape`` array, equivalent to a + ``boundless=True`` read with 0 nodata fill, but avoiding rasterio's boundless + VRT path for the common cases. ``boundless=True`` builds a VRT and serializes + it via Python's ElementTree (``_serialize_xml``) on every read, which is a + large per-read overhead on big merges. Behaviour: + - window fully outside the dataset -> all zeros (no read), matching the 0 + nodata fill a boundless read would produce here; + - window fully inside -> plain non-boundless read (no VRT), identical to a + boundless read when no out-of-bounds padding is needed; + - window partially overlapping the edge -> fall back to boundless (rare; + only the true border blocks), preserving exact fill behaviour. + """ + (r0, r1), (c0, c1) = src_window + out = np.zeros(dst_shape, dtype=dtype) + height, width = ds.height, ds.width + if r1 <= 0 or c1 <= 0 or r0 >= height or c0 >= width: + return out + if r0 >= 0 and c0 >= 0 and r1 <= height and c1 <= width: + return ds.read(out=out, window=src_window, boundless=False, masked=False) + return ds.read(out=out, window=src_window, boundless=True, masked=False) + def merge(input_ortho_and_ortho_cuts, output_orthophoto, orthophoto_vars={}, merge_skip_blending=False): """ Based on https://github.com/mapbox/rio-merge-rgba/ @@ -361,10 +383,7 @@ def merge(input_ortho_and_ortho_cuts, output_orthophoto, orthophoto_vars={}, mer src.transform, right, bottom, op=round, precision=precision ))) - temp = np.zeros(dst_shape, dtype=dtype) - temp = src.read( - out=temp, window=src_window, boundless=True, masked=False - ) + temp = _read_window_gated(src, src_window, dst_shape, dtype) # pixels without data yet are available to write write_region = np.logical_and( @@ -390,10 +409,7 @@ def merge(input_ortho_and_ortho_cuts, output_orthophoto, orthophoto_vars={}, mer src.transform, right, bottom, op=round, precision=precision ))) - temp = np.zeros(dst_shape, dtype=dtype) - temp = src.read( - out=temp, window=src_window, boundless=True, masked=False - ) + temp = _read_window_gated(src, src_window, dst_shape, dtype) where = temp[-1] != 0 for b in range(0, num_bands): @@ -414,10 +430,7 @@ def merge(input_ortho_and_ortho_cuts, output_orthophoto, orthophoto_vars={}, mer cut.transform, right, bottom, op=round, precision=precision ))) - temp = np.zeros(dst_shape, dtype=dtype) - temp = cut.read( - out=temp, window=src_window, boundless=True, masked=False - ) + temp = _read_window_gated(cut, src_window, dst_shape, dtype) # For each band, average alpha values between # destination raster and cut raster diff --git a/tests/test_orthophoto.py b/tests/test_orthophoto.py new file mode 100644 index 000000000..2792f7403 --- /dev/null +++ b/tests/test_orthophoto.py @@ -0,0 +1,103 @@ +import os +import shutil +import unittest + +import numpy as np +import rasterio +from rasterio.transform import Affine + +from opendm.orthophoto import _read_window_gated + + +class TestReadWindowGated(unittest.TestCase): + """_read_window_gated must be pixel-identical to a plain boundless read. + + It is an optimization that avoids rasterio's boundless VRT path for the + common in-bounds / fully-out-of-bounds cases, so for every window position + its result must equal what ds.read(..., boundless=True) would produce. + """ + + WIDTH = 10 + HEIGHT = 8 + COUNT = 4 # 3 bands + alpha, like the orthophotos merge() operates on + + def setUp(self): + self.tmp = "tests/assets/output" + if os.path.exists(self.tmp): + shutil.rmtree(self.tmp) + os.makedirs(self.tmp) + self.path = os.path.join(self.tmp, "gated_src.tif") + + # A distinct value per (band, row, col) so a wrong window/offset is visible. + self.data = np.arange( + self.COUNT * self.HEIGHT * self.WIDTH, dtype=np.uint8 + ).reshape(self.COUNT, self.HEIGHT, self.WIDTH) + + profile = { + "driver": "GTiff", + "width": self.WIDTH, + "height": self.HEIGHT, + "count": self.COUNT, + "dtype": "uint8", + "nodata": 0, + "transform": Affine.translation(100.0, 200.0) * Affine.scale(0.5, -0.5), + } + with rasterio.open(self.path, "w", **profile) as dst: + dst.write(self.data) + + def tearDown(self): + shutil.rmtree(self.tmp) + + def _boundless_reference(self, ds, window): + (r0, r1), (c0, c1) = window + shape = (self.COUNT, r1 - r0, c1 - c0) + out = np.zeros(shape, dtype=np.uint8) + return ds.read(out=out, window=window, boundless=True, masked=False) + + def _assert_matches_boundless(self, window): + (r0, r1), (c0, c1) = window + dst_shape = (self.COUNT, r1 - r0, c1 - c0) + with rasterio.open(self.path) as ds: + gated = _read_window_gated(ds, window, dst_shape, np.uint8) + reference = self._boundless_reference(ds, window) + self.assertEqual(gated.shape, dst_shape) + self.assertEqual(gated.dtype, np.uint8) + np.testing.assert_array_equal(gated, reference) + return gated + + def test_fully_inside(self): + gated = self._assert_matches_boundless(((2, 5), (3, 7))) + # Sanity: the fixture is non-trivial, so a real read is not all zeros. + self.assertTrue(gated.any()) + + def test_fully_inside_full_extent(self): + self._assert_matches_boundless(((0, self.HEIGHT), (0, self.WIDTH))) + + def test_fully_outside_left(self): + gated = self._assert_matches_boundless(((0, 3), (-5, -2))) + self.assertFalse(gated.any()) # nothing overlaps -> all zeros + + def test_fully_outside_right(self): + self._assert_matches_boundless(((0, 3), (self.WIDTH + 2, self.WIDTH + 5))) + + def test_fully_outside_top(self): + self._assert_matches_boundless(((-4, -1), (0, 3))) + + def test_fully_outside_bottom(self): + self._assert_matches_boundless(((self.HEIGHT + 1, self.HEIGHT + 4), (0, 3))) + + def test_partial_top_left_corner(self): + self._assert_matches_boundless(((-2, 3), (-2, 4))) + + def test_partial_bottom_right_corner(self): + self._assert_matches_boundless( + ((self.HEIGHT - 2, self.HEIGHT + 3), (self.WIDTH - 3, self.WIDTH + 2)) + ) + + def test_partial_straddles_full_width(self): + # In-bounds vertically, straddling both left and right edges at once. + self._assert_matches_boundless(((1, 4), (-2, self.WIDTH + 2))) + + +if __name__ == "__main__": + unittest.main()