diff --git a/stitching/cropper.py b/stitching/cropper.py index 4668f9a..f0a1916 100644 --- a/stitching/cropper.py +++ b/stitching/cropper.py @@ -47,6 +47,7 @@ class Cropper: def __init__(self, crop=DEFAULT_CROP): self.do_crop = crop + self.overlapping_indices = [] self.overlapping_rectangles = [] self.cropping_rectangles = [] @@ -56,14 +57,20 @@ def prepare(self, imgs, masks, corners, sizes): lir = self.estimate_largest_interior_rectangle(mask) corners = self.get_zero_center_corners(corners) rectangles = self.get_rectangles(corners, sizes) + self.overlapping_indices = self.get_overlapping_indices(rectangles, lir) + rectangles = [rectangles[idx] for idx in self.overlapping_indices] self.overlapping_rectangles = self.get_overlaps(rectangles, lir) self.intersection_rectangles = self.get_intersections( rectangles, self.overlapping_rectangles ) def crop_images(self, imgs, aspect=1): + crop_idx = 0 + overlapping_indices = set(self.overlapping_indices) for idx, img in enumerate(imgs): - yield self.crop_img(img, idx, aspect) + if not self.do_crop or idx in overlapping_indices: + yield self.crop_img(img, crop_idx, aspect) + crop_idx += 1 def crop_img(self, img, idx, aspect=1): if self.do_crop: @@ -117,6 +124,20 @@ def get_rectangles(corners, sizes): rectangles.append(rectangle) return rectangles + @staticmethod + def get_overlapping_indices(rectangles, lir): + return [ + idx + for idx, rectangle in enumerate(rectangles) + if Cropper.has_overlap(rectangle, lir) + ] + + @staticmethod + def has_overlap(rectangle1, rectangle2): + overlap_x = rectangle1.x < rectangle2.x2 and rectangle2.x < rectangle1.x2 + overlap_y = rectangle1.y < rectangle2.y2 and rectangle2.y < rectangle1.y2 + return overlap_x and overlap_y + @staticmethod def get_overlaps(rectangles, lir): return [Cropper.get_overlap(r, lir) for r in rectangles] @@ -127,7 +148,7 @@ def get_overlap(rectangle1, rectangle2): y1 = max(rectangle1.y, rectangle2.y) x2 = min(rectangle1.x2, rectangle2.x2) y2 = min(rectangle1.y2, rectangle2.y2) - if x2 < x1 or y2 < y1: + if x2 <= x1 or y2 <= y1: raise StitchingError("Rectangles do not overlap!") return Rectangle(x1, y1, x2 - x1, y2 - y1) diff --git a/tests/context.py b/tests/context.py index 8fc567d..f340d24 100644 --- a/tests/context.py +++ b/tests/context.py @@ -11,7 +11,7 @@ from stitching.camera_estimator import CameraEstimator # noqa: F401, E402 from stitching.camera_wave_corrector import WaveCorrector # noqa: F401, E402 from stitching.cli.stitch import create_parser, main # noqa: F401, E402 -from stitching.cropper import Cropper # noqa: F401, E402 +from stitching.cropper import Cropper, Rectangle # noqa: F401, E402 from stitching.exposure_error_compensator import ( # noqa: F401, E402 ExposureErrorCompensator, ) diff --git a/tests/test_cropper.py b/tests/test_cropper.py new file mode 100644 index 0000000..90c565d --- /dev/null +++ b/tests/test_cropper.py @@ -0,0 +1,37 @@ +import unittest +from unittest.mock import patch + +import numpy as np + +from .context import Cropper, Rectangle + + +class TestCropper(unittest.TestCase): + def test_prepare_ignores_rectangles_outside_lir(self): + cropper = Cropper() + imgs = [np.full((10, 10), idx, dtype=np.uint8) for idx in range(3)] + masks = [np.ones((10, 10), dtype=np.uint8) for _ in imgs] + corners = [(0, 0), (20, 20), (5, 5)] + sizes = [(10, 10)] * len(imgs) + lir = Rectangle(2, 2, 12, 12) + + with patch.object(cropper, "estimate_panorama_mask"), patch.object( + cropper, "estimate_largest_interior_rectangle", return_value=lir + ): + cropper.prepare(imgs, masks, corners, sizes) + + cropped_imgs = list(cropper.crop_images(iter(imgs))) + cropped_corners, cropped_sizes = cropper.crop_rois(corners, sizes) + + self.assertEqual(cropper.overlapping_indices, [0, 2]) + self.assertEqual(len(cropped_imgs), 2) + self.assertEqual(cropped_imgs[0].shape, (8, 8)) + self.assertEqual(cropped_imgs[1].shape, (9, 9)) + self.assertTrue(np.all(cropped_imgs[0] == 0)) + self.assertTrue(np.all(cropped_imgs[1] == 2)) + self.assertEqual(cropped_corners, [(0, 0), (3, 3)]) + self.assertEqual(cropped_sizes, [(8, 8), (9, 9)]) + + +if __name__ == "__main__": + unittest.main()