From 720a2c0899130a6fdcd56c6f7763a7ebf6638f93 Mon Sep 17 00:00:00 2001 From: Sylvester Kaczmarek <16242628+sylvesterkaczmarek@users.noreply.github.com> Date: Thu, 10 Sep 2026 17:56:16 +0100 Subject: [PATCH] Handle bounding boxes in 2D ndimage warping --- tests/warp_test.py | 62 ++++++++++++++++++++++++++++++++++++++++++++++ warp.py | 8 +++--- 2 files changed, 67 insertions(+), 3 deletions(-) diff --git a/tests/warp_test.py b/tests/warp_test.py index 0042cb4..ca6929c 100644 --- a/tests/warp_test.py +++ b/tests/warp_test.py @@ -16,6 +16,7 @@ """Tests for warp.""" from absl.testing import absltest +from absl.testing import parameterized from connectomics.common import bounding_box import numpy as np @@ -126,5 +127,66 @@ def test_warp_points(self): np.testing.assert_array_equal(warped, expected) +class NdimageWarpBoxesTest(parameterized.TestCase): + + @parameterized.product(dim=[2, 3], parallelism=[1, 2], scaled=[False, True]) + def test_cropped_output_with_boxes(self, dim, parallelism, scaled): + image_shape = (8, 12, 16)[-dim:] + image = np.zeros(image_shape, dtype=np.float32) + weights = np.array([1, 100, 10000])[:dim] + for weight, coord in zip(weights, np.indices(image_shape)[::-1]): + image += weight * coord + + scale = np.array([2., 0.5, 1.] if scaled else [1., 1., 1.])[:dim] + stride_xyz = np.array([2, 3, 1])[:dim] + image_start = np.array([20, 30, 4])[:dim] * scale + map_start = np.array([8, 8, 2])[:dim] + map_size = (12, 10, 10)[:dim] + output_start = np.array([22, 33, 5])[:dim] + output_size = (5, 4, 3)[:dim] + displacement = np.array([1., 2., 0.])[:dim] + + def box(start, size): + return bounding_box.BoundingBox( + start=tuple(start) + ((0,) if dim == 2 else ()), + size=tuple(size) + ((1,) if dim == 2 else ()), + ) + + coord_map = np.zeros((dim,) + map_size[::-1]) + coord_map[:] = displacement.reshape((dim,) + (1,) * dim) + result = warp.ndimage_warp( + image, coord_map, tuple(stride_xyz[::-1]), (4, 3, 2)[:dim], + (1,) * dim, order=1, + image_box=box(image_start, image_shape[::-1]), + map_box=box(map_start, map_size), + out_box=box(output_start, output_size), + out_scale=tuple(scale), parallelism=parallelism, + ) + expected = np.zeros(output_size[::-1]) + for axis, coord in enumerate(np.indices(expected.shape)[::-1]): + source_coord = ( + (coord + output_start[axis] + displacement[axis]) * scale[axis] + - image_start[axis] + ) + expected += weights[axis] * source_coord + self.assertEqual(result.shape, expected.shape) + self.assertEqual(result.dtype, image.dtype) + np.testing.assert_allclose(result, expected, rtol=1e-6) + + @parameterized.parameters(2, 3) + def test_output_box_without_map_box(self, dim): + shape = (4, 8, 10)[-dim:] + image = np.arange(np.prod(shape), dtype=np.float32).reshape(shape) + output_size = (6, 5, 1) + result = warp.ndimage_warp( + image, np.zeros((dim,) + shape), (1,) * dim, + (4,) * dim, (1,) * dim, + out_box=bounding_box.BoundingBox(start=(0, 0, 0), size=output_size), + ) + expected = image[tuple(slice(0, n) for n in output_size[::-1][-dim:])] + self.assertEqual(result.shape, expected.shape) + np.testing.assert_array_equal(result, expected) + + if __name__ == '__main__': absltest.main() diff --git a/warp.py b/warp.py index be17c9c..78eff16 100644 --- a/warp.py +++ b/warp.py @@ -255,7 +255,7 @@ def ndimage_warp( src_map += ( map_box.start[:dim] * stride[::-1] - image_box.start[:dim] / out_scale[:dim] - ).reshape(dim, 1, 1, 1) + ).reshape((dim,) + (1,) * dim) # Translate map to source (data) units. reshaper = tuple([slice(None)] + [np.newaxis] * dim) @@ -270,7 +270,7 @@ def ndimage_warp( sub_dim = 1 if out_box is not None: - warped = np.zeros(shape=out_box.size[::-1], dtype=image.dtype) + warped = np.zeros(shape=out_box.size[::-1][sub_dim:], dtype=image.dtype) else: warped = np.zeros_like(image) out_box = bounding_box.BoundingBox(start=(0, 0, 0), size=image_size_xyz) @@ -285,7 +285,9 @@ def ndimage_warp( # Compute the position of map_box relative to the out_box. if map_box is not None: assert out_box is not None - offset = (map_box.start * stride[::-1] - out_box.start)[::-1] + offset = ( + map_box.start[:dim] * stride[::-1] - out_box.start[:dim] + )[::-1] else: offset = (0, 0, 0)