diff --git a/README.md b/README.md index d43cb26..2121644 100644 --- a/README.md +++ b/README.md @@ -46,6 +46,7 @@ Run the powerful [FLUX](https://blackforestlabs.ai/#get-flux) models from [Black * [Configuration details](#configuration-details) * [Memory issues](#memory-issues) * [Misc](#misc) +- [🧠 Concept Attention](#-concept-attention) - [🚧 Current limitations](#-current-limitations) - [πŸ’‘Workflow tips](#workflow-tips) - [πŸ”¬ Cool research / features to support](#-cool-research--features-to-support-) @@ -233,6 +234,22 @@ The `mflux-generate-redux` command uses most of the same arguments as `mflux-gen See the [Redux](#-redux) section for more details on this feature. +#### πŸ“œ Concept Attention Command-Line Arguments + +The `mflux-concept` command supports most of the same arguments as `mflux-generate`, with these specific parameters: + +- **`--concept`** (required, `str`): The concept prompt to use for attention visualization. This defines what specific aspect or element you want to analyze within your generated image. + +- **`--heatmap-layer-indices`** (optional, `[int]`, default: `[15, 16, 17, 18]`): Transformer layer indices to use for heatmap generation. These layers capture different levels of abstraction in the attention mechanism. + +- **`--heatmap-timesteps`** (optional, `[int]`, default: `None`): Timesteps to use for heatmap generation. If not specified, all timesteps are used. Lower timestep values focus on early stages of generation. + +The `mflux-concept-from-image` command uses most of the same arguments as `mflux-concept`, with this additional parameter: + +- **`--input-image-path`** (required, `str`): Path to a reference image for concept attention analysis. The model will analyze how the concept appears in this reference image and apply similar attention patterns to the generated image. + +See the [Concept Attention](#-concept-attention) section for more details on this feature. + #### πŸ“œ Fill Tool Command-Line Arguments The `mflux-generate-fill` command supports most of the same arguments as `mflux-generate`, with these specific parameters: @@ -1313,6 +1330,80 @@ The aim is to also gradually expand the scope of this feature with alternative t - The original fine-tuning script in [Diffusers](https://huggingface.co/docs/diffusers/v0.11.0/en/training/dreambooth) +--- + +### 🧠 Concept Attention + +The Concept Attention feature allows you to visualize and understand how FLUX models focus on specific concepts within your prompts during image generation. + +![concept_example_1](src/mflux/assets/concept_example_1.jpg) + +This implementation is based on the research paper ["ConceptAttention: Diffusion Transformers Learn Highly Interpretable Features"](https://arxiv.org/abs/2502.04320) by [Helbling et al.](https://github.com/helblazer811/ConceptAttention), which demonstrates that multi-modal diffusion transformers like FLUX learn highly interpretable representations that can be leveraged to generate precise concept localization maps. + +MFLUX provides two concept attention tools: + +1. **Text-based Concept Analysis** (`mflux-concept`): Analyzes attention patterns for a specific concept within your text prompt +2. **Image-guided Concept Analysis** (`mflux-concept-from-image`): Uses a reference image to guide concept attention analysis + +#### Text-based Concept Analysis + +This approach analyzes how the model attends to a specific concept mentioned in your prompt. The model generates an image while tracking attention patterns, then creates a heatmap showing where the model focused when processing your concept. + +##### Example + +```bash +mflux-concept \ + --prompt "A dragon on a hill" \ + --concept "dragon" \ + --model schnell \ + --steps 4 \ + --seed 9643208 \ + --height 720 \ + --width 1280 \ + --heatmap-layer-indices 15 16 17 18 \ + --heatmap-timesteps 0 1 2 3 \ + -q 4 +``` +This will generate the following image + +![concept_example_1](src/mflux/assets/concept_example_2.jpg) + +This command will generate: +- The main image based on your prompt +- A concept attention heatmap showing where the model focused on the "dragon" concept +- Both images are automatically saved with appropriate naming + +#### Image-guided Concept Analysis + +This approach uses a reference image to guide the concept attention analysis. The model analyzes how a concept appears in the reference image and applies similar attention patterns to generate a new image that maintains those conceptual relationships. + +##### Example + +```bash +mflux-concept-from-image \ + --model schnell \ + --input-image-path "puffin.png" \ + --prompt "Two puffins are perched on a grassy, flower-covered cliffside, with one appearing to call out while the other looks on silently against a blurred ocean backdrop" \ + --concept "bird" \ + --steps 4 \ + --height 720 \ + --width 1280 \ + --seed 4529717 \ + --heatmap-layer-indices 15 16 17 18 \ + --heatmap-timesteps 0 1 2 3 \ + -q 4 +``` + +This will generate the following image + +![concept_example_1](src/mflux/assets/concept_example_3.jpg) + +#### Advanced Configuration + +- **`--heatmap-layer-indices`**: Controls which transformer layers to analyze (default: 15-18). Different layers capture different levels of abstraction + +- **`--heatmap-timesteps`**: Specifies which denoising steps to include in the analysis (default: all steps). Lower timestep values focus on early generation stages where broad composition is determined. + --- ### 🚧 Current limitations @@ -1357,9 +1448,7 @@ See `uv run tools/rename_images.py --help` for full CLI usage help. - Use `--stepwise-image-output-dir` to save intermediate images at each denoising step, which can be useful for debugging or creating animations of the generation process ### πŸ”¬ Cool research / features to support -- [ ] [ConceptAttention](https://github.com/helblazer811/ConceptAttention) - [ ] [PuLID](https://github.com/ToTheBeginning/PuLID) -- [ ] [RF-Inversion](https://github.com/filipstrand/mflux/issues/91) - [ ] [catvton-flux](https://github.com/nftblackmagic/catvton-flux) ### πŸŒ±β€ Related projects diff --git a/pyproject.toml b/pyproject.toml index ea71101..e074fa6 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -61,6 +61,8 @@ mflux-generate-in-context = "mflux.generate_in_context:main" mflux-generate-fill = "mflux.generate_fill:main" mflux-generate-depth = "mflux.generate_depth:main" mflux-generate-redux = "mflux.generate_redux:main" +mflux-concept = "mflux.concept:main" +mflux-concept-from-image = "mflux.concept_from_image:main" mflux-save = "mflux.save:main" mflux-save-depth = "mflux.save_depth:main" mflux-train = "mflux.train:main" diff --git a/src/mflux/assets/concept_example_1.jpg b/src/mflux/assets/concept_example_1.jpg new file mode 100644 index 0000000..f77fef2 Binary files /dev/null and b/src/mflux/assets/concept_example_1.jpg differ diff --git a/src/mflux/assets/concept_example_2.jpg b/src/mflux/assets/concept_example_2.jpg new file mode 100644 index 0000000..40d6d5d Binary files /dev/null and b/src/mflux/assets/concept_example_2.jpg differ diff --git a/src/mflux/assets/concept_example_3.jpg b/src/mflux/assets/concept_example_3.jpg new file mode 100644 index 0000000..322ea6a Binary files /dev/null and b/src/mflux/assets/concept_example_3.jpg differ diff --git a/src/mflux/callbacks/callback_manager.py b/src/mflux/callbacks/callback_manager.py new file mode 100644 index 0000000..c00021f --- /dev/null +++ b/src/mflux/callbacks/callback_manager.py @@ -0,0 +1,77 @@ +from argparse import Namespace + +from mflux.callbacks.callback_registry import CallbackRegistry +from mflux.callbacks.instances.battery_saver import BatterySaver +from mflux.callbacks.instances.canny_saver import CannyImageSaver +from mflux.callbacks.instances.depth_saver import DepthImageSaver +from mflux.callbacks.instances.memory_saver import MemorySaver +from mflux.callbacks.instances.stepwise_handler import StepwiseHandler + + +class CallbackManager: + @staticmethod + def register_callbacks( + args: Namespace, + flux, + enable_canny_saver: bool = False, + enable_depth_saver: bool = False, + ) -> MemorySaver | None: + # Battery saver (always enabled) + CallbackManager._register_battery_saver(args) + + # VAE Tiling (if requested) + CallbackManager._register_vae_tiling(args, flux) + + # Specialized savers (based on flags) + if enable_canny_saver: + CallbackManager._register_canny_saver(args) + + if enable_depth_saver: + CallbackManager._register_depth_saver(args) + + # Stepwise handler (if requested) + CallbackManager._register_stepwise_handler(args, flux) + + # Memory saver (if requested) + return CallbackManager._register_memory_saver(args, flux) + + @staticmethod + def _register_battery_saver(args: Namespace) -> None: + battery_saver = BatterySaver(battery_percentage_stop_limit=args.battery_percentage_stop_limit) + CallbackRegistry.register_before_loop(battery_saver) + + @staticmethod + def _register_vae_tiling(args: Namespace, flux) -> None: + if args.vae_tiling: + flux.vae.decoder.enable_tiling = True + flux.vae.decoder.split_direction = args.vae_tiling_split + + @staticmethod + def _register_canny_saver(args: Namespace) -> None: + if hasattr(args, "controlnet_save_canny") and args.controlnet_save_canny: + canny_image_saver = CannyImageSaver(path=args.output) + CallbackRegistry.register_before_loop(canny_image_saver) + + @staticmethod + def _register_depth_saver(args: Namespace) -> None: + if hasattr(args, "save_depth_map") and args.save_depth_map: + depth_image_saver = DepthImageSaver(path=args.output) + CallbackRegistry.register_before_loop(depth_image_saver) + + @staticmethod + def _register_stepwise_handler(args: Namespace, flux) -> None: + if args.stepwise_image_output_dir: + handler = StepwiseHandler(flux=flux, output_dir=args.stepwise_image_output_dir) + CallbackRegistry.register_before_loop(handler) + CallbackRegistry.register_in_loop(handler) + CallbackRegistry.register_interrupt(handler) + + @staticmethod + def _register_memory_saver(args: Namespace, flux) -> MemorySaver | None: + memory_saver = None + if args.low_ram: + memory_saver = MemorySaver(flux=flux, keep_transformer=len(args.seed) > 1) + CallbackRegistry.register_before_loop(memory_saver) + CallbackRegistry.register_in_loop(memory_saver) + CallbackRegistry.register_after_loop(memory_saver) + return memory_saver diff --git a/src/mflux/community/concept_attention/__init__.py b/src/mflux/community/concept_attention/__init__.py new file mode 100644 index 0000000..397e2bb --- /dev/null +++ b/src/mflux/community/concept_attention/__init__.py @@ -0,0 +1 @@ +# Concept Attention transformer models diff --git a/src/mflux/community/concept_attention/attention_data.py b/src/mflux/community/concept_attention/attention_data.py new file mode 100644 index 0000000..13697fa --- /dev/null +++ b/src/mflux/community/concept_attention/attention_data.py @@ -0,0 +1,68 @@ +from dataclasses import dataclass +from pathlib import Path +from typing import List + +import mlx.core as mx +import PIL.Image + +from mflux.community.concept_attention.joint_transformer_block_concept import ( + LayerAttentionData, +) + + +@dataclass +class TimestepAttentionData: + t: int + attention_information: List[LayerAttentionData] + + def stack_img_attentions(self) -> mx.array: + return mx.stack([layer.img_attention for layer in self.attention_information], axis=0) + + def stack_concept_attentions(self) -> mx.array: + return mx.stack([layer.concept_attention for layer in self.attention_information], axis=0) + + +class GenerationAttentionData: + def __init__(self): + self.timestep_data: List[TimestepAttentionData] = [] + + def append(self, timestep_attention: TimestepAttentionData): + self.timestep_data.append(timestep_attention) + + def stack_all_img_attentions(self) -> mx.array: + timestep_stacks = [timestep.stack_img_attentions() for timestep in self.timestep_data] + return mx.stack(timestep_stacks, axis=0) + + def stack_all_concept_attentions(self) -> mx.array: + timestep_stacks = [timestep.stack_concept_attentions() for timestep in self.timestep_data] + return mx.stack(timestep_stacks, axis=0) + + +@dataclass +class ConceptHeatmap: + concept: str + image: PIL.Image.Image + layer_indices: List[int] + timesteps: List[int] + height: int + width: int + + 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( + image=self.image, + path=path, + metadata=self.get_metadata(), + export_json_metadata=export_json_metadata, + overwrite=overwrite, + ) + + def get_metadata(self) -> dict: + return { + "concept_prompt": self.concept, + "layer_indices": self.layer_indices, + "timesteps": self.timesteps, + "height": self.height, + "width": self.width, + } diff --git a/src/mflux/community/concept_attention/concept_util.py b/src/mflux/community/concept_attention/concept_util.py new file mode 100644 index 0000000..d5c1e6b --- /dev/null +++ b/src/mflux/community/concept_attention/concept_util.py @@ -0,0 +1,93 @@ +import matplotlib.pyplot as plt +import mlx.core as mx +import numpy as np +import PIL.Image + +from mflux.community.concept_attention.attention_data import ( + ConceptHeatmap, + GenerationAttentionData, +) + + +class ConceptUtil: + @staticmethod + def create_heatmap( + concept: str, + attention_data: GenerationAttentionData, + height: int, + width: int, + layer_indices: list[int], + timesteps: list[int] = list(range(4)), + ) -> ConceptHeatmap: + heatmap = ConceptUtil._compute_heatmap( + attention_data=attention_data, + height=height, + width=width, + layer_indices=layer_indices, + timesteps=timesteps, + ) + + colorized_image = ConceptUtil._to_heatmap_image( + heatmap=heatmap, + height=height, + width=width, + ) + + return ConceptHeatmap( + concept=concept, + image=colorized_image, + layer_indices=layer_indices, + timesteps=timesteps, + height=height, + width=width, + ) + + @staticmethod + def _to_heatmap_image(heatmap: mx.array, height: int, width: int) -> PIL.Image.Image: + concept_heatmaps_min = heatmap.min() + concept_heatmaps_max = heatmap.max() + colored_heatmaps = [] + for heatmap in heatmap: + heatmap = (heatmap - concept_heatmaps_min) / (concept_heatmaps_max - concept_heatmaps_min + 1e-8) + colored_heatmap = plt.get_cmap("plasma")(heatmap) + rgb_image = (colored_heatmap[:, :, :3] * 255).astype(np.uint8) + colored_heatmaps.append(rgb_image) + base_image = colored_heatmaps[0] + scaled_heatmap = ConceptUtil._pixel_perfect_resize(base_image, height, width) + return scaled_heatmap + + @staticmethod + def _pixel_perfect_resize(image_array: np.ndarray, target_height: int, target_width: int) -> PIL.Image.Image: + patch_h, patch_w, channels = image_array.shape + block_h = target_height // patch_h + block_w = target_width // patch_w + scaled_array = np.zeros((target_height, target_width, channels), dtype=np.uint8) + for i in range(patch_h): + for j in range(patch_w): + patch_color = image_array[i, j] + start_h = i * block_h + end_h = min(start_h + block_h, target_height) + start_w = j * block_w + end_w = min(start_w + block_w, target_width) + scaled_array[start_h:end_h, start_w:end_w] = patch_color + return PIL.Image.fromarray(scaled_array) + + @staticmethod + def _compute_heatmap( + attention_data: GenerationAttentionData, + height: int, + width: int, + layer_indices: list[int], + timesteps: list[int] = list(range(4)), + ) -> mx.array: + image_vectors = attention_data.stack_all_img_attentions() + concept_vectors = attention_data.stack_all_concept_attentions() + heatmaps = concept_vectors @ mx.transpose(image_vectors, (0, 1, 2, 4, 3)) + heatmaps = mx.softmax(heatmaps, axis=-2) + heatmaps = heatmaps[timesteps] + heatmaps = heatmaps[:, layer_indices] + heatmaps = mx.mean(heatmaps, axis=(0, 1)) + batch_dim, concept_dim, patch_dim = heatmaps.shape + heatmaps = mx.reshape(heatmaps, (batch_dim, concept_dim, height // 16, width // 16)) + heatmap = np.array(heatmaps)[0] + return heatmap diff --git a/src/mflux/community/concept_attention/flux_concept.py b/src/mflux/community/concept_attention/flux_concept.py new file mode 100644 index 0000000..6fd1171 --- /dev/null +++ b/src/mflux/community/concept_attention/flux_concept.py @@ -0,0 +1,181 @@ +import mlx.core as mx +from mlx import nn +from tqdm import tqdm + +from mflux.callbacks.callbacks import Callbacks +from mflux.community.concept_attention.attention_data import ( + GenerationAttentionData, +) +from mflux.community.concept_attention.concept_util import ConceptUtil +from mflux.community.concept_attention.transformer_concept import TransformerConcept +from mflux.config.config import Config +from mflux.config.model_config import ModelConfig +from mflux.config.runtime_config import RuntimeConfig +from mflux.error.exceptions import StopImageGenerationException +from mflux.flux.flux_initializer import FluxInitializer +from mflux.latent_creator.latent_creator import Img2Img, LatentCreator +from mflux.models.text_encoder.clip_encoder.clip_encoder import CLIPEncoder +from mflux.models.text_encoder.prompt_encoder import PromptEncoder +from mflux.models.text_encoder.t5_encoder.t5_encoder import T5Encoder +from mflux.models.vae.vae import VAE +from mflux.post_processing.array_util import ArrayUtil +from mflux.post_processing.generated_image import GeneratedImage +from mflux.post_processing.image_util import ImageUtil + + +class Flux1Concept(nn.Module): + vae: VAE + transformer: TransformerConcept + t5_text_encoder: T5Encoder + clip_text_encoder: CLIPEncoder + + def __init__( + self, + model_config: ModelConfig, + quantize: int | None = None, + local_path: str | None = None, + lora_paths: list[str] | None = None, + lora_scales: list[float] | None = None, + ): + super().__init__() + FluxInitializer.init_concept( + flux_model=self, + model_config=model_config, + quantize=quantize, + local_path=local_path, + lora_paths=lora_paths, + lora_scales=lora_scales, + ) + + def generate_image( + self, + seed: int, + prompt: str, + concept: str, + config: Config, + heatmap_layer_indices: list[int] | None = None, + heatmap_timesteps: list[int] | None = None, + ) -> GeneratedImage: + # 0. Create a new runtime config based on the model type and input parameters + config = RuntimeConfig(config, self.model_config) + time_steps = tqdm(range(config.init_time_step, config.num_inference_steps)) + + # 1. Create the initial latents + latents = LatentCreator.create_for_txt2img_or_img2img( + seed=seed, + height=config.height, + width=config.width, + img2img=Img2Img( + vae=self.vae, + image_path=config.image_path, + sigmas=config.sigmas, + init_time_step=config.init_time_step, + ), + ) + + # 2. Encode the main prompt + prompt_embeds, pooled_prompt_embeds = PromptEncoder.encode_prompt( + prompt=prompt, + prompt_cache=self.prompt_cache, + t5_tokenizer=self.t5_tokenizer, + clip_tokenizer=self.clip_tokenizer, + t5_text_encoder=self.t5_text_encoder, + clip_text_encoder=self.clip_text_encoder, + ) + + # 3. Encode the concept prompt + prompt_embeds_concept, pooled_prompt_embeds_concept = PromptEncoder.encode_prompt( + prompt=concept, + prompt_cache=self.prompt_cache, + t5_tokenizer=self.t5_tokenizer, + clip_tokenizer=self.clip_tokenizer, + t5_text_encoder=self.t5_text_encoder, + clip_text_encoder=self.clip_text_encoder, + ) + + # (Optional) Call subscribers for beginning of loop + Callbacks.before_loop( + seed=seed, + prompt=prompt, + latents=latents, + config=config, + ) + + attention_data = GenerationAttentionData() + for t in time_steps: + try: + # 4.t Predict the noise + noise, attention = self.transformer( + t=t, + config=config, + hidden_states=latents, + prompt_embeds=prompt_embeds, + prompt_embeds_concept=prompt_embeds_concept, + pooled_prompt_embeds=pooled_prompt_embeds, + pooled_prompt_embeds_concept=pooled_prompt_embeds_concept, + ) + attention_data.append(attention) + + # 5.t Take one denoise step + dt = config.sigmas[t + 1] - config.sigmas[t] + latents += noise * dt + + # (Optional) Call subscribers in-loop + Callbacks.in_loop( + t=t, + seed=seed, + prompt=prompt, + latents=latents, + config=config, + time_steps=time_steps, + ) + + # (Optional) Evaluate to enable progress tracking + mx.eval(latents) + + except KeyboardInterrupt: # noqa: PERF203 + Callbacks.interruption( + t=t, + seed=seed, + prompt=prompt, + latents=latents, + config=config, + time_steps=time_steps, + ) + raise StopImageGenerationException(f"Stopping image generation at step {t + 1}/{len(time_steps)}") + + # (Optional) Call subscribers after loop + Callbacks.after_loop( + seed=seed, + prompt=prompt, + latents=latents, + config=config, + ) + + # 6. Generate concept attention heatmap + concept_heatmap = ConceptUtil.create_heatmap( + concept=concept, + attention_data=attention_data, + height=config.height, + width=config.width, + layer_indices=heatmap_layer_indices or list(range(15, 19)), + timesteps=heatmap_timesteps or list(range(config.num_inference_steps)), + ) + + # 7. Decode the latent array and return the image + latents = ArrayUtil.unpack_latents(latents=latents, height=config.height, width=config.width) + decoded = self.vae.decode(latents) + + return ImageUtil.to_image( + decoded_latents=decoded, + config=config, + seed=seed, + prompt=prompt, + quantization=self.bits, + lora_paths=self.lora_paths, + lora_scales=self.lora_scales, + image_path=config.image_path, + image_strength=config.image_strength, + generation_time=time_steps.format_dict["elapsed"], + concept_heatmap=concept_heatmap, + ) diff --git a/src/mflux/community/concept_attention/flux_concept_from_image.py b/src/mflux/community/concept_attention/flux_concept_from_image.py new file mode 100644 index 0000000..00ce94f --- /dev/null +++ b/src/mflux/community/concept_attention/flux_concept_from_image.py @@ -0,0 +1,195 @@ +import mlx.core as mx +from mlx import nn +from tqdm import tqdm + +from mflux.callbacks.callbacks import Callbacks +from mflux.community.concept_attention.attention_data import GenerationAttentionData +from mflux.community.concept_attention.concept_util import ConceptUtil +from mflux.community.concept_attention.transformer_concept import TransformerConcept +from mflux.config.config import Config +from mflux.config.model_config import ModelConfig +from mflux.config.runtime_config import RuntimeConfig +from mflux.error.exceptions import StopImageGenerationException +from mflux.flux.flux_initializer import FluxInitializer +from mflux.latent_creator.latent_creator import LatentCreator +from mflux.models.text_encoder.clip_encoder.clip_encoder import CLIPEncoder +from mflux.models.text_encoder.prompt_encoder import PromptEncoder +from mflux.models.text_encoder.t5_encoder.t5_encoder import T5Encoder +from mflux.models.vae.vae import VAE +from mflux.post_processing.array_util import ArrayUtil +from mflux.post_processing.generated_image import GeneratedImage +from mflux.post_processing.image_util import ImageUtil + + +class FluxConceptFromImage(nn.Module): + vae: VAE + transformer: TransformerConcept + t5_text_encoder: T5Encoder + clip_text_encoder: CLIPEncoder + + def __init__( + self, + model_config: ModelConfig, + quantize: int | None = None, + local_path: str | None = None, + lora_paths: list[str] | None = None, + lora_scales: list[float] | None = None, + ): + super().__init__() + FluxInitializer.init_concept( + flux_model=self, + model_config=model_config, + quantize=quantize, + local_path=local_path, + lora_paths=lora_paths, + lora_scales=lora_scales, + ) + + def generate_image( + self, + seed: int, + prompt: str, + concept: str, + image_path: str, + config: Config, + heatmap_layer_indices: list[int] | None = None, + heatmap_timesteps: list[int] | None = None, + ) -> GeneratedImage: + # 0. Create a new runtime config based on the model type and input parameters + config = RuntimeConfig(config, self.model_config) + time_steps = tqdm(range(config.init_time_step, config.num_inference_steps)) + + # 1. Create the initial latents from the reference image + encoded_image = LatentCreator.encode_image( + vae=self.vae, + image_path=image_path, + height=config.height, + width=config.width, + ) + + # Create static noise for blending at each timestep + static_noise = LatentCreator.create( + seed=seed, + height=config.height, + width=config.width, + ) + + # Start with an appropriately noised version of the encoded image + latents = LatentCreator.add_noise_by_interpolation( + clean=ArrayUtil.pack_latents(latents=encoded_image, height=config.height, width=config.width), + noise=static_noise, + sigma=config.sigmas[config.init_time_step], + ) + + # 2. Encode the main prompt + prompt_embeds, pooled_prompt_embeds = PromptEncoder.encode_prompt( + prompt=prompt, + prompt_cache=self.prompt_cache, + t5_tokenizer=self.t5_tokenizer, + clip_tokenizer=self.clip_tokenizer, + t5_text_encoder=self.t5_text_encoder, + clip_text_encoder=self.clip_text_encoder, + ) + + # 3. Encode the concept prompt + prompt_embeds_concept, pooled_prompt_embeds_concept = PromptEncoder.encode_prompt( + prompt=concept, + prompt_cache=self.prompt_cache, + t5_tokenizer=self.t5_tokenizer, + clip_tokenizer=self.clip_tokenizer, + t5_text_encoder=self.t5_text_encoder, + clip_text_encoder=self.clip_text_encoder, + ) + + # (Optional) Call subscribers for beginning of loop + Callbacks.before_loop( + seed=seed, + prompt=prompt, + latents=latents, + config=config, + ) + + attention_data = GenerationAttentionData() + for t in time_steps: + try: + # 4.t Predict the noise (we don't use the noise, only the attention) + _, attention = self.transformer( + t=t, + config=config, + hidden_states=latents, + prompt_embeds=prompt_embeds, + prompt_embeds_concept=prompt_embeds_concept, + pooled_prompt_embeds=pooled_prompt_embeds, + pooled_prompt_embeds_concept=pooled_prompt_embeds_concept, + ) + attention_data.append(attention) + + # 5.t Follow reverse diffusion trajectory + latents = LatentCreator.add_noise_by_interpolation( + clean=ArrayUtil.pack_latents(latents=encoded_image, height=config.height, width=config.width), + noise=static_noise, + sigma=config.sigmas[t + 1], + ) + + # (Optional) Call subscribers in-loop + Callbacks.in_loop( + t=t, + seed=seed, + prompt=prompt, + latents=latents, + config=config, + time_steps=time_steps, + ) + + # Evaluate attention data to force MLX computation for progress tracking + mx.eval( + [layer.img_attention for layer in attention.attention_information] + + [layer.concept_attention for layer in attention.attention_information] + ) + + except KeyboardInterrupt: # noqa: PERF203 + Callbacks.interruption( + t=t, + seed=seed, + prompt=prompt, + latents=latents, + config=config, + time_steps=time_steps, + ) + raise StopImageGenerationException(f"Stopping image generation at step {t + 1}/{len(time_steps)}") + + # (Optional) Call subscribers after loop + Callbacks.after_loop( + seed=seed, + prompt=prompt, + latents=latents, + config=config, + ) + + # 6. Generate concept attention heatmap + concept_heatmap = ConceptUtil.create_heatmap( + concept=concept, + attention_data=attention_data, + height=config.height, + width=config.width, + layer_indices=heatmap_layer_indices or list(range(15, 19)), + timesteps=heatmap_timesteps or list(range(config.num_inference_steps)), + ) + + # 7. Decode the latent array and return the image + latents = ArrayUtil.unpack_latents(latents=latents, height=config.height, width=config.width) + decoded = self.vae.decode(latents) + + return ImageUtil.to_image( + decoded_latents=decoded, + config=config, + seed=seed, + prompt=prompt, + quantization=self.bits, + lora_paths=self.lora_paths, + lora_scales=self.lora_scales, + image_path=image_path, + image_strength=config.image_strength, + generation_time=time_steps.format_dict["elapsed"], + concept_heatmap=concept_heatmap, + ) diff --git a/src/mflux/community/concept_attention/joint_attention_concept.py b/src/mflux/community/concept_attention/joint_attention_concept.py new file mode 100644 index 0000000..4cc1696 --- /dev/null +++ b/src/mflux/community/concept_attention/joint_attention_concept.py @@ -0,0 +1,125 @@ +import mlx.core as mx +from mlx import nn + +from mflux.models.transformer.common.attention_utils import AttentionUtils + + +class JointAttentionConcept(nn.Module): + def __init__(self): + super().__init__() + self.head_dimension = 128 + self.batch_size = 1 + self.num_heads = 24 + + self.to_q = nn.Linear(3072, 3072) + self.to_k = nn.Linear(3072, 3072) + self.to_v = nn.Linear(3072, 3072) + self.to_out = [nn.Linear(3072, 3072)] + self.add_q_proj = nn.Linear(3072, 3072) + self.add_k_proj = nn.Linear(3072, 3072) + self.add_v_proj = nn.Linear(3072, 3072) + self.to_add_out = nn.Linear(3072, 3072) + self.norm_q = nn.RMSNorm(128) + self.norm_k = nn.RMSNorm(128) + self.norm_added_q = nn.RMSNorm(128) + self.norm_added_k = nn.RMSNorm(128) + + def __call__( + self, + hidden_states: mx.array, + encoder_hidden_states: mx.array, + encoder_hidden_states_concept: mx.array, + image_rotary_emb: mx.array, + image_rotary_emb_concept: mx.array, + ) -> tuple[mx.array, mx.array, mx.array, mx.array, mx.array]: + # Compute Q,K,V for hidden_states once (shared between both attention computations) + image_query, image_key, image_value = AttentionUtils.process_qkv( + hidden_states=hidden_states, + to_q=self.to_q, + to_k=self.to_k, + to_v=self.to_v, + norm_q=self.norm_q, + norm_k=self.norm_k, + num_heads=self.num_heads, + head_dim=self.head_dimension, + ) + + # 1: Regular joint attention + hidden_states_final, encoder_hidden_states_final, img_attn, _ = self._compute_joint_attention_optimized( + image_query=image_query, + image_key=image_key, + image_value=image_value, + encoder_hidden_states=encoder_hidden_states, + image_rotary_emb=image_rotary_emb, + ) + + # 2: Concept-specific attention + _, encoder_hidden_states_concept_final, _, concept_attn = self._compute_joint_attention_optimized( + image_query=image_query, + image_key=image_key, + image_value=image_value, + encoder_hidden_states=encoder_hidden_states_concept, + image_rotary_emb=image_rotary_emb_concept, + ) + + return ( + hidden_states_final, + encoder_hidden_states_final, + encoder_hidden_states_concept_final, + img_attn, + concept_attn, + ) + + def _compute_joint_attention_optimized( + self, + image_query: mx.array, + image_key: mx.array, + image_value: mx.array, + encoder_hidden_states: mx.array, + image_rotary_emb: mx.array, + ) -> tuple[mx.array, mx.array, mx.array, mx.array]: + # 1. Compute Q,K,V for encoder_hidden_states + enc_query, enc_key, enc_value = AttentionUtils.process_qkv( + hidden_states=encoder_hidden_states, + to_q=self.add_q_proj, + to_k=self.add_k_proj, + to_v=self.add_v_proj, + norm_q=self.norm_added_q, + norm_k=self.norm_added_k, + num_heads=self.num_heads, + head_dim=self.head_dimension, + ) + + # 2. Concatenate results (using pre-computed hidden states QKV) + joint_query = mx.concatenate([enc_query, image_query], axis=2) + joint_key = mx.concatenate([enc_key, image_key], axis=2) + joint_value = mx.concatenate([enc_value, image_value], axis=2) + + # 3. Apply rope to Q,K + joint_query, joint_key = AttentionUtils.apply_rope( + xq=joint_query, + xk=joint_key, + freqs_cis=image_rotary_emb, + ) + + # 4. Compute attention + joint_hidden_states = AttentionUtils.compute_attention( + query=joint_query, + key=joint_key, + value=joint_value, + batch_size=self.batch_size, + num_heads=self.num_heads, + head_dim=self.head_dimension, + ) + + # 5. Separate the results + encoder_output, hidden_states_output = ( + joint_hidden_states[:, : encoder_hidden_states.shape[1]], + joint_hidden_states[:, encoder_hidden_states.shape[1] :], + ) + + # 6. Project outputs + hidden_states_final = self.to_out[0](hidden_states_output) + encoder_hidden_states_final = self.to_add_out(encoder_output) + + return hidden_states_final, encoder_hidden_states_final, hidden_states_output, encoder_output diff --git a/src/mflux/community/concept_attention/joint_transformer_block_concept.py b/src/mflux/community/concept_attention/joint_transformer_block_concept.py new file mode 100644 index 0000000..a999f71 --- /dev/null +++ b/src/mflux/community/concept_attention/joint_transformer_block_concept.py @@ -0,0 +1,131 @@ +from dataclasses import dataclass + +import mlx.core as mx +from mlx import nn + +from mflux.community.concept_attention.joint_attention_concept import JointAttentionConcept +from mflux.models.transformer.ada_layer_norm_zero import AdaLayerNormZero +from mflux.models.transformer.feed_forward import FeedForward + + +@dataclass +class LayerAttentionData: + layer: int + img_attention: mx.array + concept_attention: mx.array + + +class JointTransformerBlockConcept(nn.Module): + def __init__(self, layer): + super().__init__() + self.layer = layer + self.norm1 = AdaLayerNormZero() + self.norm1_context = AdaLayerNormZero() + self.attn = JointAttentionConcept() + self.norm2 = nn.LayerNorm(dims=3072, eps=1e-6, affine=False) + self.norm2_context = nn.LayerNorm(dims=1536, eps=1e-6, affine=False) + self.ff = FeedForward(activation_function=nn.gelu) + self.ff_context = FeedForward(activation_function=nn.gelu_approx) + + def __call__( + self, + layer_idx: int, + hidden_states: mx.array, + encoder_hidden_states: mx.array, + encoder_hidden_states_concept: mx.array, + text_embeddings: mx.array, + text_embeddings_concept: mx.array, + rotary_embeddings: mx.array, + rotary_embeddings_concept: mx.array, + ) -> tuple[mx.array, mx.array, mx.array, LayerAttentionData]: + # 1a. Compute norm for hidden_states + norm_hidden_states, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.norm1( + hidden_states=hidden_states, + text_embeddings=text_embeddings, + ) + + # 1b. Compute norm for encoder_hidden_states + norm_encoder_hidden_states, c_gate_msa, c_shift_mlp, c_scale_mlp, c_gate_mlp = self.norm1_context( + hidden_states=encoder_hidden_states, + text_embeddings=text_embeddings, + ) + + # 1c. Compute norm for encoder_hidden_states_concept + norm_encoder_hidden_states_concept, c_gate_msa_concept, c_shift_mlp_concept, c_scale_mlp_concept, c_gate_mlp_concept = self.norm1_context( + hidden_states=encoder_hidden_states_concept, + text_embeddings=text_embeddings_concept, + ) # fmt: off + + # 2. Compute attention + attn_output, context_attn_output, context_attn_output_concept, img_attn, concept_attn = self.attn( + hidden_states=norm_hidden_states, + encoder_hidden_states=norm_encoder_hidden_states, + encoder_hidden_states_concept=norm_encoder_hidden_states_concept, + image_rotary_emb=rotary_embeddings, + image_rotary_emb_concept=rotary_embeddings_concept, + ) + + # 3a. Apply norm and feed forward for hidden states + hidden_states = JointTransformerBlockConcept._apply_norm_and_feed_forward( + hidden_states=hidden_states, + attn_output=attn_output, + gate_mlp=gate_mlp, + gate_msa=gate_msa, + scale_mlp=scale_mlp, + shift_mlp=shift_mlp, + norm_layer=self.norm2, + ff_layer=self.ff, + ) + + # 3b. Apply norm and feed forward for encoder hidden states + encoder_hidden_states = JointTransformerBlockConcept._apply_norm_and_feed_forward( + hidden_states=encoder_hidden_states, + attn_output=context_attn_output, + gate_mlp=c_gate_mlp, + gate_msa=c_gate_msa, + scale_mlp=c_scale_mlp, + shift_mlp=c_shift_mlp, + norm_layer=self.norm2_context, + ff_layer=self.ff_context, + ) + + # 3c. Apply norm and feed forward for concept encoder hidden states + encoder_hidden_states_concept = JointTransformerBlockConcept._apply_norm_and_feed_forward( + hidden_states=encoder_hidden_states_concept, + attn_output=context_attn_output_concept, + gate_mlp=c_gate_mlp_concept, + gate_msa=c_gate_msa_concept, + scale_mlp=c_scale_mlp_concept, + shift_mlp=c_shift_mlp_concept, + norm_layer=self.norm2_context, + ff_layer=self.ff_context, + ) + + # 4. Package attention data with layer information + layer_attention_data = LayerAttentionData( + layer=layer_idx, + img_attention=img_attn, + concept_attention=concept_attn, + ) + + return encoder_hidden_states, hidden_states, encoder_hidden_states_concept, layer_attention_data + + @staticmethod + def _apply_norm_and_feed_forward( + hidden_states: mx.array, + attn_output: mx.array, + gate_mlp: mx.array, + gate_msa: mx.array, + scale_mlp: mx.array, + shift_mlp: mx.array, + norm_layer: nn.Module, + ff_layer: nn.Module, + ) -> mx.array: + attn_output = mx.expand_dims(gate_msa, axis=1) * attn_output + hidden_states = hidden_states + attn_output + norm_hidden_states = norm_layer(hidden_states) + norm_hidden_states = norm_hidden_states * (1 + scale_mlp[:, None]) + shift_mlp[:, None] + ff_output = ff_layer(norm_hidden_states) + ff_output = mx.expand_dims(gate_mlp, axis=1) * ff_output + hidden_states = hidden_states + ff_output + return hidden_states diff --git a/src/mflux/community/concept_attention/transformer_concept.py b/src/mflux/community/concept_attention/transformer_concept.py new file mode 100644 index 0000000..a3d22c5 --- /dev/null +++ b/src/mflux/community/concept_attention/transformer_concept.py @@ -0,0 +1,96 @@ +import mlx.core as mx +from mlx import nn + +from mflux.community.concept_attention.attention_data import ( + TimestepAttentionData, +) +from mflux.community.concept_attention.joint_transformer_block_concept import ( + JointTransformerBlockConcept, +) +from mflux.config.model_config import ModelConfig +from mflux.config.runtime_config import RuntimeConfig +from mflux.models.transformer.ada_layer_norm_continuous import ( + AdaLayerNormContinuous, +) +from mflux.models.transformer.embed_nd import EmbedND +from mflux.models.transformer.single_transformer_block import ( + SingleTransformerBlock, +) +from mflux.models.transformer.time_text_embed import TimeTextEmbed +from mflux.models.transformer.transformer import Transformer + + +class TransformerConcept(nn.Module): + def __init__( + self, + model_config: ModelConfig, + num_transformer_blocks: int = 19, + num_single_transformer_blocks: int = 38, + ): + super().__init__() + self.pos_embed = EmbedND() + self.x_embedder = nn.Linear(model_config.x_embedder_input_dim(), 3072) + self.time_text_embed = TimeTextEmbed(model_config=model_config) + self.context_embedder = nn.Linear(4096, 3072) + self.transformer_blocks = [JointTransformerBlockConcept(i) for i in range(num_transformer_blocks)] + self.single_transformer_blocks = [SingleTransformerBlock(i) for i in range(num_single_transformer_blocks)] + self.norm_out = AdaLayerNormContinuous(3072, 3072) + self.proj_out = nn.Linear(3072, 64) + + def __call__( + self, + t: int, + config: RuntimeConfig, + hidden_states: mx.array, + prompt_embeds: mx.array, + prompt_embeds_concept: mx.array, + pooled_prompt_embeds: mx.array, + pooled_prompt_embeds_concept: mx.array, + ) -> tuple[mx.array, TimestepAttentionData]: + # 1. Create embeddings + hidden_states = self.x_embedder(hidden_states) + encoder_hidden_states = self.context_embedder(prompt_embeds) + encoder_hidden_states_concept = self.context_embedder(prompt_embeds_concept) + text_embeddings = Transformer.compute_text_embeddings(t, pooled_prompt_embeds, self.time_text_embed, config) # fmt: off + text_embeddings_concept = Transformer.compute_text_embeddings(t, pooled_prompt_embeds_concept, self.time_text_embed, config) # fmt: off + image_rotary_embeddings = Transformer.compute_rotary_embeddings(prompt_embeds, self.pos_embed, config) # fmt: off + image_rotary_embeddings_concept = Transformer.compute_rotary_embeddings(prompt_embeds_concept, self.pos_embed, config) # fmt: off + + # 2. Run the joint transformer blocks + attention_information = [] + for idx, block in enumerate(self.transformer_blocks): + encoder_hidden_states, hidden_states, encoder_hidden_states_concept, attn = block( + layer_idx=idx, + hidden_states=hidden_states, + encoder_hidden_states=encoder_hidden_states, + encoder_hidden_states_concept=encoder_hidden_states_concept, + text_embeddings=text_embeddings, + text_embeddings_concept=text_embeddings_concept, + rotary_embeddings=image_rotary_embeddings, + rotary_embeddings_concept=image_rotary_embeddings_concept, + ) + attention_information.append(attn) + + # 3. Concat the hidden states + hidden_states = mx.concatenate([encoder_hidden_states, hidden_states], axis=1) + + # 4. Run the single transformer blocks + for idx, block in enumerate(self.single_transformer_blocks): + hidden_states = block( + hidden_states=hidden_states, + text_embeddings=text_embeddings, + rotary_embeddings=image_rotary_embeddings, + ) + + # 5. Project the final output + hidden_states = hidden_states[:, encoder_hidden_states.shape[1] :, ...] + hidden_states = self.norm_out(hidden_states, text_embeddings) + hidden_states = self.proj_out(hidden_states) + + # 6. Create timestep attention data structure + timestep_attention = TimestepAttentionData( + t=t, + attention_information=attention_information, + ) + + return hidden_states, timestep_attention diff --git a/src/mflux/concept.py b/src/mflux/concept.py new file mode 100644 index 0000000..7e40fb7 --- /dev/null +++ b/src/mflux/concept.py @@ -0,0 +1,61 @@ +from mflux import Config, ModelConfig, StopImageGenerationException +from mflux.callbacks.callback_manager import CallbackManager +from mflux.community.concept_attention.flux_concept import Flux1Concept +from mflux.error.exceptions import PromptFileReadError +from mflux.ui.cli.parsers import CommandLineParser +from mflux.ui.prompt_utils import get_effective_prompt + + +def main(): + # 0. Parse command line arguments + parser = CommandLineParser(description="Generate an image with concept attention based on a prompt and concept.") + parser.add_general_arguments() + parser.add_model_arguments(require_model_arg=False) + parser.add_lora_arguments() + parser.add_image_generator_arguments(supports_metadata_config=True) + parser.add_image_to_image_arguments(required=False) + parser.add_output_arguments() + parser.add_concept_attention_arguments() + args = parser.parse_args() + + # 1. Load the concept attention model + flux = Flux1Concept( + model_config=ModelConfig.from_name(model_name=args.model, base_model=args.base_model), + quantize=args.quantize, + local_path=args.path, + lora_paths=args.lora_paths, + lora_scales=args.lora_scales, + ) + + # 2. Register callbacks + memory_saver = CallbackManager.register_callbacks(args=args, flux=flux) + + try: + for seed in args.seed: + # 3. Generate an image for each seed value + image = flux.generate_image( + seed=seed, + prompt=get_effective_prompt(args), + concept=args.concept, + heatmap_timesteps=args.heatmap_timesteps, + heatmap_layer_indices=args.heatmap_layer_indices, + config=Config( + num_inference_steps=args.steps, + height=args.height, + width=args.width, + guidance=args.guidance, + image_path=args.image_path, + image_strength=args.image_strength, + ), + ) + # 4. Save the image and heatmap + image.save_with_heatmap(path=args.output.format(seed=seed), export_json_metadata=args.metadata) + except (StopImageGenerationException, PromptFileReadError) as exc: + print(exc) + finally: + if memory_saver: + print(memory_saver.memory_stats()) + + +if __name__ == "__main__": + main() diff --git a/src/mflux/concept_from_image.py b/src/mflux/concept_from_image.py new file mode 100644 index 0000000..44a3280 --- /dev/null +++ b/src/mflux/concept_from_image.py @@ -0,0 +1,62 @@ +from mflux import Config, ModelConfig, StopImageGenerationException +from mflux.callbacks.callback_manager import CallbackManager +from mflux.community.concept_attention.flux_concept_from_image import FluxConceptFromImage +from mflux.error.exceptions import PromptFileReadError +from mflux.ui.cli.parsers import CommandLineParser +from mflux.ui.prompt_utils import get_effective_prompt + + +def main(): + # 0. Parse command line arguments + parser = CommandLineParser(description="Generate an image with concept attention based on a prompt, concept, and reference image.") # fmt: off + parser.add_general_arguments() + parser.add_model_arguments(require_model_arg=False) + parser.add_lora_arguments() + parser.add_image_generator_arguments(supports_metadata_config=True) + parser.add_image_to_image_arguments(required=False) + parser.add_output_arguments() + parser.add_concept_from_image_arguments() + args = parser.parse_args() + + # 1. Load the concept attention model + flux = FluxConceptFromImage( + model_config=ModelConfig.from_name(model_name=args.model, base_model=args.base_model), + quantize=args.quantize, + local_path=args.path, + lora_paths=args.lora_paths, + lora_scales=args.lora_scales, + ) + + # 2. Register callbacks + memory_saver = CallbackManager.register_callbacks(args=args, flux=flux) + + try: + for seed in args.seed: + # 3. Generate an image for each seed value + image = flux.generate_image( + seed=seed, + prompt=get_effective_prompt(args), + concept=args.concept, + image_path=str(args.input_image_path), + heatmap_timesteps=args.heatmap_timesteps, + heatmap_layer_indices=args.heatmap_layer_indices, + config=Config( + num_inference_steps=args.steps, + height=args.height, + width=args.width, + guidance=args.guidance, + image_path=args.image_path, + image_strength=args.image_strength, + ), + ) + # 4. Save the image and heatmap + image.save_with_heatmap(path=args.output.format(seed=seed), export_json_metadata=args.metadata) + except (StopImageGenerationException, PromptFileReadError) as exc: + print(exc) + finally: + if memory_saver: + print(memory_saver.memory_stats()) + + +if __name__ == "__main__": + main() diff --git a/src/mflux/flux/flux_initializer.py b/src/mflux/flux/flux_initializer.py index 6a71362..a0ec991 100644 --- a/src/mflux/flux/flux_initializer.py +++ b/src/mflux/flux/flux_initializer.py @@ -29,6 +29,7 @@ class FluxInitializer: lora_scales: list[float] | None = None, lora_names: list[str] | None = None, lora_repo_id: str | None = None, + custom_transformer=None, ) -> None: # 0. Set paths, configs, and prompt_cache for later lora_paths = lora_paths or [] @@ -57,11 +58,17 @@ class FluxInitializer: # 3. Initialize all models flux_model.vae = VAE() - flux_model.transformer = Transformer( - model_config=model_config, - num_transformer_blocks=weights.num_transformer_blocks(), - num_single_transformer_blocks=weights.num_single_transformer_blocks(), - ) + + # Use custom transformer if provided, otherwise create default + if custom_transformer is not None: + flux_model.transformer = custom_transformer + else: + flux_model.transformer = Transformer( + model_config=model_config, + num_transformer_blocks=weights.num_transformer_blocks(), + num_single_transformer_blocks=weights.num_single_transformer_blocks(), + ) + flux_model.t5_text_encoder = T5Encoder() flux_model.clip_text_encoder = CLIPEncoder() @@ -188,3 +195,43 @@ class FluxInitializer: weights=weights_controlnet, transformer_controlnet=flux_model.transformer_controlnet, ) + + @staticmethod + def init_concept( + flux_model, + model_config: ModelConfig, + quantize: int | None, + local_path: str | None, + lora_paths: list[str] | None = None, + lora_scales: list[float] | None = None, + lora_names: list[str] | None = None, + lora_repo_id: str | None = None, + ): + # Import here to avoid circular dependency + from mflux.community.concept_attention.transformer_concept import TransformerConcept + + # 1. Load weights first to get transformer dimensions + weights = WeightHandler.load_regular_weights( + repo_id=model_config.model_name, + local_path=local_path, + ) + + # 2. Create custom TransformerConcept + custom_transformer = TransformerConcept( + model_config=model_config, + num_transformer_blocks=weights.num_transformer_blocks(), + num_single_transformer_blocks=weights.num_single_transformer_blocks(), + ) + + # 3. Use the improved FluxInitializer with custom transformer + FluxInitializer.init( + flux_model=flux_model, + model_config=model_config, + quantize=quantize, + local_path=local_path, + lora_paths=lora_paths, + lora_scales=lora_scales, + lora_names=lora_names, + lora_repo_id=lora_repo_id, + custom_transformer=custom_transformer, + ) diff --git a/src/mflux/generate.py b/src/mflux/generate.py index 432e4ff..76744bc 100644 --- a/src/mflux/generate.py +++ b/src/mflux/generate.py @@ -1,10 +1,5 @@ -from argparse import Namespace - from mflux import Config, Flux1, ModelConfig, StopImageGenerationException -from mflux.callbacks.callback_registry import CallbackRegistry -from mflux.callbacks.instances.battery_saver import BatterySaver -from mflux.callbacks.instances.memory_saver import MemorySaver -from mflux.callbacks.instances.stepwise_handler import StepwiseHandler +from mflux.callbacks.callback_manager import CallbackManager from mflux.error.exceptions import PromptFileReadError from mflux.ui.cli.parsers import CommandLineParser from mflux.ui.prompt_utils import get_effective_prompt @@ -31,7 +26,7 @@ def main(): ) # 2. Register callbacks - memory_saver = _register_callbacks(args=args, flux=flux) + memory_saver = CallbackManager.register_callbacks(args=args, flux=flux) try: for seed in args.seed: @@ -57,32 +52,5 @@ def main(): print(memory_saver.memory_stats()) -def _register_callbacks(args: Namespace, flux: Flux1) -> MemorySaver | None: - # Battery saver - battery_saver = BatterySaver(battery_percentage_stop_limit=args.battery_percentage_stop_limit) - CallbackRegistry.register_before_loop(battery_saver) - - # VAE Tiling - if args.vae_tiling: - flux.vae.decoder.enable_tiling = True - flux.vae.decoder.split_direction = args.vae_tiling_split - - # Stepwise Handler - if args.stepwise_image_output_dir: - handler = StepwiseHandler(flux=flux, output_dir=args.stepwise_image_output_dir) - CallbackRegistry.register_before_loop(handler) - CallbackRegistry.register_in_loop(handler) - CallbackRegistry.register_interrupt(handler) - - # Memory Saver - memory_saver = None - if args.low_ram: - memory_saver = MemorySaver(flux=flux, keep_transformer=len(args.seed) > 1) - CallbackRegistry.register_before_loop(memory_saver) - CallbackRegistry.register_in_loop(memory_saver) - CallbackRegistry.register_after_loop(memory_saver) - return memory_saver - - if __name__ == "__main__": main() diff --git a/src/mflux/generate_controlnet.py b/src/mflux/generate_controlnet.py index badec90..6a92e9e 100644 --- a/src/mflux/generate_controlnet.py +++ b/src/mflux/generate_controlnet.py @@ -1,11 +1,5 @@ -from argparse import Namespace - from mflux import Config, Flux1Controlnet, ModelConfig, StopImageGenerationException -from mflux.callbacks.callback_registry import CallbackRegistry -from mflux.callbacks.instances.battery_saver import BatterySaver -from mflux.callbacks.instances.canny_saver import CannyImageSaver -from mflux.callbacks.instances.memory_saver import MemorySaver -from mflux.callbacks.instances.stepwise_handler import StepwiseHandler +from mflux.callbacks.callback_manager import CallbackManager from mflux.error.exceptions import PromptFileReadError from mflux.ui.cli.parsers import CommandLineParser from mflux.ui.prompt_utils import get_effective_prompt @@ -32,7 +26,7 @@ def main(): ) # 2. Register callbacks - memory_saver = _register_callbacks(args=args, flux=flux) + memory_saver = CallbackManager.register_callbacks(args=args, flux=flux, enable_canny_saver=True) try: for seed in args.seed: @@ -65,37 +59,5 @@ def _get_controlnet_model_config(model_name: str) -> ModelConfig: return ModelConfig.dev_controlnet_canny() -def _register_callbacks(args: Namespace, flux: Flux1Controlnet) -> MemorySaver | None: - # Battery saver - battery_saver = BatterySaver(battery_percentage_stop_limit=args.battery_percentage_stop_limit) - CallbackRegistry.register_before_loop(battery_saver) - - # VAE Tiling - if args.vae_tiling: - flux.vae.decoder.enable_tiling = True - flux.vae.decoder.split_direction = args.vae_tiling_split - - # Canny Image Saver - if args.controlnet_save_canny: - canny_image_saver = CannyImageSaver(path=args.output) - CallbackRegistry.register_before_loop(canny_image_saver) - - # Stepwise Handler - if args.stepwise_image_output_dir: - handler = StepwiseHandler(flux=flux, output_dir=args.stepwise_image_output_dir) - CallbackRegistry.register_before_loop(handler) - CallbackRegistry.register_in_loop(handler) - CallbackRegistry.register_interrupt(handler) - - # Memory Saver - memory_saver = None - if args.low_ram: - memory_saver = MemorySaver(flux=flux, keep_transformer=len(args.seed) > 1) - CallbackRegistry.register_before_loop(memory_saver) - CallbackRegistry.register_in_loop(memory_saver) - CallbackRegistry.register_after_loop(memory_saver) - return memory_saver - - if __name__ == "__main__": main() diff --git a/src/mflux/generate_depth.py b/src/mflux/generate_depth.py index 4447fbb..a828af1 100644 --- a/src/mflux/generate_depth.py +++ b/src/mflux/generate_depth.py @@ -1,11 +1,5 @@ -from argparse import Namespace - from mflux import Config, StopImageGenerationException -from mflux.callbacks.callback_registry import CallbackRegistry -from mflux.callbacks.instances.battery_saver import BatterySaver -from mflux.callbacks.instances.depth_saver import DepthImageSaver -from mflux.callbacks.instances.memory_saver import MemorySaver -from mflux.callbacks.instances.stepwise_handler import StepwiseHandler +from mflux.callbacks.callback_manager import CallbackManager from mflux.error.exceptions import PromptFileReadError from mflux.flux_tools.depth.flux_depth import Flux1Depth from mflux.ui.cli.parsers import CommandLineParser @@ -36,7 +30,7 @@ def main(): ) # 2. Register callbacks - memory_saver = _register_callbacks(args=args, flux=flux) + memory_saver = CallbackManager.register_callbacks(args=args, flux=flux, enable_depth_saver=True) try: for seed in args.seed: @@ -63,37 +57,5 @@ def main(): print(memory_saver.memory_stats()) -def _register_callbacks(args: Namespace, flux: Flux1Depth) -> MemorySaver | None: - # Battery saver - battery_saver = BatterySaver(battery_percentage_stop_limit=args.battery_percentage_stop_limit) - CallbackRegistry.register_before_loop(battery_saver) - - # VAE Tiling - if args.vae_tiling: - flux.vae.decoder.enable_tiling = True - flux.vae.decoder.split_direction = args.vae_tiling_split - - # Depth Image Saver - if args.save_depth_map: - depth_image_saver = DepthImageSaver(path=args.output) - CallbackRegistry.register_before_loop(depth_image_saver) - - # Stepwise Handler - if args.stepwise_image_output_dir: - handler = StepwiseHandler(flux=flux, output_dir=args.stepwise_image_output_dir) - CallbackRegistry.register_before_loop(handler) - CallbackRegistry.register_in_loop(handler) - CallbackRegistry.register_interrupt(handler) - - # Memory Saver - memory_saver = None - if args.low_ram: - memory_saver = MemorySaver(flux=flux, keep_transformer=len(args.seed) > 1) - CallbackRegistry.register_before_loop(memory_saver) - CallbackRegistry.register_in_loop(memory_saver) - CallbackRegistry.register_after_loop(memory_saver) - return memory_saver - - if __name__ == "__main__": main() diff --git a/src/mflux/generate_fill.py b/src/mflux/generate_fill.py index 84dcbc8..85cfbde 100644 --- a/src/mflux/generate_fill.py +++ b/src/mflux/generate_fill.py @@ -1,10 +1,5 @@ -from argparse import Namespace - from mflux import Config, StopImageGenerationException -from mflux.callbacks.callback_registry import CallbackRegistry -from mflux.callbacks.instances.battery_saver import BatterySaver -from mflux.callbacks.instances.memory_saver import MemorySaver -from mflux.callbacks.instances.stepwise_handler import StepwiseHandler +from mflux.callbacks.callback_manager import CallbackManager from mflux.error.exceptions import PromptFileReadError from mflux.flux_tools.fill.flux_fill import Flux1Fill from mflux.ui.cli.parsers import CommandLineParser @@ -35,7 +30,7 @@ def main(): ) # 2. Register callbacks - memory_saver = _register_callbacks(args=args, flux=flux) + memory_saver = CallbackManager.register_callbacks(args=args, flux=flux) try: for seed in args.seed: @@ -62,32 +57,5 @@ def main(): print(memory_saver.memory_stats()) -def _register_callbacks(args: Namespace, flux: Flux1Fill) -> MemorySaver | None: - # Battery saver - battery_saver = BatterySaver(battery_percentage_stop_limit=args.battery_percentage_stop_limit) - CallbackRegistry.register_before_loop(battery_saver) - - # VAE Tiling - if args.vae_tiling: - flux.vae.decoder.enable_tiling = True - flux.vae.decoder.split_direction = args.vae_tiling_split - - # Stepwise Handler - if args.stepwise_image_output_dir: - handler = StepwiseHandler(flux=flux, output_dir=args.stepwise_image_output_dir) - CallbackRegistry.register_before_loop(handler) - CallbackRegistry.register_in_loop(handler) - CallbackRegistry.register_interrupt(handler) - - # Memory Saver - memory_saver = None - if args.low_ram: - memory_saver = MemorySaver(flux=flux, keep_transformer=len(args.seed) > 1) - CallbackRegistry.register_before_loop(memory_saver) - CallbackRegistry.register_in_loop(memory_saver) - CallbackRegistry.register_after_loop(memory_saver) - return memory_saver - - if __name__ == "__main__": main() diff --git a/src/mflux/generate_in_context.py b/src/mflux/generate_in_context.py index 0c6584e..dc518ef 100644 --- a/src/mflux/generate_in_context.py +++ b/src/mflux/generate_in_context.py @@ -1,11 +1,7 @@ -from argparse import Namespace from pathlib import Path from mflux import Config, StopImageGenerationException -from mflux.callbacks.callback_registry import CallbackRegistry -from mflux.callbacks.instances.battery_saver import BatterySaver -from mflux.callbacks.instances.memory_saver import MemorySaver -from mflux.callbacks.instances.stepwise_handler import StepwiseHandler +from mflux.callbacks.callback_manager import CallbackManager from mflux.community.in_context_lora.flux_in_context_lora import Flux1InContextLoRA from mflux.community.in_context_lora.in_context_loras import LORA_REPO_ID, get_lora_filename from mflux.config.model_config import ModelConfig @@ -42,7 +38,7 @@ def main(): ) # 2. Register callbacks - memory_saver = _register_callbacks(args=args, flux=flux) + memory_saver = CallbackManager.register_callbacks(args=args, flux=flux) try: for seed in args.seed: @@ -70,32 +66,5 @@ def main(): print(memory_saver.memory_stats()) -def _register_callbacks(args: Namespace, flux: Flux1InContextLoRA) -> MemorySaver | None: - # Battery saver - battery_saver = BatterySaver(battery_percentage_stop_limit=args.battery_percentage_stop_limit) - CallbackRegistry.register_before_loop(battery_saver) - - # VAE Tiling - if args.vae_tiling: - flux.vae.decoder.enable_tiling = True - flux.vae.decoder.split_direction = "vertical" - - # Stepwise Handler - if args.stepwise_image_output_dir: - handler = StepwiseHandler(flux=flux, output_dir=args.stepwise_image_output_dir) - CallbackRegistry.register_before_loop(handler) - CallbackRegistry.register_in_loop(handler) - CallbackRegistry.register_interrupt(handler) - - # Memory Saver - memory_saver = None - if args.low_ram: - memory_saver = MemorySaver(flux=flux, keep_transformer=len(args.seed) > 1) - CallbackRegistry.register_before_loop(memory_saver) - CallbackRegistry.register_in_loop(memory_saver) - CallbackRegistry.register_after_loop(memory_saver) - return memory_saver - - if __name__ == "__main__": main() diff --git a/src/mflux/generate_redux.py b/src/mflux/generate_redux.py index e2340ba..6459b9e 100644 --- a/src/mflux/generate_redux.py +++ b/src/mflux/generate_redux.py @@ -1,11 +1,7 @@ -from argparse import Namespace from pathlib import Path from mflux import Config, ModelConfig, StopImageGenerationException -from mflux.callbacks.callback_registry import CallbackRegistry -from mflux.callbacks.instances.battery_saver import BatterySaver -from mflux.callbacks.instances.memory_saver import MemorySaver -from mflux.callbacks.instances.stepwise_handler import StepwiseHandler +from mflux.callbacks.callback_manager import CallbackManager from mflux.error.exceptions import PromptFileReadError from mflux.flux_tools.redux.flux_redux import Flux1Redux from mflux.ui.cli.parsers import CommandLineParser @@ -39,7 +35,7 @@ def main(): ) # 2. Register callbacks - memory_saver = _register_callbacks(args=args, flux=flux) + memory_saver = CallbackManager.register_callbacks(args=args, flux=flux) try: for seed in args.seed: @@ -66,33 +62,6 @@ def main(): print(memory_saver.memory_stats()) -def _register_callbacks(args: Namespace, flux: Flux1Redux) -> MemorySaver | None: - # Battery saver - battery_saver = BatterySaver(battery_percentage_stop_limit=args.battery_percentage_stop_limit) - CallbackRegistry.register_before_loop(battery_saver) - - # VAE Tiling - if args.vae_tiling: - flux.vae.decoder.enable_tiling = True - flux.vae.decoder.split_direction = args.vae_tiling_split - - # Stepwise Handler - if args.stepwise_image_output_dir: - handler = StepwiseHandler(flux=flux, output_dir=args.stepwise_image_output_dir) - CallbackRegistry.register_before_loop(handler) - CallbackRegistry.register_in_loop(handler) - CallbackRegistry.register_interrupt(handler) - - # Memory Saver - memory_saver = None - if args.low_ram: - memory_saver = MemorySaver(flux=flux, keep_transformer=len(args.seed) > 1) - CallbackRegistry.register_before_loop(memory_saver) - CallbackRegistry.register_in_loop(memory_saver) - CallbackRegistry.register_after_loop(memory_saver) - return memory_saver - - def _validate_redux_image_strengths( redux_image_paths: list[Path], redux_image_strengths: list[float] | None, diff --git a/src/mflux/post_processing/generated_image.py b/src/mflux/post_processing/generated_image.py index 87b8bc5..26e853c 100644 --- a/src/mflux/post_processing/generated_image.py +++ b/src/mflux/post_processing/generated_image.py @@ -5,6 +5,7 @@ import mlx.core as mx import PIL.Image import toml +from mflux.community.concept_attention.attention_data import ConceptHeatmap from mflux.config.model_config import ModelConfig @@ -30,6 +31,7 @@ class GeneratedImage: 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 @@ -50,6 +52,7 @@ class GeneratedImage: 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 @@ -75,6 +78,7 @@ class GeneratedImage: 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( @@ -87,15 +91,43 @@ class GeneratedImage: 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 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: - """Generate metadata for reference as well as input data for - command line --config-from-metadata arg in future generations. - """ return { "mflux_version": GeneratedImage.get_version(), "model": self.model_config.model_name, diff --git a/src/mflux/post_processing/image_util.py b/src/mflux/post_processing/image_util.py index 55a27b4..9164e75 100644 --- a/src/mflux/post_processing/image_util.py +++ b/src/mflux/post_processing/image_util.py @@ -8,6 +8,7 @@ import piexif import PIL.Image import PIL.ImageDraw +from mflux.community.concept_attention.attention_data import ConceptHeatmap from mflux.config.runtime_config import RuntimeConfig from mflux.post_processing.generated_image import GeneratedImage from mflux.ui.box_values import AbsoluteBoxValues, BoxValues @@ -33,6 +34,7 @@ class ImageUtil: image_strength: float | None = None, masked_image_path: str | Path | None = None, depth_image_path: str | Path | None = None, + concept_heatmap: ConceptHeatmap | None = None, ) -> GeneratedImage: normalized = ImageUtil._denormalize(decoded_latents) normalized_numpy = ImageUtil._to_numpy(normalized) @@ -57,6 +59,7 @@ class ImageUtil: depth_image_path=depth_image_path, redux_image_paths=redux_image_paths, redux_image_strengths=redux_image_strengths, + concept_heatmap=concept_heatmap, ) @staticmethod diff --git a/src/mflux/ui/cli/parsers.py b/src/mflux/ui/cli/parsers.py index a6d512c..2108999 100644 --- a/src/mflux/ui/cli/parsers.py +++ b/src/mflux/ui/cli/parsers.py @@ -126,6 +126,20 @@ class CommandLineParser(argparse.ArgumentParser): self.add_argument("--controlnet-strength", type=float, default=ui_defaults.CONTROLNET_STRENGTH, help=f"Controls how strongly the control image influences the output image. A value of 0.0 means no influence. (Default is {ui_defaults.CONTROLNET_STRENGTH})") self.add_argument("--controlnet-save-canny", action="store_true", help="If set, save the Canny edge detection reference input image.") + def add_concept_attention_arguments(self) -> None: + concept_group = self.add_argument_group("Concept Attention configuration") + concept_group.add_argument("--concept", type=str, required=True, help="The concept prompt to use for attention visualization") + concept_group.add_argument("--input-image-path", type=Path, required=False, default=None, help="Local path to reference image for concept attention analysis (uses FluxConceptFromImage instead of text-based concept)") + concept_group.add_argument("--heatmap-layer-indices", type=int, nargs="*", default=list(range(15, 19)), help="Layer indices to use for heatmap generation (default: 15-18)") + concept_group.add_argument("--heatmap-timesteps", type=int, nargs="*", default=None, help="Timesteps to use for heatmap generation (default: all timesteps)") + + def add_concept_from_image_arguments(self) -> None: + concept_group = self.add_argument_group("Concept Attention from Image configuration") + concept_group.add_argument("--concept", type=str, required=True, help="The concept prompt to use for attention visualization") + concept_group.add_argument("--input-image-path", type=Path, required=True, help="Local path to reference image for concept attention analysis") + concept_group.add_argument("--heatmap-layer-indices", type=int, nargs="*", default=list(range(15, 19)), help="Layer indices to use for heatmap generation (default: 15-18)") + concept_group.add_argument("--heatmap-timesteps", type=int, nargs="*", default=None, help="Timesteps to use for heatmap generation (default: all timesteps)") + def add_metadata_config(self) -> None: self.supports_metadata_config = True self.add_argument("--config-from-metadata", "-C", type=Path, required=False, default=argparse.SUPPRESS, help="Re-use the parameters from prior metadata. Params from metadata are secondary to other args you provide.") diff --git a/src/mflux/upscale.py b/src/mflux/upscale.py index 1fc7c8b..a730a0c 100644 --- a/src/mflux/upscale.py +++ b/src/mflux/upscale.py @@ -1,10 +1,5 @@ -from argparse import Namespace - from mflux import Config, Flux1Controlnet, ModelConfig, StopImageGenerationException -from mflux.callbacks.callback_registry import CallbackRegistry -from mflux.callbacks.instances.battery_saver import BatterySaver -from mflux.callbacks.instances.memory_saver import MemorySaver -from mflux.callbacks.instances.stepwise_handler import StepwiseHandler +from mflux.callbacks.callback_manager import CallbackManager from mflux.error.exceptions import PromptFileReadError from mflux.ui.cli.parsers import CommandLineParser from mflux.ui.prompt_utils import get_effective_prompt @@ -31,7 +26,7 @@ def main(): ) # 2. Register the optional callbacks - memory_saver = _register_callbacks(args=args, flux=flux) + memory_saver = CallbackManager.register_callbacks(args=args, flux=flux) try: for seed in args.seed: @@ -57,32 +52,5 @@ def main(): print(memory_saver.memory_stats()) -def _register_callbacks(args: Namespace, flux: Flux1Controlnet) -> MemorySaver | None: - # Battery saver - battery_saver = BatterySaver(battery_percentage_stop_limit=args.battery_percentage_stop_limit) - CallbackRegistry.register_before_loop(battery_saver) - - # VAE Tiling - if args.vae_tiling: - flux.vae.decoder.enable_tiling = True - flux.vae.decoder.split_direction = args.vae_tiling_split - - # Stepwise Handler - if args.stepwise_image_output_dir: - handler = StepwiseHandler(flux=flux, output_dir=args.stepwise_image_output_dir) - CallbackRegistry.register_before_loop(handler) - CallbackRegistry.register_in_loop(handler) - CallbackRegistry.register_interrupt(handler) - - # Memory Saver - memory_saver = None - if args.low_ram: - memory_saver = MemorySaver(flux=flux, keep_transformer=len(args.seed) > 1) - CallbackRegistry.register_before_loop(memory_saver) - CallbackRegistry.register_in_loop(memory_saver) - CallbackRegistry.register_after_loop(memory_saver) - return memory_saver - - if __name__ == "__main__": main() diff --git a/tests/arg_parser/test_cli_argparser.py b/tests/arg_parser/test_cli_argparser.py index cde36e1..f2fb76e 100644 --- a/tests/arg_parser/test_cli_argparser.py +++ b/tests/arg_parser/test_cli_argparser.py @@ -140,6 +140,30 @@ def mflux_redux_minimal_argv() -> list[str]: return ["mflux-generate-redux", "--redux-image-paths", "image1.png", "image2.png"] +@pytest.fixture +def mflux_concept_parser() -> CommandLineParser: + parser = CommandLineParser(description="Generate an image with concept attention based on a prompt and concept.") + parser.add_general_arguments() + parser.add_model_arguments(require_model_arg=False) + parser.add_lora_arguments() + parser.add_image_generator_arguments(supports_metadata_config=True) + parser.add_image_to_image_arguments(required=False) + parser.add_output_arguments() + parser.add_concept_attention_arguments() + return parser + + +@pytest.fixture +def mflux_concept_minimal_argv() -> list[str]: + return [ + "mflux-concept", + "--prompt", + "a beautiful landscape with a car", + "--concept", + "car", + ] + + def test_model_path_requires_model_arg(mflux_generate_parser): # when loading a model via --path, the model name still need to be specified with patch("sys.argv", "mflux-generate", "--path", "/some/saved/model"): @@ -657,3 +681,40 @@ def test_redux_args(mflux_redux_parser, mflux_redux_minimal_argv): args = mflux_redux_parser.parse_args() assert len(args.redux_image_paths) == 2 assert args.output == "redux_result.png" + + +def test_concept_attention_args(mflux_concept_parser, mflux_concept_minimal_argv): + # Test required arguments + with patch("sys.argv", mflux_concept_minimal_argv): + args = mflux_concept_parser.parse_args() + assert args.prompt == "a beautiful landscape with a car" + assert args.concept == "car" + # Test defaults + assert args.heatmap_layer_indices == list(range(15, 19)) + assert args.heatmap_timesteps is None + + # Test with missing required concept - should raise SystemExit + with patch("sys.argv", ["mflux-concept", "--prompt", "test"]): + pytest.raises(SystemExit, mflux_concept_parser.parse_args) + + # Test with missing regular prompt - should raise SystemExit + with patch("sys.argv", ["mflux-concept", "--concept", "test concept"]): + pytest.raises(SystemExit, mflux_concept_parser.parse_args) + + # Test with custom heatmap parameters + custom_argv = mflux_concept_minimal_argv + [ + "--heatmap-layer-indices", + "10", + "11", + "12", + "--heatmap-timesteps", + "0", + "1", + "2", + ] + with patch("sys.argv", custom_argv): + args = mflux_concept_parser.parse_args() + assert args.prompt == "a beautiful landscape with a car" + assert args.concept == "car" + assert args.heatmap_layer_indices == [10, 11, 12] + assert args.heatmap_timesteps == [0, 1, 2] diff --git a/tests/image_generation/helpers/image_generation_concept_test_helper.py b/tests/image_generation/helpers/image_generation_concept_test_helper.py new file mode 100644 index 0000000..eec42d1 --- /dev/null +++ b/tests/image_generation/helpers/image_generation_concept_test_helper.py @@ -0,0 +1,134 @@ +import os +from pathlib import Path + +import numpy as np +from PIL import Image + +from mflux import Config, ModelConfig +from mflux.community.concept_attention.flux_concept import Flux1Concept +from mflux.community.concept_attention.flux_concept_from_image import FluxConceptFromImage + + +class ImageGenerationConceptTestHelper: + @staticmethod + def assert_matches_reference_image_concept( + reference_heatmap_path: str, + output_heatmap_path: str, + model_config: ModelConfig, + prompt: str, + concept: str, + steps: int, + seed: int, + height: int | None = None, + width: int | None = None, + heatmap_layer_indices: list[int] | None = None, + heatmap_timesteps: list[int] | None = None, + lora_paths: list[str] | None = None, + lora_scales: list[float] | None = None, + ): + # resolve paths + reference_heatmap_path = ImageGenerationConceptTestHelper.resolve_path(reference_heatmap_path) + output_heatmap_path = ImageGenerationConceptTestHelper.resolve_path(output_heatmap_path) + lora_paths = [str(ImageGenerationConceptTestHelper.resolve_path(p)) for p in lora_paths] if lora_paths else None + + try: + # given + flux = Flux1Concept( + model_config=model_config, + quantize=4, + lora_paths=lora_paths, + lora_scales=lora_scales, + ) + # when + image = flux.generate_image( + seed=seed, + prompt=prompt, + concept=concept, + heatmap_layer_indices=heatmap_layer_indices, + heatmap_timesteps=heatmap_timesteps, + config=Config( + num_inference_steps=steps, + height=height, + width=width, + ), + ) + # Save only the heatmap (we don't need the original image for testing) + image.save_concept_heatmap(path=output_heatmap_path, overwrite=True) + + # then - verify the heatmap matches reference + np.testing.assert_array_equal( + np.array(Image.open(output_heatmap_path)), + np.array(Image.open(reference_heatmap_path)), + err_msg=f"Generated concept heatmap doesn't match reference heatmap. Check {output_heatmap_path} vs {reference_heatmap_path}", + ) + + finally: + # cleanup + if os.path.exists(output_heatmap_path): + os.remove(output_heatmap_path) + + @staticmethod + def assert_matches_reference_image_concept_from_image( + reference_heatmap_path: str, + output_heatmap_path: str, + input_image_path: str, + model_config: ModelConfig, + prompt: str, + concept: str, + steps: int, + seed: int, + height: int | None = None, + width: int | None = None, + heatmap_layer_indices: list[int] | None = None, + heatmap_timesteps: list[int] | None = None, + lora_paths: list[str] | None = None, + lora_scales: list[float] | None = None, + ): + # resolve paths + reference_heatmap_path = ImageGenerationConceptTestHelper.resolve_path(reference_heatmap_path) + output_heatmap_path = ImageGenerationConceptTestHelper.resolve_path(output_heatmap_path) + input_image_path = ImageGenerationConceptTestHelper.resolve_path(input_image_path) + lora_paths = [str(ImageGenerationConceptTestHelper.resolve_path(p)) for p in lora_paths] if lora_paths else None + + try: + # given + flux = FluxConceptFromImage( + model_config=model_config, + quantize=8, + lora_paths=lora_paths, + lora_scales=lora_scales, + ) + # when + image = flux.generate_image( + seed=seed, + prompt=prompt, + concept=concept, + image_path=str(input_image_path), + heatmap_layer_indices=heatmap_layer_indices, + heatmap_timesteps=heatmap_timesteps, + config=Config( + num_inference_steps=steps, + height=height, + width=width, + ), + ) + # Save only the heatmap (we don't need the original image for testing) + image.save_concept_heatmap(path=output_heatmap_path, overwrite=True) + + # then - verify the heatmap matches reference + np.testing.assert_array_equal( + np.array(Image.open(output_heatmap_path)), + np.array(Image.open(reference_heatmap_path)), + err_msg=f"Generated concept from image heatmap doesn't match reference heatmap. Check {output_heatmap_path} vs {reference_heatmap_path}", + ) + + finally: + # cleanup + if os.path.exists(output_heatmap_path): + os.remove(output_heatmap_path) + + @staticmethod + def resolve_path(path) -> Path | None: + if path is None: + return None + return Path(__file__).parent.parent.parent / "resources" / path diff --git a/tests/image_generation/test_generate_concept.py b/tests/image_generation/test_generate_concept.py new file mode 100644 index 0000000..4d186e0 --- /dev/null +++ b/tests/image_generation/test_generate_concept.py @@ -0,0 +1,35 @@ +from mflux import ModelConfig +from tests.image_generation.helpers.image_generation_concept_test_helper import ImageGenerationConceptTestHelper + + +class TestImageGeneratorConcept: + def test_concept_attention_generation(self): + ImageGenerationConceptTestHelper.assert_matches_reference_image_concept( + reference_heatmap_path="reference_concept_schnell_heatmap.png", + output_heatmap_path="output_concept_schnell_heatmap.png", + model_config=ModelConfig.schnell(), + prompt="A dragon on a hill", + concept="dragon", + steps=4, + seed=44, + height=512, + width=512, + heatmap_layer_indices=[15, 16, 17, 18], + heatmap_timesteps=[0, 1, 2, 3], + ) + + def test_concept_attention_from_image(self): + ImageGenerationConceptTestHelper.assert_matches_reference_image_concept_from_image( + reference_heatmap_path="reference_concept_from_image_schnell_heatmap.png", + output_heatmap_path="output_concept_from_image_schnell_heatmap.png", + input_image_path="reference_depth_dev_from_image.png", + model_config=ModelConfig.schnell(), + prompt="A photo of cartoon of Albert Einstein", + concept="man", + steps=4, + seed=42, + height=512, + width=320, + heatmap_layer_indices=[15, 16, 17, 18], + heatmap_timesteps=[0, 1, 2, 3], + )