Qwen-Image-Layered-MRP-MLX/tests/image_generation/helpers/image_compare.py
Anthony Wu 8c2cefdd51
Update tests/image_generation image comparison logic to allow similar images that are _close enough_ (#263)
Co-authored-by: Anthony Wu <pls-file-gh-issue@users.noreply.github.com>
Co-authored-by: filipstrand <strand.filip@gmail.com>
2025-09-08 16:49:54 +02:00

47 lines
1.8 KiB
Python

import os
from pathlib import Path
import cv2
import numpy as np
from PIL import Image
class ReferenceVsOutputImageError(AssertionError): ...
# How we determined DEFAULT_MISMATCH_THRESHOLD value: by eye test
# successive mlx/mflux versions have generated images that are visually close enough
# to consider as a valid upgrade. Minor visual differences are attributable to mlx updates
DEFAULT_MISMATCH_THRESHOLD = 0.15
ENV_MISMATCH_THRESHOLD = float(os.environ.get("MFLUX_IMAGE_MISMATCH_THRESHOLD", DEFAULT_MISMATCH_THRESHOLD))
def check_images_close_enough(
image1_path: str | Path,
image2_path: str | Path,
error_message_prefix: str,
mismatch_threshold: float = ENV_MISMATCH_THRESHOLD,
):
image1_path = Path(image1_path)
image2_path = Path(image2_path)
image1_data = np.array(Image.open(image1_path))
image2_data = np.array(Image.open(image2_path))
closeness_array = np.isclose(
image1_data,
image2_data,
rtol=float(os.environ.get("MFLUX_IMAGE_ALLCLOSE_RTOL", 0.1)),
atol=0, # ignore absolute tolerance
)
num_mismatched = np.count_nonzero(~closeness_array)
total_elements = image1_data.size
mismatch_ratio = num_mismatched / total_elements
if mismatch_ratio > mismatch_threshold:
diff = image1_data - image2_data
# Scale the difference so it's visible. A difference of 50 will become bright.
# We'll scale it so that a difference of 50 or more becomes pure white (255).
diff_visual = (diff / 50 * 255).clip(0, 255).astype(np.uint8)
cv2.imwrite(image2_path.with_stem(image1_path.stem + "_diff"), diff_visual)
raise AssertionError(
f"{error_message_prefix} Check {image1_path} vs {image2_path} :: their elements are {mismatch_ratio:.1%} different. Fails assertion for {mismatch_threshold=}"
)