From 7b4aab9c1b31ca40aaed0ca97e0ea79fffc973fc Mon Sep 17 00:00:00 2001 From: filipstrand Date: Sun, 22 Sep 2024 20:04:46 +0200 Subject: [PATCH] Move actual image saving logic to ImageUtil --- src/mflux/post_processing/generated_image.py | 63 +----------------- src/mflux/post_processing/image_util.py | 67 +++++++++++++++++++- 2 files changed, 67 insertions(+), 63 deletions(-) diff --git a/src/mflux/post_processing/generated_image.py b/src/mflux/post_processing/generated_image.py index 64c12e6..85a21e4 100644 --- a/src/mflux/post_processing/generated_image.py +++ b/src/mflux/post_processing/generated_image.py @@ -1,16 +1,11 @@ import importlib -import json -import logging -from pathlib import Path import PIL.Image import mlx.core as mx -import piexif +from mflux import ImageUtil from mflux.config.model_config import ModelConfig -log = logging.getLogger(__name__) - class GeneratedImage: def __init__( @@ -44,61 +39,7 @@ class GeneratedImage: self.controlnet_strength = controlnet_strength def save(self, path: str, export_json_metadata: bool = False) -> None: - file_path = Path(path) - file_path.parent.mkdir(parents=True, exist_ok=True) - file_name = file_path.stem - file_extension = file_path.suffix - - # If a file already exists, create a new name with a counter - counter = 1 - while file_path.exists(): - new_name = f"{file_name}_{counter}{file_extension}" - file_path = file_path.with_name(new_name) - counter += 1 - - try: - # Save image without metadata first - self.image.save(file_path) - log.info(f"Image saved successfully at: {file_path}") - - # Optionally save json metadata file - if export_json_metadata: - with open(f"{file_path.with_suffix('.json')}", "w") as json_file: - json.dump(self._get_metadata(), json_file, indent=4) - - # Embed metadata - self._embed_metadata(file_path) - log.info(f"Metadata embedded successfully at: {file_path}") - except Exception as e: - log.error(f"Error saving image: {e}") - - def _embed_metadata(self, path: str) -> None: - try: - # Prepare metadata - metadata = self._get_metadata() - - # Convert metadata dictionary to a string - metadata_str = str(metadata) - - # Convert the string to bytes (using UTF-8 encoding) - user_comment_bytes = metadata_str.encode("utf-8") - - # Define the UserComment tag ID - USER_COMMENT_TAG_ID = 0x9286 - - # Create a piexif-compatible dictionary structure - exif_piexif_dict = {"Exif": {USER_COMMENT_TAG_ID: user_comment_bytes}} - - # Load the image and embed the EXIF data - image = PIL.Image.open(path) - exif_bytes = piexif.dump(exif_piexif_dict) - image.info["exif"] = exif_bytes - - # Save the image with metadata - image.save(path, exif=exif_bytes) - - except Exception as e: - log.error(f"Error embedding metadata: {e}") + ImageUtil.save(self.image, path, self._get_metadata(), export_json_metadata) def _get_metadata(self) -> dict: return { diff --git a/src/mflux/post_processing/image_util.py b/src/mflux/post_processing/image_util.py index 79552b6..190b4ad 100644 --- a/src/mflux/post_processing/image_util.py +++ b/src/mflux/post_processing/image_util.py @@ -1,12 +1,19 @@ +import json +import logging +from pathlib import Path + import PIL -from PIL import Image +import PIL.Image import mlx.core as mx import numpy as np - +from PIL import Image +import piexif from mflux.config.config import ConfigControlnet from mflux.config.runtime_config import RuntimeConfig from mflux.post_processing.generated_image import GeneratedImage +log = logging.getLogger(__name__) + class ImageUtil: @staticmethod @@ -78,3 +85,59 @@ class ImageUtil: @staticmethod def load_image(path: str) -> Image.Image: return Image.open(path) + + @staticmethod + def save(image: PIL.Image.Image, path: str, metadata: dict, export_json_metadata: bool = False) -> None: + file_path = Path(path) + file_path.parent.mkdir(parents=True, exist_ok=True) + file_name = file_path.stem + file_extension = file_path.suffix + + # If a file already exists, create a new name with a counter + counter = 1 + while file_path.exists(): + new_name = f"{file_name}_{counter}{file_extension}" + file_path = file_path.with_name(new_name) + counter += 1 + + try: + # Save image without metadata first + image.save(file_path) + log.info(f"Image saved successfully at: {file_path}") + + # Optionally save json metadata file + if export_json_metadata: + with open(f"{file_path.with_suffix('.json')}", "w") as json_file: + json.dump(metadata, json_file, indent=4) + + # Embed metadata + ImageUtil._embed_metadata(metadata, file_path) + log.info(f"Metadata embedded successfully at: {file_path}") + except Exception as e: + log.error(f"Error saving image: {e}") + + @staticmethod + def _embed_metadata(metadata: dict, path: str) -> None: + try: + # Convert metadata dictionary to a string + metadata_str = str(metadata) + + # Convert the string to bytes (using UTF-8 encoding) + user_comment_bytes = metadata_str.encode("utf-8") + + # Define the UserComment tag ID + USER_COMMENT_TAG_ID = 0x9286 + + # Create a piexif-compatible dictionary structure + exif_piexif_dict = {"Exif": {USER_COMMENT_TAG_ID: user_comment_bytes}} + + # Load the image and embed the EXIF data + image = PIL.Image.open(path) + exif_bytes = piexif.dump(exif_piexif_dict) + image.info["exif"] = exif_bytes + + # Save the image with metadata + image.save(path, exif=exif_bytes) + + except Exception as e: + log.error(f"Error embedding metadata: {e}")