diff --git a/src/diffusers/image_processor.py b/src/diffusers/image_processor.py index 4f6f4bd52b9c..6dccc02dd5ef 100644 --- a/src/diffusers/image_processor.py +++ b/src/diffusers/image_processor.py @@ -14,6 +14,7 @@ import math import warnings +from typing import Any import numpy as np import PIL.Image @@ -885,7 +886,7 @@ def preprocess( height: int | None = None, width: int | None = None, padding_mask_crop: int | None = None, - ) -> tuple[torch.Tensor, torch.Tensor]: + ) -> tuple[torch.Tensor, torch.Tensor | None, dict[str, Any]]: """ Preprocess the image and mask. """ @@ -894,7 +895,13 @@ def preprocess( # if mask is None, same behavior as regular image processor if mask is None: - return self._image_processor.preprocess(image, height=height, width=width) + processed_image = self._image_processor.preprocess(image, height=height, width=width) + postprocessing_kwargs = { + "crops_coords": None, + "original_image": None, + "original_mask": None, + } + return processed_image, None, postprocessing_kwargs if padding_mask_crop is not None: crops_coords = self._image_processor.get_crop_region(mask, width, height, pad=padding_mask_crop) diff --git a/tests/others/test_image_processor.py b/tests/others/test_image_processor.py index 0d358699f105..05cbb8a27a2d 100644 --- a/tests/others/test_image_processor.py +++ b/tests/others/test_image_processor.py @@ -15,9 +15,10 @@ import numpy as np import PIL.Image +import pytest import torch -from diffusers.image_processor import VaeImageProcessor +from diffusers.image_processor import InpaintProcessor, VaeImageProcessor class TestImageProcessor: @@ -306,3 +307,38 @@ def test_vae_image_processor_resize_np(self): assert out_np.shape == exp_np_shape, ( f"resized image output shape '{out_np.shape}' didn't match expected shape '{exp_np_shape}'." ) + + @pytest.mark.parametrize( + "has_mask, padding_mask_crop", + [ + (True, None), + (False, None), + (True, 8), + ], + ) + def test_inpaint_processor_preprocess(self, has_mask, padding_mask_crop): + processor = InpaintProcessor() + image = PIL.Image.fromarray(np.zeros((64, 64, 3), dtype=np.uint8)) + mask = PIL.Image.fromarray(np.ones((64, 64), dtype=np.uint8) * 255) if has_mask else None + + kwargs = {"height": 64, "width": 64} + if padding_mask_crop is not None: + kwargs["padding_mask_crop"] = padding_mask_crop + + out_img, out_mask, postprocessing_kwargs = processor.preprocess(image, mask=mask, **kwargs) + + assert isinstance(out_img, torch.Tensor) + if has_mask: + assert isinstance(out_mask, torch.Tensor) + else: + assert out_mask is None + assert isinstance(postprocessing_kwargs, dict) + + if padding_mask_crop is not None: + assert postprocessing_kwargs["crops_coords"] is not None + assert postprocessing_kwargs["original_image"] == image + assert postprocessing_kwargs["original_mask"] == mask + else: + assert postprocessing_kwargs["crops_coords"] is None + assert postprocessing_kwargs["original_image"] is None + assert postprocessing_kwargs["original_mask"] is None