Merge pull request #21527 from kostrykin/patch-1

Extend `image_diff` and image content assertions to handle boolean images
This commit is contained in:
Marius van den Beek
2026-01-12 12:39:18 +01:00
committed by GitHub
6 changed files with 49 additions and 4 deletions
+6
View File
@@ -595,6 +595,12 @@ def files_image_diff(file1: str, file2: str, attributes: Optional[Dict[str, Any]
if arr1.shape != arr2.shape:
raise AssertionError(f"Image dimensions did not match ({arr1.shape}, {arr2.shape}).")
# Handle `bool` images by converting them to `uint8`
if numpy.issubdtype(arr1.dtype, bool):
arr1 = arr1.astype(numpy.uint8)
if numpy.issubdtype(arr2.dtype, bool):
arr2 = arr2.astype(numpy.uint8)
distance = get_image_metric(attributes)(arr1, arr2)
distance_eps = attributes.get("eps", DEFAULT_EPS)
if distance > distance_eps:
+7 -1
View File
@@ -747,8 +747,14 @@ def _get_image_labels(
label = label.strip()
if numpy.issubdtype(im_arr.dtype, numpy.integer):
return int(label)
if numpy.issubdtype(im_arr.dtype, float):
if numpy.issubdtype(im_arr.dtype, numpy.floating):
return float(label)
if numpy.issubdtype(im_arr.dtype, bool):
label_lower = label.lower()
if label_lower in ("0", "1", "false", "true"):
return label_lower in ("1", "true")
else:
raise AssertionError(f'Label "{label}" incompatible with label type "{im_arr.dtype}"')
raise AssertionError(f'Unsupported image label type: "{im_arr.dtype}"')
# Determine labels present in the image.
Binary file not shown.
+5
View File
@@ -27,6 +27,11 @@
<param name="in" value="im4_float.tif" />
<output name="out" value="im4_float.tif" compare="image_diff" />
</test>
<!-- test pair of equal images (bool tiff) -->
<test>
<param name="in" value="im3_b_bool.tiff" />
<output name="out" value="im3_b_bool.tiff" compare="image_diff" metric="iou" pin_labels="0" />
</test>
<!-- test pair of different images -->
<test>
<param name="in" value="im2_a.png" />
@@ -181,6 +181,20 @@
</assert_contents>
</output>
</test>
<test>
<param name="input" value="im3_b_bool.tiff" />
<output name="output">
<assert_contents>
<has_image_width width="32" />
<has_image_height height="32" />
<has_image_channels channels="1" />
<has_image_n_labels n="1" exclude_labels="0" />
<has_image_n_labels n="1" exclude_labels="1" />
<has_image_n_labels n="0" exclude_labels="0,1" />
<has_image_mean_object_size mean_object_size="256" exclude_labels="0" />
</assert_contents>
</output>
</test>
<test>
<param name="input" value="im2_b.png" />
<output name="output">
+17 -3
View File
@@ -87,6 +87,17 @@ F9 = _encode_image(
),
format="PNG",
)
F10 = _encode_image(
numpy.array(
[
[True, True, True],
[True, False, True],
[True, True, True],
],
dtype=bool,
),
format="TIFF",
)
def _test_file_list():
@@ -102,6 +113,7 @@ def _test_file_list():
(F7, ".tiff"),
(F8, ".tiff"),
(F9, ".png"),
(F10, ".tiff"),
]:
with tempfile.NamedTemporaryFile(mode="wb", suffix=ext, delete=False) as out:
if ext == ".txt.gz":
@@ -112,7 +124,7 @@ def _test_file_list():
def generate_tests(multiline=False):
f1, f2, f3, f4, multiline_match, f5, f6, f7, f8, f9 = _test_file_list()
f1, f2, f3, f4, multiline_match, f5, f6, f7, f8, f9, f10 = _test_file_list()
tests: List[TestDef]
if multiline:
tests = [(multiline_match, f1, {"lines_diff": 0, "sort": True}, None)]
@@ -130,7 +142,7 @@ def generate_tests(multiline=False):
def generate_tests_sim_size():
f1, f2, f3, f4, multiline_match, f5, f6, f7, f8, f9 = _test_file_list()
f1, f2, f3, f4, multiline_match, f5, f6, f7, f8, f9, f10 = _test_file_list()
# tests for equal files
tests: List[TestDef] = [
(f1, f1, None, None), # pass default values
@@ -156,7 +168,7 @@ def generate_tests_sim_size():
def generate_tests_image_diff():
f1, f2, f3, f4, multiline_match, f5, f6, f7, f8, f9 = _test_file_list()
f1, f2, f3, f4, multiline_match, f5, f6, f7, f8, f9, f10 = _test_file_list()
metrics = ["mae", "mse", "rms", "fro", "iou"]
# tests for equal files (uint8, PNG)
tests: List[TestDef] = [(f6, f6, {"metric": metric}, None) for metric in metrics]
@@ -164,6 +176,8 @@ def generate_tests_image_diff():
tests += [(f7, f7, {"metric": metric}, None) for metric in metrics]
# tests for equal files (float, TIFF)
tests += [(f8, f8, {"metric": metric}, None) for metric in metrics]
# tests for equal files (bool, TIFF)
tests += [(f10, f10, {"metric": metric}, None) for metric in metrics]
# tests for pairs of different files
tests += [(f6, f8, {"metric": metric}, AssertionError) for metric in metrics] # uint8 vs float
tests += [(f7, f8, {"metric": metric}, AssertionError) for metric in metrics] # uint8 vs float