Add Flux1.Kontext[dev] support (#219)

This commit is contained in:
Filip Strand 2025-06-27 23:37:19 +02:00 committed by GitHub
parent b26a1d66b3
commit 7388a856e5
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
28 changed files with 629 additions and 69 deletions

View File

@ -7,6 +7,20 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
## [Unreleased]
### 📸 FLUX.1 Kontext
- **Added FLUX.1 Kontext support**: Official Black Forest Labs model for character consistency, local editing, and style reference
- **New command**: `mflux-generate-kontext` for image-guided generation with text instructions
## Breaking Changes
### 🔧 Python API Class Naming Standardization
- **Class rename**: `FluxInContextFill` is now `Flux1InContextFill` to follow consistent naming convention
- **Class rename**: `FluxConceptFromImage` is now `Flux1ConceptFromImage` to follow consistent naming convention
- **Breaking change for library users**: If you import these classes directly in Python code, you may need to update your imports
- **CLI tools unaffected**: All command-line tools (`mflux-generate-*`) continue to work without changes
## [0.8.0] - 2025-06-14
# MFLUX v.0.8.0 Release Notes

104
README.md
View File

@ -22,6 +22,7 @@ Run the powerful [FLUX](https://blackforestlabs.ai/#get-flux) models from [Black
- [🎨 Image-to-Image](#-image-to-image)
- [🔌 LoRA](#-lora)
- [🎭 In-Context Generation](#-in-context-generation)
* [📸 Kontext](#-kontext)
* [🎨 In-Context LoRA](#-in-context-lora)
* [👕 CatVTON (Virtual Try-On)](#-catvton-virtual-try-on)
* [✏️ IC-Edit (In-Context Editing)](#%EF%B8%8F-ic-edit-in-context-editing)
@ -274,6 +275,23 @@ For more advanced Python usage and additional configuration options, you can exp
</details>
#### Kontext Command-Line Arguments
<details>
<summary>Click to expand Kontext arguments</summary>
The `mflux-generate-kontext` command supports most of the same arguments as `mflux-generate`, with these specific parameters:
- **`--image-path`** (required, `str`): Path to the reference image that will be used as the base for editing. This image provides the visual context for the text-guided modifications.
- **`--guidance`** (optional, `float`, default: `2.5`): Guidance scale for the generation. Kontext typically works well with values between 2.0-4.0, with lower values allowing more creative freedom and higher values adhering more closely to the prompt.
**Note**: The Kontext tool automatically uses the appropriate model (`dev-kontext`), so you don't need to specify `--model`.
See the [Kontext](#-kontext) section for more details on this feature and comprehensive prompting strategies.
</details>
#### In-Context LoRA Command-Line Arguments
<details>
@ -993,7 +1011,91 @@ To report additional formats, examples or other any suggestions related to LoRA
In-Context Generation is a powerful collection of techniques that allows you to generate images based on reference images and context, without requiring model fine-tuning. MFLUX supports multiple in-context approaches, each designed for specific use cases ranging from style transfer to virtual try-on and image editing.
**Note**: Black Forest Labs has announced [FLUX.1 Kontext](https://bfl.ai/models/flux-kontext), a new model suite designed specifically for text-and-image-driven generation and editing with advanced in-context capabilities. This official model may replace existing third-party implementations in the future. FLUX.1 Kontext offers character consistency, local editing, style reference, and interactive speed optimizations. The dev variant is coming soon as an open-weights model.
#### 📸 Kontext
FLUX.1 Kontext, released in June 2025, is Black Forest Labs' [official model](https://bfl.ai/models/flux-kontext) that goes beyond text-to-image generation. Unlike previous flow models that only allow text-based generation, Kontext understands and creates from existing images, enabling flexible image editing through simple text instructions without complex workflows or finetuning.
Out of the various in-context generation techniques, Kontext is the most advanced and versatile, allowing for a wide range of image editing and generation tasks. It is recommended to try Kontext first before exploring other in-context methods.
![kontext example](src/mflux/assets/kontext_example.jpg)
*Original images credit: [Jon Tyson (sign)](https://unsplash.com/photos/the-sign-urges-viewers-to-get-more-rest-a4Cm6EnYtls), [Daniel Seßler (town)](https://unsplash.com/photos/colorful-buildings-line-a-coastal-towns-harbor-haVcuj4buqE), [Liga Alksne (house)](https://unsplash.com/photos/a-building-with-a-mountain-in-the-background-jZV4jaycoxE) on Unsplash*
Kontext's core capabilities include **character consistency** (preserving unique elements across scenes), **local editing** (targeted modifications without affecting the rest), **style reference** (generating novel scenes while preserving styles), and **interactive speed** for iterative refinement with minimal latency.
##### Example Usage with sequential edits
![kontext sequential edit example](src/mflux/assets/kontext_sequential_edit_example.jpg)
The image above illustrates **sequential editing** with Kontext - making multiple individual edits to progressively refine an image. Each edit builds upon the previous result, allowing for precise, iterative modifications.
Here's an example of one step in a sequential editing workflow, where we add atmospheric mist to the mountains:
```bash
mflux-generate-kontext \
--image-path "house_without_background_objects.png" \
--prompt "Make the mountains in the background be covered in white mist so that they are barely visible. Preserve the original camera angle, house placement etc" \
--steps 20 \
--seed 9375333 \
--guidance 2.5 \
--width 912 \
--height 1360 \
-q 8
```
##### Tips for Best Results
For comprehensive prompting strategies, see the [official Black Forest Labs Kontext prompting guide](https://docs.bfl.ai/guides/prompting_guide_kontext_i2i).
**Core Principles:**
- **Be Specific and Clear**: Use precise descriptions, avoid vague terms
- **Step-by-step Editing**: Break complex modifications into multiple simple steps
- **Explicit Preservation**: State what should remain unchanged
- **Verb Selection**: Use "change", "replace" rather than "transform"
**Prompting Techniques:**
*Basic Modifications:*
- Simple: `"Change the car color to red"`
- Maintain style: `"Change to daytime while maintaining the same style of the painting"`
*Character Consistency:*
- Use specific descriptions: `"The woman with short black hair"` instead of `"she"`
- Preserve features: `"while maintaining the same facial features, hairstyle, and expression"`
*Style Transfer:*
- Name the style: `"Transform to Bauhaus art style"`
- Describe characteristics: `"Transform to oil painting with visible brushstrokes, thick paint texture"`
- Preserve composition: `"Change to Bauhaus style while maintaining the original composition"`
*Text Editing:*
- Use quotes: `"Replace 'joy' with 'BFL'"`
- Maintain format: `"Replace text while maintaining the same font style"`
**Common Problems & Solutions:**
**Character changes too much**: `"Transform the person into a Viking"`
**Better**: `"Change the clothes to be a viking warrior while preserving facial features"`
**Position changes**: `"Put him on a beach"`
**Better**: `"Change the background to a beach while keeping the person in the exact same position, scale, and pose"`
**Vague style**: `"Make it a sketch"`
**Better**: `"Convert to pencil sketch with natural graphite lines, cross-hatching, and visible paper texture"`
**Ready-to-Use Templates:**
- **Object Modification**: `"Change [object] to [new state], keep [content to preserve] unchanged"`
- **Style Transfer**: `"Transform to [specific style], while maintaining [composition/character/other] unchanged"`
- **Background Replacement**: `"Change the background to [new background], keep the subject in the exact same position and pose"`
- **Text Editing**: `"Replace '[original text]' with '[new text]', maintain the same font style"`
**Technical Tips:**
1. **High-Quality References**: Use clear, well-lit reference images with distinct characteristics
2. **Guidance Scale**: Experiment with guidance values between 2.5-4.0 for optimal balance
3. **Multiple Seeds**: Try different seeds to explore variations while maintaining consistency
4. **Resolution**: Use standard resolutions (1024x1024, 768x1024, etc.) for best results
⚠️ *Note: Using the Kontext tool requires an additional ~34 GB download from [black-forest-labs/FLUX.1-Kontext-dev](https://huggingface.co/black-forest-labs/FLUX.1-Kontext-dev). The download happens automatically on first use.*
#### 🎨 In-Context LoRA

View File

@ -64,6 +64,7 @@ mflux-generate-in-context-catvton = "mflux.generate_in_context_catvton:main"
mflux-generate-fill = "mflux.generate_fill:main"
mflux-generate-depth = "mflux.generate_depth:main"
mflux-generate-redux = "mflux.generate_redux:main"
mflux-generate-kontext = "mflux.generate_kontext:main"
mflux-concept = "mflux.concept:main"
mflux-concept-from-image = "mflux.concept_from_image:main"
mflux-save = "mflux.save:main"

Binary file not shown.

After

Width:  |  Height:  |  Size: 959 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 555 KiB

View File

@ -21,7 +21,7 @@ from mflux.post_processing.generated_image import GeneratedImage
from mflux.post_processing.image_util import ImageUtil
class FluxConceptFromImage(nn.Module):
class Flux1ConceptFromImage(nn.Module):
vae: VAE
transformer: TransformerConcept
t5_text_encoder: T5Encoder

View File

@ -53,8 +53,8 @@ class TransformerConcept(nn.Module):
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
image_rotary_embeddings = Transformer.compute_rotary_embeddings(prompt_embeds, self.pos_embed, config, None) # fmt: off
image_rotary_embeddings_concept = Transformer.compute_rotary_embeddings(prompt_embeds_concept, self.pos_embed, config, None) # fmt: off
# 2. Run the joint transformer blocks
attention_information = []

View File

@ -20,7 +20,7 @@ from mflux.post_processing.generated_image import GeneratedImage
from mflux.post_processing.image_util import ImageUtil
class FluxInContextFill(nn.Module):
class Flux1InContextFill(nn.Module):
vae: VAE
transformer: Transformer
t5_text_encoder: T5Encoder

View File

@ -15,6 +15,7 @@ class CompletionGenerator:
self.commands = [
"mflux-generate",
"mflux-generate-controlnet",
"mflux-generate-kontext",
"mflux-generate-in-context",
"mflux-generate-in-context-edit",
"mflux-generate-in-context-catvton",
@ -51,6 +52,14 @@ class CompletionGenerator:
parser.add_image_generator_arguments()
parser.add_output_arguments()
elif command == "mflux-generate-kontext":
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=True)
parser.add_output_arguments()
elif command == "mflux-generate-in-context":
parser.add_general_arguments()
parser.add_model_arguments()

View File

@ -1,6 +1,6 @@
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.community.concept_attention.flux_concept_from_image import Flux1ConceptFromImage
from mflux.error.exceptions import PromptFileReadError
from mflux.ui import defaults as ui_defaults
from mflux.ui.cli.parsers import CommandLineParser
@ -24,7 +24,7 @@ def main():
args.guidance = ui_defaults.GUIDANCE_SCALE
# 1. Load the concept attention model
flux = FluxConceptFromImage(
flux = Flux1ConceptFromImage(
model_config=ModelConfig.from_name(model_name=args.model, base_model=args.base_model),
quantize=args.quantize,
local_path=args.path,

View File

@ -34,6 +34,16 @@ class ModelConfig:
def dev() -> "ModelConfig":
return AVAILABLE_MODELS["dev"]
@staticmethod
@lru_cache
def schnell() -> "ModelConfig":
return AVAILABLE_MODELS["schnell"]
@staticmethod
@lru_cache
def dev_kontext() -> "ModelConfig":
return AVAILABLE_MODELS["dev-kontext"]
@staticmethod
@lru_cache
def dev_fill() -> "ModelConfig":
@ -41,19 +51,14 @@ class ModelConfig:
@staticmethod
@lru_cache
def dev_fill_catvton() -> "ModelConfig":
return AVAILABLE_MODELS["dev-fill-catvton"]
def dev_redux() -> "ModelConfig":
return AVAILABLE_MODELS["dev-redux"]
@staticmethod
@lru_cache
def dev_depth() -> "ModelConfig":
return AVAILABLE_MODELS["dev-depth"]
@staticmethod
@lru_cache
def dev_redux() -> "ModelConfig":
return AVAILABLE_MODELS["dev-redux"]
@staticmethod
@lru_cache
def dev_controlnet_canny() -> "ModelConfig":
@ -71,8 +76,8 @@ class ModelConfig:
@staticmethod
@lru_cache
def schnell() -> "ModelConfig":
return AVAILABLE_MODELS["schnell"]
def dev_fill_catvton() -> "ModelConfig":
return AVAILABLE_MODELS["dev-fill-catvton"]
def x_embedder_input_dim(self) -> int:
if self.alias and "dev-fill" in self.alias:
@ -112,8 +117,13 @@ class ModelConfig:
# Find by explicit base_model name
default_base = next((b for b in base_models if base_model in (b.alias, b.model_name)), None)
else:
# Infer from model_name substring (priority order via sorted base_models)
default_base = next((b for b in base_models if b.alias and b.alias in model_name), None)
# Infer from model_name substring - prefer longer matches (more specific)
matching_bases = [b for b in base_models if b.alias and b.alias in model_name]
if matching_bases:
# Sort by alias length descending, then by priority ascending
default_base = sorted(matching_bases, key=lambda x: (-len(x.alias), x.priority))[0]
else:
default_base = None
if not default_base:
raise ModelConfigError(f"Cannot infer base_model from {model_name}")
@ -133,18 +143,6 @@ class ModelConfig:
AVAILABLE_MODELS = {
"schnell": ModelConfig(
alias="schnell",
model_name="black-forest-labs/FLUX.1-schnell",
base_model=None,
controlnet_model=None,
custom_transformer_model=None,
num_train_steps=1000,
max_sequence_length=256,
supports_guidance=False,
requires_sigma_shift=False,
priority=2,
),
"dev": ModelConfig(
alias="dev",
model_name="black-forest-labs/FLUX.1-dev",
@ -155,8 +153,32 @@ AVAILABLE_MODELS = {
max_sequence_length=512,
supports_guidance=True,
requires_sigma_shift=True,
priority=0,
),
"schnell": ModelConfig(
alias="schnell",
model_name="black-forest-labs/FLUX.1-schnell",
base_model=None,
controlnet_model=None,
custom_transformer_model=None,
num_train_steps=1000,
max_sequence_length=256,
supports_guidance=False,
requires_sigma_shift=False,
priority=1,
),
"dev-kontext": ModelConfig(
alias="dev-kontext",
model_name="black-forest-labs/FLUX.1-Kontext-dev",
base_model=None,
controlnet_model=None,
custom_transformer_model=None,
num_train_steps=1000,
max_sequence_length=512,
supports_guidance=True,
requires_sigma_shift=True,
priority=2,
),
"dev-fill": ModelConfig(
alias="dev-fill",
model_name="black-forest-labs/FLUX.1-Fill-dev",
@ -167,19 +189,7 @@ AVAILABLE_MODELS = {
max_sequence_length=512,
supports_guidance=True,
requires_sigma_shift=True,
priority=0,
),
"dev-depth": ModelConfig(
alias="dev-depth",
model_name="black-forest-labs/FLUX.1-Depth-dev",
base_model=None,
controlnet_model=None,
custom_transformer_model=None,
num_train_steps=1000,
max_sequence_length=512,
supports_guidance=True,
requires_sigma_shift=True,
priority=4,
priority=3,
),
"dev-redux": ModelConfig(
alias="dev-redux",
@ -191,19 +201,19 @@ AVAILABLE_MODELS = {
max_sequence_length=512,
supports_guidance=True,
requires_sigma_shift=True,
priority=3,
priority=4,
),
"dev-fill-catvton": ModelConfig(
alias="dev-fill-catvton",
model_name="black-forest-labs/FLUX.1-Fill-dev",
"dev-depth": ModelConfig(
alias="dev-depth",
model_name="black-forest-labs/FLUX.1-Depth-dev",
base_model=None,
controlnet_model=None,
custom_transformer_model="xiaozaa/catvton-flux-beta",
custom_transformer_model=None,
num_train_steps=1000,
max_sequence_length=512,
supports_guidance=True,
requires_sigma_shift=False, # Not sure why, but produced better results this way...
priority=6,
requires_sigma_shift=True,
priority=5,
),
"dev-controlnet-canny": ModelConfig(
alias="dev-controlnet-canny",
@ -215,7 +225,7 @@ AVAILABLE_MODELS = {
max_sequence_length=512,
supports_guidance=True,
requires_sigma_shift=True,
priority=5,
priority=6,
),
"schnell-controlnet-canny": ModelConfig(
alias="schnell-controlnet-canny",
@ -227,7 +237,7 @@ AVAILABLE_MODELS = {
max_sequence_length=256,
supports_guidance=False,
requires_sigma_shift=False,
priority=6,
priority=7,
),
"dev-controlnet-upscaler": ModelConfig(
alias="dev-controlnet-upscaler",
@ -239,6 +249,18 @@ AVAILABLE_MODELS = {
max_sequence_length=512,
supports_guidance=False,
requires_sigma_shift=False,
priority=7,
priority=8,
),
"dev-fill-catvton": ModelConfig(
alias="dev-fill-catvton",
model_name="black-forest-labs/FLUX.1-Fill-dev",
base_model=None,
controlnet_model=None,
custom_transformer_model="xiaozaa/catvton-flux-beta",
num_train_steps=1000,
max_sequence_length=512,
supports_guidance=True,
requires_sigma_shift=False, # Not sure why, but produced better results this way...
priority=9,
),
}

View File

@ -42,7 +42,7 @@ class TransformerControlnet(nn.Module):
hidden_states = self.x_embedder(hidden_states) + self.controlnet_x_embedder(controlnet_condition)
encoder_hidden_states = self.context_embedder(prompt_embeds)
text_embeddings = Transformer.compute_text_embeddings(t, pooled_prompt_embeds, self.time_text_embed, config)
image_rotary_embeddings = Transformer.compute_rotary_embeddings(prompt_embeds, self.pos_embed, config)
image_rotary_embeddings = Transformer.compute_rotary_embeddings(prompt_embeds, self.pos_embed, config, None)
# 2. Run the joint transformer blocks
controlnet_block_samples = []

View File

@ -2,7 +2,7 @@ from pathlib import Path
from mflux import Config, ModelConfig, StopImageGenerationException
from mflux.callbacks.callback_manager import CallbackManager
from mflux.community.in_context.flux_in_context_fill import FluxInContextFill
from mflux.community.in_context.flux_in_context_fill import Flux1InContextFill
from mflux.error.exceptions import PromptFileReadError
from mflux.ui import defaults as ui_defaults
from mflux.ui.cli.parsers import CommandLineParser
@ -34,7 +34,7 @@ def main():
args.vae_tiling_split = "vertical"
# 1. Load the model
flux = FluxInContextFill(
flux = Flux1InContextFill(
model_config=ModelConfig.dev_fill_catvton(),
quantize=args.quantize,
local_path=args.path,

View File

@ -4,7 +4,7 @@ from PIL import Image
from mflux import Config, ModelConfig, StopImageGenerationException
from mflux.callbacks.callback_manager import CallbackManager
from mflux.community.in_context.flux_in_context_fill import FluxInContextFill
from mflux.community.in_context.flux_in_context_fill import Flux1InContextFill
from mflux.community.in_context.utils.in_context_loras import prepare_ic_edit_loras
from mflux.error.exceptions import PromptFileReadError
from mflux.ui import defaults as ui_defaults
@ -36,7 +36,7 @@ def main():
width, height = _resize_for_ic_edit_optimal_width(args)
# 1. Load the model with IC-Edit LoRA
flux = FluxInContextFill(
flux = Flux1InContextFill(
model_config=ModelConfig.dev_fill(),
quantize=args.quantize,
local_path=args.path,

View File

@ -0,0 +1,65 @@
from pathlib import Path
from mflux import Config, StopImageGenerationException
from mflux.callbacks.callback_manager import CallbackManager
from mflux.error.exceptions import PromptFileReadError
from mflux.kontext.flux_kontext import Flux1Kontext
from mflux.ui import defaults as ui_defaults
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 using Flux Kontext with image conditioning.")
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=True)
parser.add_output_arguments()
args = parser.parse_args()
# 0. Set default guidance value if not provided by user
if args.guidance is None:
args.guidance = ui_defaults.GUIDANCE_SCALE_KONTEXT
# 1. Load the model
flux = Flux1Kontext(
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),
config=Config(
num_inference_steps=args.steps,
height=args.height,
width=args.width,
guidance=args.guidance,
image_path=args.image_path,
),
)
# 4. Save the image
output_path = Path(args.output.format(seed=seed))
image.save(path=output_path, 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

View File

@ -0,0 +1,156 @@
import mlx.core as mx
from mlx import nn
from tqdm import tqdm
from mflux.callbacks.callbacks import Callbacks
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.kontext.utils.kontext_util import KontextUtil
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.transformer.transformer import Transformer
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 Flux1Kontext(nn.Module):
vae: VAE
transformer: Transformer
t5_text_encoder: T5Encoder
clip_text_encoder: CLIPEncoder
def __init__(
self,
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(
flux_model=self,
model_config=ModelConfig.dev_kontext(),
quantize=quantize,
local_path=local_path,
lora_paths=lora_paths,
lora_scales=lora_scales,
)
def generate_image(
self,
seed: int,
prompt: str,
config: Config,
) -> 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(
seed=seed,
height=config.height,
width=config.width,
)
# 2. Encode the 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. Create the static image conditioning latents and IDs
static_image_latents, kontext_image_ids = KontextUtil.create_image_conditioning_latents(
vae=self.vae,
height=config.height,
width=config.width,
image_path=config.image_path,
)
# (Optional) Call subscribers for beginning of loop
Callbacks.before_loop(
seed=seed,
prompt=prompt,
latents=latents,
config=config,
)
for t in time_steps:
try:
# 4.t Concatenate the updated latents with the static image latents
hidden_states = mx.concatenate([latents, static_image_latents], axis=1)
# 5.t Predict the noise
noise = self.transformer(
t=t,
config=config,
hidden_states=hidden_states,
prompt_embeds=prompt_embeds,
pooled_prompt_embeds=pooled_prompt_embeds,
kontext_image_ids=kontext_image_ids,
)
# 6.t Extract only the noise for the generation latents (first part)
noise = noise[:, : latents.shape[1]]
# 7.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,
)
# 8. 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,
generation_time=time_steps.format_dict["elapsed"],
)

View File

View File

@ -0,0 +1,84 @@
import mlx.core as mx
from mflux.latent_creator.latent_creator import LatentCreator
from mflux.post_processing.array_util import ArrayUtil
class KontextUtil:
@staticmethod
def create_image_conditioning_latents(
vae,
height: int,
width: int,
image_path: str,
) -> tuple[mx.array, mx.array]:
# Load and encode the input image
input_image = LatentCreator.encode_image(
vae=vae,
image_path=image_path,
height=height,
width=width,
)
# Pack image latents for conditioning
image_latents = ArrayUtil.pack_latents(
latents=input_image,
height=height,
width=width,
num_channels_latents=16,
)
# Create image IDs for positional embeddings
image_ids = KontextUtil._create_image_ids(
height=height,
width=width,
)
return image_latents, image_ids
@staticmethod
def _create_image_ids(
height: int,
width: int,
) -> mx.array:
# Create image IDs similar to the reference implementation
latent_height = height // 16 # VAE downsampling factor
latent_width = width // 16
# Create coordinate grid for image positioning
image_ids = mx.zeros((latent_height, latent_width, 3))
# Add row coordinates
row_coords = mx.arange(0, latent_height)[:, None]
row_coords = mx.broadcast_to(row_coords, (latent_height, latent_width))
image_ids = mx.concatenate(
[
image_ids[:, :, :1], # Keep first dimension as 0 for now
row_coords[:, :, None], # Set row coordinates
image_ids[:, :, 2:], # Keep remaining dimensions
],
axis=2,
)
# Add column coordinates
col_coords = mx.arange(0, latent_width)[None, :]
col_coords = mx.broadcast_to(col_coords, (latent_height, latent_width))
image_ids = mx.concatenate(
[
image_ids[:, :, :2], # Keep first two dimensions
col_coords[:, :, None], # Set column coordinates
],
axis=2,
)
# Reshape to sequence format
image_ids = mx.reshape(image_ids, (latent_height * latent_width, 3))
# Set the first dimension to 1 to distinguish from generation latents (which use 0)
first_dim = mx.ones((image_ids.shape[0], 1))
image_ids = mx.concatenate([first_dim, image_ids[:, 1:]], axis=1)
# Add batch dimension
image_ids = mx.expand_dims(image_ids, axis=0)
return image_ids

View File

@ -44,12 +44,13 @@ class Transformer(nn.Module):
pooled_prompt_embeds: mx.array,
controlnet_block_samples: list[mx.array] | None = None,
controlnet_single_block_samples: list[mx.array] | None = None,
kontext_image_ids: mx.array | None = None,
) -> mx.array:
# 1. Create embeddings
hidden_states = self.x_embedder(hidden_states)
encoder_hidden_states = self.context_embedder(prompt_embeds)
text_embeddings = Transformer.compute_text_embeddings(t, pooled_prompt_embeds, self.time_text_embed, config)
image_rotary_embeddings = Transformer.compute_rotary_embeddings(prompt_embeds, self.pos_embed, config)
image_rotary_embeddings = Transformer.compute_rotary_embeddings(prompt_embeds, self.pos_embed, config, kontext_image_ids) # fmt: off
# 2. Run the joint transformer blocks
for idx, block in enumerate(self.transformer_blocks):
@ -132,9 +133,18 @@ class Transformer(nn.Module):
return encoder_hidden_states, hidden_states
@staticmethod
def compute_rotary_embeddings(prompt_embeds: mx.array, pos_embed: EmbedND, config: RuntimeConfig) -> mx.array:
def compute_rotary_embeddings(
prompt_embeds: mx.array,
pos_embed: EmbedND,
config: RuntimeConfig,
kontext_image_ids: mx.array | None = None,
) -> mx.array:
txt_ids = Transformer._prepare_text_ids(seq_len=prompt_embeds.shape[1])
img_ids = Transformer._prepare_latent_image_ids(height=config.height, width=config.width)
if kontext_image_ids is not None:
img_ids = mx.concatenate([img_ids, kontext_image_ids], axis=1)
ids = mx.concatenate((txt_ids, img_ids), axis=1)
image_rotary_emb = pos_embed(ids)
return image_rotary_emb

View File

@ -171,7 +171,7 @@ class CommandLineParser(argparse.ArgumentParser):
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("--input-image-path", type=Path, required=False, default=None, help="Local path to reference image for concept attention analysis (uses Flux1ConceptFromImage 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)")

View File

@ -13,15 +13,17 @@ DEFAULT_DEV_FILL_GUIDANCE = 30
DEFAULT_DEPTH_GUIDANCE = 10
DIMENSION_STEP_PIXELS = 16
GUIDANCE_SCALE = 3.5
GUIDANCE_SCALE_KONTEXT = 2.5
HEIGHT, WIDTH = 1024, 1024
MAX_PIXELS_WARNING_THRESHOLD = 2048 * 2048
IMAGE_STRENGTH = 0.4
MODEL_CHOICES = ["dev", "dev-fill", "schnell"]
MODEL_CHOICES = ["dev", "schnell", "dev-kontext", "dev-fill"]
MODEL_INFERENCE_STEPS = {
"dev": 14,
"dev-fill": 14,
"dev-depth": 14,
"dev-redux": 14,
"dev-kontext": 14,
"schnell": 4,
}
QUANTIZE_CHOICES = [3, 4, 6, 8]

View File

@ -44,9 +44,20 @@ class WeightHandlerLoRAHuggingFace:
# Check if the file already exists in the cache
cached_file_path = cache_path / lora_name
if cached_file_path.exists():
print(f"Using cached LoRA: {cached_file_path}")
return str(cached_file_path)
if cached_file_path.exists() and cached_file_path.is_file():
try:
# Verify the file is actually readable (catches broken symlinks)
with open(cached_file_path, "rb") as f:
f.read(1) # Try to read just 1 byte to verify it works
print(f"Using cached LoRA: {cached_file_path}")
return str(cached_file_path)
except (OSError, IOError):
# File exists but is not readable (broken symlink, permissions, etc.)
print(f"Cached LoRA file is corrupted or inaccessible, re-downloading: {cached_file_path}")
try:
cached_file_path.unlink() # Remove the broken file/symlink
except OSError:
pass # Ignore if we can't remove it
# Download the LoRA from Hugging Face
print(f"Downloading LoRA '{lora_name}' from {repo_id}...")

View File

@ -6,7 +6,7 @@ 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
from mflux.community.concept_attention.flux_concept_from_image import Flux1ConceptFromImage
class ImageGenerationConceptTestHelper:
@ -92,7 +92,7 @@ class ImageGenerationConceptTestHelper:
try:
# given
flux = FluxConceptFromImage(
flux = Flux1ConceptFromImage(
model_config=model_config,
quantize=8,
lora_paths=lora_paths,

View File

@ -5,7 +5,7 @@ import numpy as np
from PIL import Image
from mflux import Config, ModelConfig
from mflux.community.in_context.flux_in_context_fill import FluxInContextFill
from mflux.community.in_context.flux_in_context_fill import Flux1InContextFill
from mflux.community.in_context.utils.in_context_loras import prepare_ic_edit_loras
@ -31,7 +31,7 @@ class ImageGeneratorICEditTestHelper:
try:
# given
flux = FluxInContextFill(
flux = Flux1InContextFill(
model_config=ModelConfig.dev_fill(),
quantize=8,
lora_paths=prepare_ic_edit_loras(lora_paths),

View File

@ -0,0 +1,66 @@
import os
from pathlib import Path
import numpy as np
from PIL import Image
from mflux import Config, ModelConfig
from mflux.kontext.flux_kontext import Flux1Kontext
class ImageGeneratorKontextTestHelper:
@staticmethod
def assert_matches_reference_image(
reference_image_path: str,
output_image_path: str,
model_config: ModelConfig,
steps: int,
seed: int,
height: int,
width: int,
prompt: str,
kontext_image_path: str,
guidance: float = 2.5,
):
# resolve paths
reference_image_path = ImageGeneratorKontextTestHelper.resolve_path(reference_image_path)
output_image_path = ImageGeneratorKontextTestHelper.resolve_path(output_image_path)
kontext_image_path = ImageGeneratorKontextTestHelper.resolve_path(kontext_image_path)
try:
# given
flux = Flux1Kontext(
quantize=8,
)
# when
image = flux.generate_image(
seed=seed,
prompt=prompt,
config=Config(
num_inference_steps=steps,
height=height,
width=width,
guidance=guidance,
image_path=kontext_image_path,
),
)
image.save(path=output_image_path)
# then
np.testing.assert_array_equal(
np.array(Image.open(output_image_path)),
np.array(Image.open(reference_image_path)),
err_msg=f"Generated image doesn't match reference image. Check {output_image_path} vs {reference_image_path}",
)
finally:
# cleanup
if os.path.exists(output_image_path):
os.remove(output_image_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,18 @@
from mflux import ModelConfig
from tests.image_generation.helpers.image_generation_kontext_test_helper import ImageGeneratorKontextTestHelper
class TestImageGeneratorKontext:
def test_image_generation_kontext(self):
ImageGeneratorKontextTestHelper.assert_matches_reference_image(
reference_image_path="reference_dev_kontext.png",
output_image_path="output_dev_kontext.png",
model_config=ModelConfig.dev_kontext(),
steps=20,
seed=4869845,
height=384,
width=640,
guidance=2.5,
prompt="Make the hand fistbump the camera instead of showing a flat palm",
kontext_image_path="reference_upscaled.png",
)

Binary file not shown.

After

Width:  |  Height:  |  Size: 272 KiB