Add XMP/IPTC metadata support (#267)
Co-authored-by: filipstrand <strand.filip@gmail.com>
This commit is contained in:
parent
6947a23d6c
commit
83354c3921
@ -254,7 +254,7 @@ mflux-generate \
|
|||||||
|
|
||||||
- **`--lora-scales`** (optional, `[float]`, default: `None`): The scale for each respective [LoRA](#-LoRA) (will default to `1.0` if not specified and only one LoRA weight is loaded.)
|
- **`--lora-scales`** (optional, `[float]`, default: `None`): The scale for each respective [LoRA](#-LoRA) (will default to `1.0` if not specified and only one LoRA weight is loaded.)
|
||||||
|
|
||||||
- **`--metadata`** (optional): Exports a `.json` file containing the metadata for the image with the same name. (Even without this flag, the image metadata is saved and can be viewed using `exiftool image.png`)
|
- **`--metadata`** (optional): Exports a `.json` file containing the metadata for the image with the same name. (Even without this flag, the image metadata is saved and can be viewed using `mflux-info image.png` or `exiftool image.png`)
|
||||||
|
|
||||||
- **`--image-path`** (optional, `str`, default: `None`): Local path to the initial image for image-to-image generation.
|
- **`--image-path`** (optional, `str`, default: `None`): Local path to the initial image for image-to-image generation.
|
||||||
|
|
||||||
@ -643,7 +643,7 @@ system_profiler SPHardwareDataType SPDisplaysDataType
|
|||||||
|
|
||||||
|
|
||||||
*Note that these numbers includes starting the application from scratch, which means doing model i/o, setting/quantizing weights etc.
|
*Note that these numbers includes starting the application from scratch, which means doing model i/o, setting/quantizing weights etc.
|
||||||
If we assume that the model is already loaded, you can inspect the image metadata using `exiftool image.png` and see the total duration of the denoising loop (excluding text embedding).*
|
If we assume that the model is already loaded, you can inspect the image metadata using `mflux-info image.png` (or `exiftool image.png`) and see the total duration of the denoising loop (excluding text embedding).*
|
||||||
|
|
||||||
*These benchmarks are not very scientific and is only intended to give ballpark numbers. They were performed during different times with different MFLUX and MLX-versions etc. Additional hardware information such as number of GPU cores, Mac device etc. are not always known.*
|
*These benchmarks are not very scientific and is only intended to give ballpark numbers. They were performed during different times with different MFLUX and MLX-versions etc. Additional hardware information such as number of GPU cores, Mac device etc. are not always known.*
|
||||||
|
|
||||||
|
|||||||
@ -91,6 +91,7 @@ mflux-save-depth = "mflux.save_depth:main"
|
|||||||
mflux-train = "mflux.train:main"
|
mflux-train = "mflux.train:main"
|
||||||
mflux-upscale = "mflux.upscale:main"
|
mflux-upscale = "mflux.upscale:main"
|
||||||
mflux-lora-library = "mflux.lora_library:main"
|
mflux-lora-library = "mflux.lora_library:main"
|
||||||
|
mflux-info = "mflux.info:main"
|
||||||
mflux-completions = "mflux.ui.cli.completions.install:main"
|
mflux-completions = "mflux.ui.cli.completions.install:main"
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
124
src/mflux/info.py
Normal file
124
src/mflux/info.py
Normal file
@ -0,0 +1,124 @@
|
|||||||
|
"""Display metadata information from MFLUX generated images."""
|
||||||
|
|
||||||
|
import sys
|
||||||
|
from datetime import datetime
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
from mflux.post_processing.metadata_reader import MetadataReader
|
||||||
|
from mflux.ui.cli.parsers import CommandLineParser
|
||||||
|
|
||||||
|
|
||||||
|
def format_metadata(metadata: dict) -> str:
|
||||||
|
"""Format metadata in a clean, readable format."""
|
||||||
|
exif = metadata.get("exif", {})
|
||||||
|
if not exif:
|
||||||
|
return "No metadata found"
|
||||||
|
|
||||||
|
lines = []
|
||||||
|
lines.append("=" * 60)
|
||||||
|
lines.append("MFLUX Image Information")
|
||||||
|
lines.append("=" * 60)
|
||||||
|
|
||||||
|
# Prompt
|
||||||
|
if prompt := exif.get("prompt"):
|
||||||
|
lines.append(f"\nPrompt: {prompt}")
|
||||||
|
|
||||||
|
if negative_prompt := exif.get("negative_prompt"):
|
||||||
|
lines.append(f"Negative Prompt: {negative_prompt}")
|
||||||
|
|
||||||
|
# Model information
|
||||||
|
lines.append("")
|
||||||
|
if model := exif.get("model"):
|
||||||
|
lines.append(f"Model: {model}")
|
||||||
|
|
||||||
|
# Image dimensions
|
||||||
|
if width := exif.get("width"):
|
||||||
|
lines.append(f"Width: {width}")
|
||||||
|
if height := exif.get("height"):
|
||||||
|
lines.append(f"Height: {height}")
|
||||||
|
|
||||||
|
# Generation parameters
|
||||||
|
lines.append("")
|
||||||
|
if seed := exif.get("seed"):
|
||||||
|
lines.append(f"Seed: {seed}")
|
||||||
|
if steps := exif.get("steps"):
|
||||||
|
lines.append(f"Steps: {steps}")
|
||||||
|
if guidance := exif.get("guidance"):
|
||||||
|
lines.append(f"Guidance: {guidance}")
|
||||||
|
|
||||||
|
# Technical settings
|
||||||
|
if quantize := exif.get("quantize"):
|
||||||
|
lines.append(f"Quantization: {quantize}-bit")
|
||||||
|
if precision := exif.get("precision"):
|
||||||
|
lines.append(f"Precision: {precision}")
|
||||||
|
|
||||||
|
# LoRA information
|
||||||
|
if lora_paths := exif.get("lora_paths"):
|
||||||
|
lines.append("")
|
||||||
|
lines.append(f"LoRAs ({len(lora_paths)}):")
|
||||||
|
lora_scales = exif.get("lora_scales", [])
|
||||||
|
for i, lora in enumerate(lora_paths):
|
||||||
|
scale = lora_scales[i] if i < len(lora_scales) else 1.0
|
||||||
|
lora_name = Path(lora).name
|
||||||
|
lines.append(f" - {lora_name} (scale: {scale})")
|
||||||
|
|
||||||
|
# Image-to-image parameters
|
||||||
|
if image_path := exif.get("image_path"):
|
||||||
|
lines.append("")
|
||||||
|
lines.append(f"Source Image: {Path(image_path).name}")
|
||||||
|
if image_strength := exif.get("image_strength"):
|
||||||
|
lines.append(f"Image Strength: {image_strength}")
|
||||||
|
|
||||||
|
# ControlNet parameters
|
||||||
|
if controlnet_path := exif.get("controlnet_image_path"):
|
||||||
|
lines.append("")
|
||||||
|
lines.append(f"ControlNet Image: {Path(controlnet_path).name}")
|
||||||
|
if controlnet_strength := exif.get("controlnet_strength"):
|
||||||
|
lines.append(f"ControlNet Strength: {controlnet_strength}")
|
||||||
|
|
||||||
|
# Generation metadata
|
||||||
|
lines.append("")
|
||||||
|
if gen_time := exif.get("generation_time_seconds"):
|
||||||
|
lines.append(f"Generation Time: {gen_time:.2f}s")
|
||||||
|
|
||||||
|
if created_at := exif.get("created_at"):
|
||||||
|
try:
|
||||||
|
dt = datetime.fromisoformat(created_at)
|
||||||
|
lines.append(f"Created: {dt.strftime('%Y-%m-%d %H:%M:%S')}")
|
||||||
|
except (ValueError, AttributeError):
|
||||||
|
lines.append(f"Created: {created_at}")
|
||||||
|
|
||||||
|
if version := exif.get("mflux_version"):
|
||||||
|
lines.append(f"MFLUX Version: {version}")
|
||||||
|
|
||||||
|
lines.append("=" * 60)
|
||||||
|
return "\n".join(lines)
|
||||||
|
|
||||||
|
|
||||||
|
def main():
|
||||||
|
# Parse command line arguments
|
||||||
|
parser = CommandLineParser(description="Display metadata from MFLUX generated images")
|
||||||
|
parser.add_info_arguments()
|
||||||
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
# Check if file exists
|
||||||
|
image_path = Path(args.image_path)
|
||||||
|
if not image_path.exists():
|
||||||
|
print(f"Error: Image file not found: {image_path}")
|
||||||
|
sys.exit(1)
|
||||||
|
|
||||||
|
# Read metadata
|
||||||
|
metadata = MetadataReader.read_all_metadata(image_path)
|
||||||
|
|
||||||
|
# Check if metadata was found
|
||||||
|
if not metadata or (not metadata.get("exif") and not metadata.get("xmp")):
|
||||||
|
print("No metadata found")
|
||||||
|
sys.exit(1)
|
||||||
|
|
||||||
|
# Format and display
|
||||||
|
print(format_metadata(metadata))
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
main()
|
||||||
|
|
||||||
@ -1,3 +1,4 @@
|
|||||||
|
from datetime import datetime
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
import mlx.core as mx
|
import mlx.core as mx
|
||||||
@ -22,6 +23,8 @@ class GeneratedImage:
|
|||||||
generation_time: float,
|
generation_time: float,
|
||||||
lora_paths: list[str],
|
lora_paths: list[str],
|
||||||
lora_scales: list[float],
|
lora_scales: list[float],
|
||||||
|
height: int | None = None,
|
||||||
|
width: int | None = None,
|
||||||
controlnet_image_path: str | Path | None = None,
|
controlnet_image_path: str | Path | None = None,
|
||||||
controlnet_strength: float | None = None,
|
controlnet_strength: float | None = None,
|
||||||
image_path: str | Path | None = None,
|
image_path: str | Path | None = None,
|
||||||
@ -44,6 +47,8 @@ class GeneratedImage:
|
|||||||
self.generation_time = generation_time
|
self.generation_time = generation_time
|
||||||
self.lora_paths = lora_paths
|
self.lora_paths = lora_paths
|
||||||
self.lora_scales = lora_scales
|
self.lora_scales = lora_scales
|
||||||
|
self.height = height
|
||||||
|
self.width = width
|
||||||
self.controlnet_image_path = controlnet_image_path
|
self.controlnet_image_path = controlnet_image_path
|
||||||
self.controlnet_strength = controlnet_strength
|
self.controlnet_strength = controlnet_strength
|
||||||
self.image_path = image_path
|
self.image_path = image_path
|
||||||
@ -73,6 +78,8 @@ class GeneratedImage:
|
|||||||
generation_time=self.generation_time,
|
generation_time=self.generation_time,
|
||||||
lora_paths=self.lora_paths,
|
lora_paths=self.lora_paths,
|
||||||
lora_scales=self.lora_scales,
|
lora_scales=self.lora_scales,
|
||||||
|
height=self.height,
|
||||||
|
width=self.width,
|
||||||
controlnet_image_path=self.controlnet_image_path,
|
controlnet_image_path=self.controlnet_image_path,
|
||||||
controlnet_strength=self.controlnet_strength,
|
controlnet_strength=self.controlnet_strength,
|
||||||
image_path=self.image_path,
|
image_path=self.image_path,
|
||||||
@ -139,9 +146,12 @@ class GeneratedImage:
|
|||||||
"seed": self.seed,
|
"seed": self.seed,
|
||||||
"steps": self.steps,
|
"steps": self.steps,
|
||||||
"guidance": self.guidance if self.model_config.supports_guidance else None,
|
"guidance": self.guidance if self.model_config.supports_guidance else None,
|
||||||
|
"height": self.height,
|
||||||
|
"width": self.width,
|
||||||
"precision": str(self.precision),
|
"precision": str(self.precision),
|
||||||
"quantize": self.quantization,
|
"quantize": self.quantization,
|
||||||
"generation_time_seconds": round(self.generation_time, 2),
|
"generation_time_seconds": round(self.generation_time, 2),
|
||||||
|
"created_at": datetime.now().isoformat(),
|
||||||
"lora_paths": [str(p) for p in self.lora_paths] if self.lora_paths else None,
|
"lora_paths": [str(p) for p in self.lora_paths] if self.lora_paths else None,
|
||||||
"lora_scales": [round(scale, 2) for scale in self.lora_scales] if self.lora_scales else None,
|
"lora_scales": [round(scale, 2) for scale in self.lora_scales] if self.lora_scales else None,
|
||||||
"image_path": str(self.image_path) if self.image_path else None,
|
"image_path": str(self.image_path) if self.image_path else None,
|
||||||
|
|||||||
@ -12,6 +12,7 @@ from PIL._typing import StrOrBytesPath
|
|||||||
from mflux.config.runtime_config import RuntimeConfig
|
from mflux.config.runtime_config import RuntimeConfig
|
||||||
from mflux.models.flux.variants.concept_attention.attention_data import ConceptHeatmap
|
from mflux.models.flux.variants.concept_attention.attention_data import ConceptHeatmap
|
||||||
from mflux.post_processing.generated_image import GeneratedImage
|
from mflux.post_processing.generated_image import GeneratedImage
|
||||||
|
from mflux.post_processing.metadata_builder import MetadataBuilder
|
||||||
from mflux.ui.box_values import AbsoluteBoxValues, BoxValues
|
from mflux.ui.box_values import AbsoluteBoxValues, BoxValues
|
||||||
|
|
||||||
log = logging.getLogger(__name__)
|
log = logging.getLogger(__name__)
|
||||||
@ -53,6 +54,8 @@ class ImageUtil:
|
|||||||
generation_time=generation_time,
|
generation_time=generation_time,
|
||||||
lora_paths=lora_paths,
|
lora_paths=lora_paths,
|
||||||
lora_scales=lora_scales,
|
lora_scales=lora_scales,
|
||||||
|
height=config.height,
|
||||||
|
width=config.width,
|
||||||
image_path=image_path,
|
image_path=image_path,
|
||||||
image_strength=image_strength,
|
image_strength=image_strength,
|
||||||
controlnet_image_path=controlnet_image_path,
|
controlnet_image_path=controlnet_image_path,
|
||||||
@ -242,21 +245,24 @@ class ImageUtil:
|
|||||||
with open(f"{file_path.with_suffix('.json')}", "w") as json_file:
|
with open(f"{file_path.with_suffix('.json')}", "w") as json_file:
|
||||||
json.dump(metadata, json_file, indent=4)
|
json.dump(metadata, json_file, indent=4)
|
||||||
|
|
||||||
# Embed metadata
|
# Embed metadata in multiple formats for maximum compatibility
|
||||||
if metadata is not None:
|
if metadata is not None:
|
||||||
ImageUtil._embed_metadata(metadata, file_path)
|
ImageUtil._embed_metadata(metadata, file_path)
|
||||||
|
MetadataBuilder.embed_metadata(metadata, file_path)
|
||||||
log.info(f"Metadata embedded successfully at: {file_path}")
|
log.info(f"Metadata embedded successfully at: {file_path}")
|
||||||
except Exception as e: # noqa: BLE001
|
except Exception as e: # noqa: BLE001
|
||||||
log.error(f"Error saving image: {e}")
|
log.error(f"Error saving image: {e}")
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _embed_metadata(metadata: dict, path: str | Path) -> None:
|
def _embed_metadata(metadata: dict, path: str | Path) -> None:
|
||||||
|
"""Original EXIF metadata embedding - preserved for compatibility"""
|
||||||
try:
|
try:
|
||||||
# Convert metadata dictionary to a string
|
# Convert metadata dictionary to a string
|
||||||
metadata_str = json.dumps(metadata)
|
metadata_str = json.dumps(metadata)
|
||||||
|
|
||||||
# Convert the string to bytes (using UTF-8 encoding)
|
# Convert the string to bytes (using UTF-8 encoding)
|
||||||
user_comment_bytes = metadata_str.encode("utf-8")
|
# Add the ASCII character code prefix required by EXIF spec
|
||||||
|
user_comment_bytes = b"ASCII\x00\x00\x00" + metadata_str.encode("utf-8")
|
||||||
|
|
||||||
# Define the UserComment tag ID
|
# Define the UserComment tag ID
|
||||||
USER_COMMENT_TAG_ID = 0x9286
|
USER_COMMENT_TAG_ID = 0x9286
|
||||||
@ -273,7 +279,7 @@ class ImageUtil:
|
|||||||
image.save(path, exif=exif_bytes)
|
image.save(path, exif=exif_bytes)
|
||||||
|
|
||||||
except Exception as e: # noqa: BLE001
|
except Exception as e: # noqa: BLE001
|
||||||
log.error(f"Error embedding metadata: {e}")
|
log.error(f"Error embedding EXIF metadata: {e}")
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def preprocess_for_model(
|
def preprocess_for_model(
|
||||||
|
|||||||
231
src/mflux/post_processing/metadata_builder.py
Normal file
231
src/mflux/post_processing/metadata_builder.py
Normal file
@ -0,0 +1,231 @@
|
|||||||
|
"""Metadata builder for XMP and IPTC formats."""
|
||||||
|
|
||||||
|
import logging
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
import PIL.Image
|
||||||
|
|
||||||
|
log = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class MetadataBuilder:
|
||||||
|
"""Builds XMP and IPTC metadata packets for image embedding."""
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def embed_metadata(metadata: dict, path: str | Path) -> None:
|
||||||
|
"""
|
||||||
|
Embed XMP and IPTC metadata into an image file without touching existing EXIF.
|
||||||
|
Only supports PNG format.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
metadata: Dictionary containing image generation metadata
|
||||||
|
path: Path to the image file to embed metadata into
|
||||||
|
|
||||||
|
Raises:
|
||||||
|
Exception: If there's an error during metadata embedding
|
||||||
|
"""
|
||||||
|
# Check if file is PNG format
|
||||||
|
path_obj = Path(path) if isinstance(path, str) else path
|
||||||
|
if path_obj.suffix.lower() != ".png":
|
||||||
|
log.warning(f"XMP/IPTC metadata embedding is only supported for PNG files, skipping: {path_obj}")
|
||||||
|
return
|
||||||
|
|
||||||
|
try:
|
||||||
|
from PIL import PngImagePlugin
|
||||||
|
|
||||||
|
# Load the image preserving existing metadata
|
||||||
|
image = PIL.Image.open(path)
|
||||||
|
|
||||||
|
# Get existing PNG info to preserve it
|
||||||
|
existing_info = image.info if hasattr(image, "info") else {}
|
||||||
|
|
||||||
|
# Preserve existing EXIF separately (if it exists)
|
||||||
|
existing_exif = existing_info.get("exif")
|
||||||
|
|
||||||
|
# Create new PngInfo preserving existing data
|
||||||
|
pnginfo = PngImagePlugin.PngInfo()
|
||||||
|
|
||||||
|
# Copy existing metadata
|
||||||
|
for key, value in existing_info.items():
|
||||||
|
if key not in ["XML:com.adobe.xmp", "IPTC", "exif"]: # Handle these separately
|
||||||
|
pnginfo.add_text(key, str(value))
|
||||||
|
|
||||||
|
# Build XMP and IPTC metadata using builder methods
|
||||||
|
xmp_packet = MetadataBuilder.build_xmp_packet(metadata)
|
||||||
|
iptc_binary = MetadataBuilder.build_iptc_binary(metadata)
|
||||||
|
|
||||||
|
# Add XMP and IPTC to PNG info
|
||||||
|
pnginfo.add_text("XML:com.adobe.xmp", xmp_packet)
|
||||||
|
if iptc_binary:
|
||||||
|
pnginfo.add_text("IPTC", iptc_binary.hex())
|
||||||
|
|
||||||
|
# Save preserving ALL existing metadata + adding XMP/IPTC
|
||||||
|
# Pass exif separately to preserve it correctly
|
||||||
|
if existing_exif:
|
||||||
|
image.save(path, pnginfo=pnginfo, exif=existing_exif)
|
||||||
|
else:
|
||||||
|
image.save(path, pnginfo=pnginfo)
|
||||||
|
|
||||||
|
except Exception as e: # noqa: BLE001
|
||||||
|
log.error(f"Error embedding XMP/IPTC metadata: {e}")
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def build_xmp_packet(metadata: dict) -> str:
|
||||||
|
"""
|
||||||
|
Build an XMP metadata packet from the provided metadata dictionary.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
metadata: Dictionary containing image generation metadata
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
XMP packet as a string
|
||||||
|
"""
|
||||||
|
# Escape prompt for XML
|
||||||
|
prompt_escaped = metadata.get("prompt", "").replace("&", "&").replace("<", "<").replace(">", ">")
|
||||||
|
|
||||||
|
# Build LoRA info for XMP
|
||||||
|
lora_info = MetadataBuilder._build_lora_string(metadata)
|
||||||
|
|
||||||
|
# Get version from metadata
|
||||||
|
version = metadata.get("mflux_version", "unknown")
|
||||||
|
|
||||||
|
xmp_packet = f"""<?xpacket begin="" id="W5M0MpCehiHzreSzNTczkc9d"?>
|
||||||
|
<x:xmpmeta xmlns:x="adobe:ns:meta/">
|
||||||
|
<rdf:RDF xmlns:rdf="http://www.w3.org/1999/02/22-rdf-syntax-ns#">
|
||||||
|
<rdf:Description rdf:about=""
|
||||||
|
xmlns:dc="http://purl.org/dc/elements/1.1/"
|
||||||
|
xmlns:xmp="http://ns.adobe.com/xap/1.0/"
|
||||||
|
xmlns:photoshop="http://ns.adobe.com/photoshop/1.0/"
|
||||||
|
xmlns:mflux="http://ns.mflux.ai/1.0/">
|
||||||
|
<dc:description><rdf:Alt><rdf:li xml:lang="x-default">{prompt_escaped}</rdf:li></rdf:Alt></dc:description>
|
||||||
|
<dc:creator><rdf:Seq><rdf:li>MFLUX</rdf:li></rdf:Seq></dc:creator>
|
||||||
|
<dc:rights><rdf:Alt><rdf:li xml:lang="x-default">AI Generated Content</rdf:li></rdf:Alt></dc:rights>
|
||||||
|
<xmp:CreatorTool>MFLUX {version}</xmp:CreatorTool>
|
||||||
|
<photoshop:Category>ART</photoshop:Category>
|
||||||
|
<photoshop:Credit>Generated by MFLUX</photoshop:Credit>"""
|
||||||
|
|
||||||
|
# Add technical parameters to XMP
|
||||||
|
if "seed" in metadata:
|
||||||
|
xmp_packet += f"\n <mflux:seed>{metadata['seed']}</mflux:seed>"
|
||||||
|
if "steps" in metadata:
|
||||||
|
xmp_packet += f"\n <mflux:steps>{metadata['steps']}</mflux:steps>"
|
||||||
|
if "guidance" in metadata:
|
||||||
|
xmp_packet += f"\n <mflux:guidance>{metadata['guidance']}</mflux:guidance>"
|
||||||
|
if "model_config" in metadata:
|
||||||
|
xmp_packet += f"\n <mflux:model>{metadata['model_config']}</mflux:model>"
|
||||||
|
if lora_info:
|
||||||
|
xmp_packet += f"\n <mflux:loras>{lora_info}</mflux:loras>"
|
||||||
|
if "generation_time" in metadata:
|
||||||
|
xmp_packet += f"\n <mflux:generationTime>{metadata['generation_time']}</mflux:generationTime>"
|
||||||
|
|
||||||
|
xmp_packet += """
|
||||||
|
</rdf:Description>
|
||||||
|
</rdf:RDF>
|
||||||
|
</x:xmpmeta>
|
||||||
|
<?xpacket end="w"?>"""
|
||||||
|
|
||||||
|
return xmp_packet
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def build_iptc_binary(metadata: dict) -> bytes:
|
||||||
|
"""
|
||||||
|
Build IPTC metadata in binary format from the provided metadata dictionary.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
metadata: Dictionary containing image generation metadata
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
IPTC binary data
|
||||||
|
"""
|
||||||
|
# Build LoRA info for IPTC
|
||||||
|
lora_info = MetadataBuilder._build_lora_string(metadata)
|
||||||
|
|
||||||
|
iptc_data = {}
|
||||||
|
|
||||||
|
# Add prompt information
|
||||||
|
if "prompt" in metadata:
|
||||||
|
prompt = metadata["prompt"]
|
||||||
|
prompt_encoded = prompt.encode("utf-8")
|
||||||
|
if len(prompt_encoded) > 2000:
|
||||||
|
log.warning(f"Prompt is too long ({len(prompt_encoded)} bytes), truncating to 2000 bytes for IPTC")
|
||||||
|
iptc_data[120] = prompt_encoded[:2000] # Caption/Description
|
||||||
|
else:
|
||||||
|
iptc_data[120] = prompt_encoded # Caption/Description
|
||||||
|
iptc_data[5] = f"AI: {prompt[:50]}...".encode("utf-8") # Object Name/Title
|
||||||
|
iptc_data[105] = f"AI Generated: {prompt[:80]}...".encode("utf-8") # Headline
|
||||||
|
|
||||||
|
# Add standard fields
|
||||||
|
iptc_data[80] = b"MFLUX" # By-line (Creator)
|
||||||
|
iptc_data[85] = b"AI Artist" # By-line Title
|
||||||
|
iptc_data[15] = b"ART" # Category
|
||||||
|
iptc_data[110] = b"Generated by MFLUX" # Credit
|
||||||
|
iptc_data[115] = b"AI Generation" # Source
|
||||||
|
iptc_data[116] = b"AI Generated Content" # Copyright Notice
|
||||||
|
iptc_data[118] = b"AI Generated using MFLUX" # Contact
|
||||||
|
iptc_data[103] = b"AI" # Instructions/Special Instructions
|
||||||
|
|
||||||
|
# Add seed and model info in specific fields
|
||||||
|
if "seed" in metadata:
|
||||||
|
iptc_data[122] = f"Seed: {metadata['seed']}".encode("utf-8") # Writer/Editor
|
||||||
|
|
||||||
|
if "model_config" in metadata:
|
||||||
|
iptc_data[90] = f"Model: {metadata['model_config']}".encode("utf-8") # City
|
||||||
|
|
||||||
|
# Add LoRA info in Province/State field
|
||||||
|
if lora_info:
|
||||||
|
iptc_data[95] = f"LoRA: {lora_info}".encode("utf-8") # Province/State
|
||||||
|
|
||||||
|
# Add generation parameters in Country field
|
||||||
|
if "steps" in metadata and "guidance" in metadata:
|
||||||
|
iptc_data[101] = f"Steps:{metadata['steps']} CFG:{metadata['guidance']}".encode("utf-8") # Country
|
||||||
|
|
||||||
|
# Build keywords including LoRA info
|
||||||
|
keywords = ["AI", "Generated", "MFLUX"]
|
||||||
|
if "seed" in metadata:
|
||||||
|
keywords.append(f"seed-{metadata['seed']}")
|
||||||
|
if "steps" in metadata:
|
||||||
|
keywords.append(f"steps-{metadata['steps']}")
|
||||||
|
if "guidance" in metadata:
|
||||||
|
keywords.append(f"guidance-{metadata['guidance']}")
|
||||||
|
if "model_config" in metadata:
|
||||||
|
keywords.append(f"model-{metadata['model_config']}")
|
||||||
|
if lora_info:
|
||||||
|
keywords.append(f"loras-{lora_info}")
|
||||||
|
|
||||||
|
iptc_data[25] = ";".join(keywords).encode("utf-8") # Keywords
|
||||||
|
|
||||||
|
# Build IPTC binary (IPTC field length limit is 32767 bytes)
|
||||||
|
iptc_binary = b""
|
||||||
|
for tag_id, value in iptc_data.items():
|
||||||
|
length = len(value)
|
||||||
|
if length < 32768:
|
||||||
|
iptc_binary += bytes([0x1C, 0x02, tag_id]) + length.to_bytes(2, "big") + value
|
||||||
|
|
||||||
|
return iptc_binary
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _build_lora_string(metadata: dict) -> str:
|
||||||
|
"""
|
||||||
|
Build a LoRA information string from metadata.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
metadata: Dictionary containing lora_paths and lora_scales
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Comma-separated string of LoRA names and scales, or empty string
|
||||||
|
"""
|
||||||
|
if "lora_paths" not in metadata or not metadata["lora_paths"]:
|
||||||
|
return ""
|
||||||
|
|
||||||
|
lora_list = []
|
||||||
|
lora_paths = metadata["lora_paths"]
|
||||||
|
lora_scales = metadata.get("lora_scales", [])
|
||||||
|
|
||||||
|
for i, lora_path in enumerate(lora_paths):
|
||||||
|
# Use Path for OS-agnostic path handling (works on Windows, Mac, Linux)
|
||||||
|
lora_name = Path(lora_path).name
|
||||||
|
scale = lora_scales[i] if i < len(lora_scales) else "1.0"
|
||||||
|
lora_list.append(f"{lora_name}:{scale}")
|
||||||
|
|
||||||
|
return ", ".join(lora_list)
|
||||||
|
|
||||||
126
src/mflux/post_processing/metadata_reader.py
Normal file
126
src/mflux/post_processing/metadata_reader.py
Normal file
@ -0,0 +1,126 @@
|
|||||||
|
"""Metadata reader for extracting and parsing image metadata."""
|
||||||
|
|
||||||
|
import json
|
||||||
|
import logging
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
import piexif
|
||||||
|
import PIL.Image
|
||||||
|
|
||||||
|
log = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class MetadataReader:
|
||||||
|
"""Reads and parses metadata from MFLUX generated images."""
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def read_exif_metadata(image_path: str | Path) -> dict | None:
|
||||||
|
"""
|
||||||
|
Extract EXIF metadata from an image.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
image_path: Path to the image file
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Dictionary containing parsed EXIF metadata, or None if not found
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
img = PIL.Image.open(image_path)
|
||||||
|
exif_bytes = img.info.get("exif")
|
||||||
|
|
||||||
|
if not exif_bytes:
|
||||||
|
return None
|
||||||
|
|
||||||
|
exif_dict = piexif.load(exif_bytes)
|
||||||
|
user_comment = exif_dict["Exif"].get(0x9286, b"")
|
||||||
|
|
||||||
|
if user_comment:
|
||||||
|
# Try to parse as JSON (strip the ASCII prefix if present)
|
||||||
|
if user_comment.startswith(b"ASCII\x00\x00\x00"):
|
||||||
|
metadata_str = user_comment[8:].decode("utf-8")
|
||||||
|
else:
|
||||||
|
metadata_str = user_comment.decode("utf-8")
|
||||||
|
return json.loads(metadata_str)
|
||||||
|
|
||||||
|
return None
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
log.debug(f"Error reading EXIF metadata: {e}")
|
||||||
|
return None
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def read_xmp_metadata(image_path: str | Path) -> dict | None:
|
||||||
|
"""
|
||||||
|
Extract XMP metadata from an image.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
image_path: Path to the image file
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Dictionary containing parsed XMP metadata, or None if not found
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
img = PIL.Image.open(image_path)
|
||||||
|
xmp_data = img.info.get("XML:com.adobe.xmp")
|
||||||
|
|
||||||
|
if not xmp_data:
|
||||||
|
return None
|
||||||
|
|
||||||
|
# Parse XMP XML to extract key fields
|
||||||
|
xmp_dict = {}
|
||||||
|
|
||||||
|
# Simple XML parsing for common fields
|
||||||
|
fields = {
|
||||||
|
"description": "<dc:description><rdf:Alt><rdf:li xml:lang=\"x-default\">",
|
||||||
|
"creator": "<dc:creator><rdf:Seq><rdf:li>",
|
||||||
|
"rights": "<dc:rights><rdf:Alt><rdf:li xml:lang=\"x-default\">",
|
||||||
|
"creator_tool": "<xmp:CreatorTool>",
|
||||||
|
"category": "<photoshop:Category>",
|
||||||
|
"credit": "<photoshop:Credit>",
|
||||||
|
"seed": "<mflux:seed>",
|
||||||
|
"steps": "<mflux:steps>",
|
||||||
|
"guidance": "<mflux:guidance>",
|
||||||
|
"model": "<mflux:model>",
|
||||||
|
"loras": "<mflux:loras>",
|
||||||
|
"generation_time": "<mflux:generationTime>",
|
||||||
|
}
|
||||||
|
|
||||||
|
for key, start_tag in fields.items():
|
||||||
|
if start_tag in xmp_data:
|
||||||
|
start_idx = xmp_data.index(start_tag) + len(start_tag)
|
||||||
|
# Find the closing tag
|
||||||
|
if key in ["description", "rights"]:
|
||||||
|
end_tag = "</rdf:li>"
|
||||||
|
elif key == "creator":
|
||||||
|
end_tag = "</rdf:li>"
|
||||||
|
else:
|
||||||
|
# Extract tag name from start_tag
|
||||||
|
tag_name = start_tag.split(":")[1].rstrip(">")
|
||||||
|
end_tag = f"</{start_tag.split(':')[0]}:{tag_name}>"
|
||||||
|
|
||||||
|
end_idx = xmp_data.index(end_tag, start_idx)
|
||||||
|
value = xmp_data[start_idx:end_idx]
|
||||||
|
xmp_dict[key] = value
|
||||||
|
|
||||||
|
return xmp_dict if xmp_dict else None
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
log.debug(f"Error reading XMP metadata: {e}")
|
||||||
|
return None
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def read_all_metadata(image_path: str | Path) -> dict:
|
||||||
|
"""
|
||||||
|
Read all available metadata from an image.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
image_path: Path to the image file
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Dictionary with 'exif' and 'xmp' keys containing metadata
|
||||||
|
"""
|
||||||
|
return {
|
||||||
|
"exif": MetadataReader.read_exif_metadata(image_path),
|
||||||
|
"xmp": MetadataReader.read_xmp_metadata(image_path),
|
||||||
|
}
|
||||||
|
|
||||||
@ -192,6 +192,9 @@ class CommandLineParser(argparse.ArgumentParser):
|
|||||||
self.add_argument("--train-config", type=str, required=False, help="Local path of the training configuration file")
|
self.add_argument("--train-config", type=str, required=False, help="Local path of the training configuration file")
|
||||||
self.add_argument("--train-checkpoint", type=str, required=False, help="Local path of the checkpoint file which specifies how to continue the training process")
|
self.add_argument("--train-checkpoint", type=str, required=False, help="Local path of the checkpoint file which specifies how to continue the training process")
|
||||||
|
|
||||||
|
def add_info_arguments(self) -> None:
|
||||||
|
self.add_argument("image_path", type=str, help="Path to the image file to inspect")
|
||||||
|
|
||||||
def parse_args(self) -> argparse.Namespace: # type: ignore
|
def parse_args(self) -> argparse.Namespace: # type: ignore
|
||||||
namespace = super().parse_args()
|
namespace = super().parse_args()
|
||||||
|
|
||||||
|
|||||||
2
tests/metadata/__init__.py
Normal file
2
tests/metadata/__init__.py
Normal file
@ -0,0 +1,2 @@
|
|||||||
|
# Metadata tests
|
||||||
|
|
||||||
184
tests/metadata/test_metadata.py
Normal file
184
tests/metadata/test_metadata.py
Normal file
@ -0,0 +1,184 @@
|
|||||||
|
"""
|
||||||
|
Test metadata embedding and reading functionality.
|
||||||
|
|
||||||
|
This test generates a single small image using the schnell model and verifies
|
||||||
|
that all metadata fields are correctly embedded and can be read back.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import json
|
||||||
|
import os
|
||||||
|
import tempfile
|
||||||
|
from datetime import datetime
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
from mflux.config.config import Config
|
||||||
|
from mflux.config.model_config import ModelConfig
|
||||||
|
from mflux.models.flux.variants.txt2img.flux import Flux1
|
||||||
|
from mflux.post_processing.metadata_reader import MetadataReader
|
||||||
|
|
||||||
|
|
||||||
|
class TestMetadata:
|
||||||
|
"""Test suite for image metadata functionality."""
|
||||||
|
|
||||||
|
def test_metadata_complete(self):
|
||||||
|
"""
|
||||||
|
Comprehensive test that generates one image using img2img and verifies:
|
||||||
|
- EXIF metadata (all fields including img2img-specific ones)
|
||||||
|
- mflux-info command output formatting
|
||||||
|
- Creation timestamp validity
|
||||||
|
- Optional fields handling (both None and populated)
|
||||||
|
- Metadata reader handles missing files gracefully
|
||||||
|
"""
|
||||||
|
# Use a temporary file - don't keep it open to avoid conflicts
|
||||||
|
fd, temp_path = tempfile.mkstemp(suffix=".png")
|
||||||
|
os.close(fd) # Close the file descriptor immediately
|
||||||
|
output_path = Path(temp_path)
|
||||||
|
|
||||||
|
try:
|
||||||
|
# Use img2img to test image_path and image_strength metadata fields
|
||||||
|
reference_image = Path(__file__).parent.parent / "resources" / "reference_schnell.png"
|
||||||
|
|
||||||
|
# Generate a small, fast image with schnell img2img (256x256, 2 steps, quantized)
|
||||||
|
flux = Flux1(
|
||||||
|
model_config=ModelConfig.schnell(),
|
||||||
|
quantize=8,
|
||||||
|
)
|
||||||
|
|
||||||
|
config = Config(
|
||||||
|
num_inference_steps=2,
|
||||||
|
height=256,
|
||||||
|
width=256,
|
||||||
|
image_path=reference_image,
|
||||||
|
image_strength=0.3,
|
||||||
|
)
|
||||||
|
|
||||||
|
image = flux.generate_image(
|
||||||
|
seed=42,
|
||||||
|
prompt="A simple test image",
|
||||||
|
config=config,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Save with metadata (overwrite=True since mkstemp creates an empty file)
|
||||||
|
image.save(path=output_path, overwrite=True)
|
||||||
|
|
||||||
|
# =================================================================
|
||||||
|
# Test 1: Read metadata and verify structure
|
||||||
|
# =================================================================
|
||||||
|
metadata = MetadataReader.read_all_metadata(output_path)
|
||||||
|
|
||||||
|
assert metadata is not None, "Metadata should not be None"
|
||||||
|
assert "exif" in metadata, "EXIF metadata should be present"
|
||||||
|
|
||||||
|
exif = metadata["exif"]
|
||||||
|
|
||||||
|
# =================================================================
|
||||||
|
# Test 2: Core generation parameters
|
||||||
|
# =================================================================
|
||||||
|
assert exif.get("seed") == 42, "Seed should match"
|
||||||
|
assert exif.get("steps") == 2, "Steps should match"
|
||||||
|
assert exif.get("prompt") == "A simple test image", "Prompt should match"
|
||||||
|
assert exif.get("model") == "black-forest-labs/FLUX.1-schnell", "Model should match"
|
||||||
|
|
||||||
|
# =================================================================
|
||||||
|
# Test 3: Dimensions (NEW FEATURE)
|
||||||
|
# =================================================================
|
||||||
|
assert exif.get("width") == 256, "Width should be saved"
|
||||||
|
assert exif.get("height") == 256, "Height should be saved"
|
||||||
|
|
||||||
|
# =================================================================
|
||||||
|
# Test 4: Technical parameters
|
||||||
|
# =================================================================
|
||||||
|
assert exif.get("quantize") == 8, "Quantization should match"
|
||||||
|
assert exif.get("precision") is not None, "Precision should be set"
|
||||||
|
|
||||||
|
# =================================================================
|
||||||
|
# Test 5: MFLUX version (not hardcoded)
|
||||||
|
# =================================================================
|
||||||
|
assert exif.get("mflux_version") is not None, "MFLUX version should be present"
|
||||||
|
assert exif.get("mflux_version") != "unknown", "MFLUX version should not be unknown"
|
||||||
|
|
||||||
|
# =================================================================
|
||||||
|
# Test 6: Generation time
|
||||||
|
# =================================================================
|
||||||
|
assert exif.get("generation_time_seconds") is not None, "Generation time should be present"
|
||||||
|
assert isinstance(exif.get("generation_time_seconds"), (int, float)), "Generation time should be numeric"
|
||||||
|
assert exif.get("generation_time_seconds") > 0, "Generation time should be positive"
|
||||||
|
|
||||||
|
# =================================================================
|
||||||
|
# Test 7: Creation timestamp (NEW FEATURE)
|
||||||
|
# =================================================================
|
||||||
|
assert exif.get("created_at") is not None, "Creation timestamp should be present"
|
||||||
|
created_at = exif.get("created_at")
|
||||||
|
|
||||||
|
# Verify it's valid ISO format
|
||||||
|
dt = datetime.fromisoformat(created_at)
|
||||||
|
now = datetime.now()
|
||||||
|
time_diff = abs((now - dt).total_seconds())
|
||||||
|
assert time_diff < 120, f"Timestamp should be recent (within 120s), but was {time_diff}s ago"
|
||||||
|
|
||||||
|
# =================================================================
|
||||||
|
# Test 8: Img2img-specific fields are populated
|
||||||
|
# =================================================================
|
||||||
|
assert exif.get("image_path") is not None, "Image path should be set for img2img"
|
||||||
|
assert str(reference_image) in exif.get("image_path"), "Image path should contain reference image name"
|
||||||
|
assert exif.get("image_strength") == 0.3, "Image strength should match config"
|
||||||
|
|
||||||
|
# =================================================================
|
||||||
|
# Test 9: Optional fields are None when not used
|
||||||
|
# =================================================================
|
||||||
|
assert exif.get("lora_paths") is None, "LoRA paths should be None when not used"
|
||||||
|
assert exif.get("lora_scales") is None, "LoRA scales should be None when not used"
|
||||||
|
assert exif.get("controlnet_image_path") is None, "ControlNet path should be None when not used"
|
||||||
|
assert exif.get("negative_prompt") is None, "Negative prompt should be None when not used"
|
||||||
|
assert exif.get("guidance") is None, "Guidance should be None for schnell model"
|
||||||
|
|
||||||
|
# =================================================================
|
||||||
|
# Test 10: EXIF JSON is valid and parseable
|
||||||
|
# =================================================================
|
||||||
|
import piexif
|
||||||
|
from PIL import Image
|
||||||
|
|
||||||
|
img = Image.open(output_path)
|
||||||
|
exif_bytes = img.info.get("exif")
|
||||||
|
assert exif_bytes is not None, "EXIF bytes should be present"
|
||||||
|
|
||||||
|
exif_dict = piexif.load(exif_bytes)
|
||||||
|
user_comment = exif_dict.get("Exif", {}).get(piexif.ExifIFD.UserComment)
|
||||||
|
assert user_comment is not None, "UserComment should be present"
|
||||||
|
|
||||||
|
# Decode JSON
|
||||||
|
if user_comment.startswith(b"ASCII\x00\x00\x00"):
|
||||||
|
json_str = user_comment[8:].decode("utf-8")
|
||||||
|
else:
|
||||||
|
json_str = user_comment.decode("utf-8")
|
||||||
|
|
||||||
|
metadata_parsed = json.loads(json_str)
|
||||||
|
# Just verify it's valid JSON
|
||||||
|
assert metadata_parsed.get("prompt") == "A simple test image", "EXIF JSON should contain correct prompt"
|
||||||
|
|
||||||
|
# =================================================================
|
||||||
|
# Test 11: mflux-info command output
|
||||||
|
# =================================================================
|
||||||
|
from mflux.info import format_metadata
|
||||||
|
|
||||||
|
output = format_metadata(metadata)
|
||||||
|
assert "A simple test image" in output, "Prompt should be in output"
|
||||||
|
assert "MFLUX" in output, "MFLUX should be mentioned"
|
||||||
|
assert "42" in output, "Seed should be in output"
|
||||||
|
assert "256" in output, "Dimensions should be in output"
|
||||||
|
assert "Generation Time:" in output, "Generation time should be shown"
|
||||||
|
assert "Created:" in output, "Creation timestamp should be shown"
|
||||||
|
assert "Source Image:" in output, "Source image should be shown for img2img"
|
||||||
|
assert "Image Strength:" in output, "Image strength should be shown for img2img"
|
||||||
|
|
||||||
|
# =================================================================
|
||||||
|
# Test 12: Metadata reader handles nonexistent files
|
||||||
|
# =================================================================
|
||||||
|
nonexistent_metadata = MetadataReader.read_all_metadata(Path("/nonexistent/file.png"))
|
||||||
|
assert nonexistent_metadata.get("exif") is None, "Nonexistent file should return None for EXIF"
|
||||||
|
assert nonexistent_metadata.get("xmp") is None, "Nonexistent file should return None for XMP"
|
||||||
|
|
||||||
|
finally:
|
||||||
|
# Cleanup: Always remove the temporary file
|
||||||
|
if output_path.exists():
|
||||||
|
output_path.unlink()
|
||||||
Loading…
Reference in New Issue
Block a user