Concept Attention (#199)
This commit is contained in:
parent
aaae64ad6e
commit
13bfb24964
93
README.md
93
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.
|
||||
|
||||

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

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

|
||||
|
||||
#### 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
|
||||
|
||||
@ -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"
|
||||
|
||||
BIN
src/mflux/assets/concept_example_1.jpg
Normal file
BIN
src/mflux/assets/concept_example_1.jpg
Normal file
Binary file not shown.
|
After Width: | Height: | Size: 845 KiB |
BIN
src/mflux/assets/concept_example_2.jpg
Normal file
BIN
src/mflux/assets/concept_example_2.jpg
Normal file
Binary file not shown.
|
After Width: | Height: | Size: 324 KiB |
BIN
src/mflux/assets/concept_example_3.jpg
Normal file
BIN
src/mflux/assets/concept_example_3.jpg
Normal file
Binary file not shown.
|
After Width: | Height: | Size: 342 KiB |
77
src/mflux/callbacks/callback_manager.py
Normal file
77
src/mflux/callbacks/callback_manager.py
Normal 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
|
||||
1
src/mflux/community/concept_attention/__init__.py
Normal file
1
src/mflux/community/concept_attention/__init__.py
Normal file
@ -0,0 +1 @@
|
||||
# Concept Attention transformer models
|
||||
68
src/mflux/community/concept_attention/attention_data.py
Normal file
68
src/mflux/community/concept_attention/attention_data.py
Normal 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,
|
||||
}
|
||||
93
src/mflux/community/concept_attention/concept_util.py
Normal file
93
src/mflux/community/concept_attention/concept_util.py
Normal 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
|
||||
181
src/mflux/community/concept_attention/flux_concept.py
Normal file
181
src/mflux/community/concept_attention/flux_concept.py
Normal 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,
|
||||
)
|
||||
195
src/mflux/community/concept_attention/flux_concept_from_image.py
Normal file
195
src/mflux/community/concept_attention/flux_concept_from_image.py
Normal 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,
|
||||
)
|
||||
125
src/mflux/community/concept_attention/joint_attention_concept.py
Normal file
125
src/mflux/community/concept_attention/joint_attention_concept.py
Normal 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
|
||||
@ -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
|
||||
96
src/mflux/community/concept_attention/transformer_concept.py
Normal file
96
src/mflux/community/concept_attention/transformer_concept.py
Normal 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
61
src/mflux/concept.py
Normal 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()
|
||||
62
src/mflux/concept_from_image.py
Normal file
62
src/mflux/concept_from_image.py
Normal 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()
|
||||
@ -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,
|
||||
)
|
||||
|
||||
@ -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()
|
||||
|
||||
@ -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()
|
||||
|
||||
@ -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()
|
||||
|
||||
@ -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()
|
||||
|
||||
@ -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()
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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.")
|
||||
|
||||
@ -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()
|
||||
|
||||
@ -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]
|
||||
|
||||
@ -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
|
||||
35
tests/image_generation/test_generate_concept.py
Normal file
35
tests/image_generation/test_generate_concept.py
Normal 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],
|
||||
)
|
||||
Loading…
Reference in New Issue
Block a user