from pathlib import Path import mlx.core as mx import PIL.Image from mflux.community.concept_attention.attention_data import ConceptHeatmap from mflux.config.model_config import ModelConfig from mflux.utils.version_util import VersionUtil class GeneratedImage: def __init__( self, image: PIL.Image.Image, model_config: ModelConfig, seed: int, prompt: str, steps: int, guidance: float | None, precision: mx.Dtype, quantization: int, generation_time: float, lora_paths: list[str], lora_scales: list[float], controlnet_image_path: str | Path | None = None, controlnet_strength: float | None = None, image_path: str | Path | None = None, image_strength: float | None = None, masked_image_path: str | Path | None = None, depth_image_path: str | Path | None = None, redux_image_paths: list[str] | list[Path] | None = None, redux_image_strengths: list[float] | None = None, concept_heatmap: ConceptHeatmap | None = None, ): self.image = image self.model_config = model_config self.seed = seed self.prompt = prompt self.steps = steps self.guidance = guidance self.precision = precision self.quantization = quantization self.generation_time = generation_time self.lora_paths = lora_paths self.lora_scales = lora_scales self.controlnet_image_path = controlnet_image_path self.controlnet_strength = controlnet_strength self.image_path = image_path self.image_strength = image_strength self.masked_image_path = masked_image_path self.depth_image_path = depth_image_path self.redux_image_paths = redux_image_paths self.redux_image_strengths = redux_image_strengths self.concept_heatmap = concept_heatmap def get_right_half(self) -> "GeneratedImage": # Calculate the coordinates for the right half width, height = self.image.size right_half = self.image.crop((width // 2, 0, width, height)) # Create a new GeneratedImage with the right half and the same metadata return GeneratedImage( image=right_half, model_config=self.model_config, seed=self.seed, prompt=self.prompt, steps=self.steps, guidance=self.guidance, precision=self.precision, quantization=self.quantization, generation_time=self.generation_time, lora_paths=self.lora_paths, lora_scales=self.lora_scales, controlnet_image_path=self.controlnet_image_path, controlnet_strength=self.controlnet_strength, image_path=self.image_path, image_strength=self.image_strength, masked_image_path=self.masked_image_path, depth_image_path=self.depth_image_path, concept_heatmap=self.concept_heatmap, ) def save( self, path: str | Path, export_json_metadata: bool = False, overwrite: bool = False, ) -> None: from mflux.post_processing.image_util import ImageUtil ImageUtil.save_image(self.image, path, self._get_metadata(), export_json_metadata, overwrite) def save_with_heatmap( self, path: str | Path, export_json_metadata: bool = False, overwrite: bool = False, ) -> None: # Save the main image self.save(path=path, export_json_metadata=export_json_metadata, overwrite=overwrite) # Save the concept heatmap if available if self.concept_heatmap: file_path = Path(path) heatmap_path = file_path.with_stem(file_path.stem + "_heatmap") self.save_concept_heatmap(path=heatmap_path, export_json_metadata=export_json_metadata, overwrite=overwrite) def save_concept_heatmap( self, path: str | Path, export_json_metadata: bool = False, overwrite: bool = False, ) -> None: if self.concept_heatmap: from mflux.post_processing.image_util import ImageUtil ImageUtil.save_image( image=self.concept_heatmap.image, path=path, metadata=self.concept_heatmap.get_metadata(), export_json_metadata=export_json_metadata, overwrite=overwrite, ) else: raise ValueError("No concept heatmap available to save") def _format_redux_strengths(self) -> list[float] | None: if not self.redux_image_strengths: return None return [round(scale, 2) for scale in self.redux_image_strengths] def _get_metadata(self) -> dict: return { "mflux_version": VersionUtil.get_mflux_version(), "model": self.model_config.model_name, "base_model": str(self.model_config.base_model), "seed": self.seed, "steps": self.steps, "guidance": self.guidance if self.model_config.supports_guidance else None, "precision": str(self.precision), "quantize": self.quantization, "generation_time_seconds": round(self.generation_time, 2), "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, "image_path": str(self.image_path) if self.image_path else None, "image_strength": self.image_strength if self.image_path else None, "controlnet_image_path": str(self.controlnet_image_path) if self.controlnet_image_path else None, "controlnet_strength": round(self.controlnet_strength, 2) if self.controlnet_strength else None, "masked_image_path": str(self.masked_image_path) if self.masked_image_path else None, "depth_image_path": str(self.depth_image_path) if self.depth_image_path else None, "redux_image_paths": str(self.redux_image_paths) if self.redux_image_paths else None, "redux_image_strengths": self._format_redux_strengths(), "prompt": self.prompt, }