Concept Attention (#199)

This commit is contained in:
Filip Strand 2025-06-04 19:27:03 +02:00 committed by GitHub
parent aaae64ad6e
commit 13bfb24964
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
30 changed files with 1531 additions and 258 deletions

View File

@ -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

View File

@ -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"

Binary file not shown.

After

Width:  |  Height:  |  Size: 845 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 324 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 342 KiB

View File

@ -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

View File

@ -0,0 +1 @@
# Concept Attention transformer models

View File

@ -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,
}

View File

@ -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

View File

@ -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,
)

View File

@ -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,
)

View File

@ -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

View File

@ -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

View File

@ -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

61
src/mflux/concept.py Normal file
View File

@ -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()

View File

@ -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()

View File

@ -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,
)

View File

@ -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()

View File

@ -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()

View File

@ -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()

View File

@ -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()

View File

@ -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()

View File

@ -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,

View File

@ -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,

View File

@ -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

View File

@ -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.")

View File

@ -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()

View File

@ -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]

View File

@ -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

View File

@ -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],
)