diff --git a/.gitignore b/.gitignore index 15979672..3afebe2e 100644 --- a/.gitignore +++ b/.gitignore @@ -35,3 +35,6 @@ hosts_*.json # bench files bench/**/*.json + +# tmp +tmp/models diff --git a/.mlx_typings/mflux/models/flux/variants/kontext/__init__.pyi b/.mlx_typings/mflux/models/flux/variants/kontext/__init__.pyi new file mode 100644 index 00000000..36c5443a --- /dev/null +++ b/.mlx_typings/mflux/models/flux/variants/kontext/__init__.pyi @@ -0,0 +1,7 @@ +""" +This type stub file was generated by pyright. +""" + +from mflux.models.flux.variants.kontext.flux_kontext import Flux1Kontext + +__all__ = ["Flux1Kontext"] diff --git a/.mlx_typings/mflux/models/flux/variants/kontext/flux_kontext.pyi b/.mlx_typings/mflux/models/flux/variants/kontext/flux_kontext.pyi new file mode 100644 index 00000000..8050e68b --- /dev/null +++ b/.mlx_typings/mflux/models/flux/variants/kontext/flux_kontext.pyi @@ -0,0 +1,49 @@ +""" +This type stub file was generated by pyright. +""" + +from pathlib import Path +from typing import Any + +from mlx import nn + +from mflux.models.common.config.model_config import ModelConfig +from mflux.models.flux.model.flux_text_encoder.clip_encoder.clip_encoder import ( + CLIPEncoder, +) +from mflux.models.flux.model.flux_text_encoder.t5_encoder.t5_encoder import T5Encoder +from mflux.models.flux.model.flux_transformer.transformer import Transformer +from mflux.models.flux.model.flux_vae.vae import VAE +from mflux.utils.generated_image import GeneratedImage + +class Flux1Kontext(nn.Module): + vae: VAE + transformer: Transformer + t5_text_encoder: T5Encoder + clip_text_encoder: CLIPEncoder + bits: int | None + lora_paths: list[str] | None + lora_scales: list[float] | None + prompt_cache: dict[str, Any] + tokenizers: dict[str, Any] + + def __init__( + self, + quantize: int | None = ..., + model_path: str | None = ..., + lora_paths: list[str] | None = ..., + lora_scales: list[float] | None = ..., + model_config: ModelConfig = ..., + ) -> None: ... + def generate_image( + self, + seed: int, + prompt: str, + num_inference_steps: int = ..., + height: int = ..., + width: int = ..., + guidance: float = ..., + image_path: Path | str | None = ..., + image_strength: float | None = ..., + scheduler: str = ..., + ) -> GeneratedImage: ... diff --git a/.mlx_typings/mflux/models/flux/variants/kontext/kontext_util.pyi b/.mlx_typings/mflux/models/flux/variants/kontext/kontext_util.pyi new file mode 100644 index 00000000..c7588ec6 --- /dev/null +++ b/.mlx_typings/mflux/models/flux/variants/kontext/kontext_util.pyi @@ -0,0 +1,16 @@ +""" +This type stub file was generated by pyright. +""" + +import mlx.core as mx + +from mflux.models.flux.model.flux_vae.vae import VAE + +class KontextUtil: + @staticmethod + def create_image_conditioning_latents( + vae: VAE, + height: int, + width: int, + image_path: str, + ) -> tuple[mx.array, mx.array]: ... diff --git a/pyproject.toml b/pyproject.toml index dd84cf17..ee219b74 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -26,7 +26,7 @@ dependencies = [ "httpx>=0.28.1", "tomlkit>=0.14.0", "pillow>=11.0,<12.0", # compatibility with mflux - "mflux==0.15.4", + "mflux==0.15.5", "python-multipart>=0.0.21", ] diff --git a/resources/image_model_cards/exolabs--FLUX.1-Kontext-dev-4bit.toml b/resources/image_model_cards/exolabs--FLUX.1-Kontext-dev-4bit.toml new file mode 100644 index 00000000..baf305ca --- /dev/null +++ b/resources/image_model_cards/exolabs--FLUX.1-Kontext-dev-4bit.toml @@ -0,0 +1,45 @@ +model_id = "exolabs/FLUX.1-Kontext-dev-4bit" +n_layers = 57 +hidden_size = 1 +supports_tensor = false +tasks = ["ImageToImage"] + +[storage_size] +in_bytes = 15475325472 + +[[components]] +component_name = "text_encoder" +component_path = "text_encoder/" +n_layers = 12 +can_shard = false + +[components.storage_size] +in_bytes = 0 + +[[components]] +component_name = "text_encoder_2" +component_path = "text_encoder_2/" +n_layers = 24 +can_shard = false +safetensors_index_filename = "model.safetensors.index.json" + +[components.storage_size] +in_bytes = 9524621312 + +[[components]] +component_name = "transformer" +component_path = "transformer/" +n_layers = 57 +can_shard = true +safetensors_index_filename = "diffusion_pytorch_model.safetensors.index.json" + +[components.storage_size] +in_bytes = 5950704160 + +[[components]] +component_name = "vae" +component_path = "vae/" +can_shard = false + +[components.storage_size] +in_bytes = 0 diff --git a/resources/image_model_cards/exolabs--FLUX.1-Kontext-dev-8bit.toml b/resources/image_model_cards/exolabs--FLUX.1-Kontext-dev-8bit.toml new file mode 100644 index 00000000..ce0809c2 --- /dev/null +++ b/resources/image_model_cards/exolabs--FLUX.1-Kontext-dev-8bit.toml @@ -0,0 +1,45 @@ +model_id = "exolabs/FLUX.1-Kontext-dev-8bit" +n_layers = 57 +hidden_size = 1 +supports_tensor = false +tasks = ["ImageToImage"] + +[storage_size] +in_bytes = 21426029632 + +[[components]] +component_name = "text_encoder" +component_path = "text_encoder/" +n_layers = 12 +can_shard = false + +[components.storage_size] +in_bytes = 0 + +[[components]] +component_name = "text_encoder_2" +component_path = "text_encoder_2/" +n_layers = 24 +can_shard = false +safetensors_index_filename = "model.safetensors.index.json" + +[components.storage_size] +in_bytes = 9524621312 + +[[components]] +component_name = "transformer" +component_path = "transformer/" +n_layers = 57 +can_shard = true +safetensors_index_filename = "diffusion_pytorch_model.safetensors.index.json" + +[components.storage_size] +in_bytes = 11901408320 + +[[components]] +component_name = "vae" +component_path = "vae/" +can_shard = false + +[components.storage_size] +in_bytes = 0 diff --git a/resources/image_model_cards/exolabs--FLUX.1-Kontext-dev.toml b/resources/image_model_cards/exolabs--FLUX.1-Kontext-dev.toml new file mode 100644 index 00000000..2ebb0c43 --- /dev/null +++ b/resources/image_model_cards/exolabs--FLUX.1-Kontext-dev.toml @@ -0,0 +1,45 @@ +model_id = "exolabs/FLUX.1-Kontext-dev" +n_layers = 57 +hidden_size = 1 +supports_tensor = false +tasks = ["ImageToImage"] + +[storage_size] +in_bytes = 33327437952 + +[[components]] +component_name = "text_encoder" +component_path = "text_encoder/" +n_layers = 12 +can_shard = false + +[components.storage_size] +in_bytes = 0 + +[[components]] +component_name = "text_encoder_2" +component_path = "text_encoder_2/" +n_layers = 24 +can_shard = false +safetensors_index_filename = "model.safetensors.index.json" + +[components.storage_size] +in_bytes = 9524621312 + +[[components]] +component_name = "transformer" +component_path = "transformer/" +n_layers = 57 +can_shard = true +safetensors_index_filename = "diffusion_pytorch_model.safetensors.index.json" + +[components.storage_size] +in_bytes = 23802816640 + +[[components]] +component_name = "vae" +component_path = "vae/" +can_shard = false + +[components.storage_size] +in_bytes = 0 diff --git a/src/exo/worker/engines/image/models/__init__.py b/src/exo/worker/engines/image/models/__init__.py index dc0a9d8c..9f58cf9d 100644 --- a/src/exo/worker/engines/image/models/__init__.py +++ b/src/exo/worker/engines/image/models/__init__.py @@ -5,7 +5,9 @@ from exo.worker.engines.image.config import ImageModelConfig from exo.worker.engines.image.models.base import ModelAdapter from exo.worker.engines.image.models.flux import ( FLUX_DEV_CONFIG, + FLUX_KONTEXT_CONFIG, FLUX_SCHNELL_CONFIG, + FluxKontextModelAdapter, FluxModelAdapter, ) from exo.worker.engines.image.models.qwen import ( @@ -26,13 +28,16 @@ AdapterFactory = Callable[ # Registry maps model_family string to adapter factory _ADAPTER_REGISTRY: dict[str, AdapterFactory] = { "flux": FluxModelAdapter, + "flux-kontext": FluxKontextModelAdapter, "qwen-edit": QwenEditModelAdapter, "qwen": QwenModelAdapter, } # Config registry: maps model ID patterns to configs +# Order matters: longer/more-specific patterns must come before shorter ones _CONFIG_REGISTRY: dict[str, ImageModelConfig] = { "flux.1-schnell": FLUX_SCHNELL_CONFIG, + "flux.1-kontext": FLUX_KONTEXT_CONFIG, # Must come before "flux.1-dev" for pattern matching "flux.1-krea-dev": FLUX_DEV_CONFIG, # Must come before "flux.1-dev" for pattern matching "flux.1-dev": FLUX_DEV_CONFIG, "qwen-image-edit": QWEN_IMAGE_EDIT_CONFIG, # Must come before "qwen-image" for pattern matching diff --git a/src/exo/worker/engines/image/models/base.py b/src/exo/worker/engines/image/models/base.py index f77ea882..0aa660e2 100644 --- a/src/exo/worker/engines/image/models/base.py +++ b/src/exo/worker/engines/image/models/base.py @@ -66,6 +66,19 @@ class PromptData(ABC): """ ... + @property + @abstractmethod + def kontext_image_ids(self) -> mx.array | None: + """Kontext-style position IDs for image conditioning. + + For FLUX.1-Kontext models, returns position IDs with first_coord=1 + to distinguish conditioning tokens from generation tokens (first_coord=0). + + Returns: + Position IDs array [1, seq_len, 3] for Kontext, None for other models. + """ + ... + @abstractmethod def get_batched_cfg_data( self, diff --git a/src/exo/worker/engines/image/models/flux/__init__.py b/src/exo/worker/engines/image/models/flux/__init__.py index 3adc2626..ac3e9335 100644 --- a/src/exo/worker/engines/image/models/flux/__init__.py +++ b/src/exo/worker/engines/image/models/flux/__init__.py @@ -1,11 +1,17 @@ from exo.worker.engines.image.models.flux.adapter import FluxModelAdapter from exo.worker.engines.image.models.flux.config import ( FLUX_DEV_CONFIG, + FLUX_KONTEXT_CONFIG, FLUX_SCHNELL_CONFIG, ) +from exo.worker.engines.image.models.flux.kontext_adapter import ( + FluxKontextModelAdapter, +) __all__ = [ "FluxModelAdapter", + "FluxKontextModelAdapter", "FLUX_DEV_CONFIG", + "FLUX_KONTEXT_CONFIG", "FLUX_SCHNELL_CONFIG", ] diff --git a/src/exo/worker/engines/image/models/flux/adapter.py b/src/exo/worker/engines/image/models/flux/adapter.py index 1aa510da..90b7bacd 100644 --- a/src/exo/worker/engines/image/models/flux/adapter.py +++ b/src/exo/worker/engines/image/models/flux/adapter.py @@ -59,6 +59,10 @@ class FluxPromptData(PromptData): def conditioning_latents(self) -> mx.array | None: return None + @property + def kontext_image_ids(self) -> mx.array | None: + return None + def get_batched_cfg_data( self, ) -> tuple[mx.array, mx.array, mx.array | None, mx.array | None] | None: diff --git a/src/exo/worker/engines/image/models/flux/config.py b/src/exo/worker/engines/image/models/flux/config.py index d37161b0..a85bd0c8 100644 --- a/src/exo/worker/engines/image/models/flux/config.py +++ b/src/exo/worker/engines/image/models/flux/config.py @@ -32,3 +32,19 @@ FLUX_DEV_CONFIG = ImageModelConfig( default_steps={"low": 10, "medium": 25, "high": 50}, num_sync_steps=4, ) + + +FLUX_KONTEXT_CONFIG = ImageModelConfig( + model_family="flux-kontext", + block_configs=( + TransformerBlockConfig( + block_type=BlockType.JOINT, count=19, has_separate_text_output=True + ), + TransformerBlockConfig( + block_type=BlockType.SINGLE, count=38, has_separate_text_output=False + ), + ), + default_steps={"low": 10, "medium": 25, "high": 50}, + num_sync_steps=4, + guidance_scale=4.0, +) diff --git a/src/exo/worker/engines/image/models/flux/kontext_adapter.py b/src/exo/worker/engines/image/models/flux/kontext_adapter.py new file mode 100644 index 00000000..19d0be56 --- /dev/null +++ b/src/exo/worker/engines/image/models/flux/kontext_adapter.py @@ -0,0 +1,348 @@ +import math +from pathlib import Path +from typing import Any, final + +import mlx.core as mx +from mflux.models.common.config.config import Config +from mflux.models.common.config.model_config import ModelConfig +from mflux.models.flux.latent_creator.flux_latent_creator import FluxLatentCreator +from mflux.models.flux.model.flux_text_encoder.prompt_encoder import PromptEncoder +from mflux.models.flux.model.flux_transformer.transformer import Transformer +from mflux.models.flux.variants.kontext.flux_kontext import Flux1Kontext +from mflux.models.flux.variants.kontext.kontext_util import KontextUtil + +from exo.worker.engines.image.config import ImageModelConfig +from exo.worker.engines.image.models.base import ( + ModelAdapter, + PromptData, + RotaryEmbeddings, +) +from exo.worker.engines.image.models.flux.wrappers import ( + FluxJointBlockWrapper, + FluxSingleBlockWrapper, +) +from exo.worker.engines.image.pipeline.block_wrapper import ( + JointBlockWrapper, + SingleBlockWrapper, +) + + +@final +class FluxKontextPromptData(PromptData): + """Prompt data for FLUX.1-Kontext image editing. + + Stores text embeddings along with conditioning latents and position IDs + for the input image. + """ + + def __init__( + self, + prompt_embeds: mx.array, + pooled_prompt_embeds: mx.array, + conditioning_latents: mx.array, + kontext_image_ids: mx.array, + ): + self._prompt_embeds = prompt_embeds + self._pooled_prompt_embeds = pooled_prompt_embeds + self._conditioning_latents = conditioning_latents + self._kontext_image_ids = kontext_image_ids + + @property + def prompt_embeds(self) -> mx.array: + return self._prompt_embeds + + @property + def pooled_prompt_embeds(self) -> mx.array: + return self._pooled_prompt_embeds + + @property + def negative_prompt_embeds(self) -> mx.array | None: + return None + + @property + def negative_pooled_prompt_embeds(self) -> mx.array | None: + return None + + def get_encoder_hidden_states_mask(self, positive: bool = True) -> mx.array | None: + return None + + @property + def cond_image_grid( + self, + ) -> tuple[int, int, int] | list[tuple[int, int, int]] | None: + return None + + @property + def conditioning_latents(self) -> mx.array | None: + """VAE-encoded input image latents for Kontext conditioning.""" + return self._conditioning_latents + + @property + def kontext_image_ids(self) -> mx.array | None: + """Position IDs for Kontext conditioning (first_coord=1).""" + return self._kontext_image_ids + + def get_cfg_branch_data( + self, positive: bool + ) -> tuple[mx.array, mx.array | None, mx.array | None, mx.array | None]: + """Kontext doesn't use CFG, but we return positive data for compatibility.""" + return ( + self._prompt_embeds, + None, + self._pooled_prompt_embeds, + self._conditioning_latents, + ) + + def get_batched_cfg_data( + self, + ) -> tuple[mx.array, mx.array, mx.array | None, mx.array | None] | None: + # Kontext doesn't use CFG + return None + + +@final +class FluxKontextModelAdapter(ModelAdapter[Flux1Kontext, Transformer]): + """Adapter for FLUX.1-Kontext image editing model. + + Key differences from standard FluxModelAdapter: + - Takes an input image and computes output dimensions from it + - Creates conditioning latents from the input image via VAE + - Creates special position IDs (kontext_image_ids) for conditioning tokens + - Creates pure noise latents (not img2img blending) + """ + + def __init__( + self, + config: ImageModelConfig, + model_id: str, + local_path: Path, + quantize: int | None = None, + ): + self._config = config + self._model = Flux1Kontext( + model_config=ModelConfig.from_name(model_name=model_id, base_model=None), + model_path=str(local_path), + quantize=quantize, + ) + self._transformer = self._model.transformer + + # Stores image path and computed dimensions after set_image_dimensions + self._image_path: str | None = None + self._output_height: int | None = None + self._output_width: int | None = None + + @property + def hidden_dim(self) -> int: + return self._transformer.x_embedder.weight.shape[0] # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType] + + @property + def needs_cfg(self) -> bool: + return False + + def _get_latent_creator(self) -> type: + return FluxLatentCreator + + def get_joint_block_wrappers( + self, + text_seq_len: int, + encoder_hidden_states_mask: mx.array | None = None, + ) -> list[JointBlockWrapper[Any]]: + """Create wrapped joint blocks for Flux Kontext.""" + return [ + FluxJointBlockWrapper(block, text_seq_len) + for block in self._transformer.transformer_blocks + ] + + def get_single_block_wrappers( + self, + text_seq_len: int, + ) -> list[SingleBlockWrapper[Any]]: + """Create wrapped single blocks for Flux Kontext.""" + return [ + FluxSingleBlockWrapper(block, text_seq_len) + for block in self._transformer.single_transformer_blocks + ] + + def slice_transformer_blocks( + self, + start_layer: int, + end_layer: int, + ): + all_joint = list(self._transformer.transformer_blocks) + all_single = list(self._transformer.single_transformer_blocks) + total_joint_blocks = len(all_joint) + if end_layer <= total_joint_blocks: + # All assigned are joint blocks + joint_start, joint_end = start_layer, end_layer + single_start, single_end = 0, 0 + elif start_layer >= total_joint_blocks: + # All assigned are single blocks + joint_start, joint_end = 0, 0 + single_start = start_layer - total_joint_blocks + single_end = end_layer - total_joint_blocks + else: + # Spans both joint and single + joint_start, joint_end = start_layer, total_joint_blocks + single_start = 0 + single_end = end_layer - total_joint_blocks + + self._transformer.transformer_blocks = all_joint[joint_start:joint_end] + self._transformer.single_transformer_blocks = all_single[ + single_start:single_end + ] + + def set_image_dimensions(self, image_path: Path) -> tuple[int, int]: + """Compute and store dimensions from input image. + + Also stores image_path for use in encode_prompt(). + + Args: + image_path: Path to the input image + + Returns: + (output_width, output_height) for runtime config + """ + from mflux.utils.image_util import ImageUtil + + pil_image = ImageUtil.load_image(str(image_path)).convert("RGB") + image_size = pil_image.size + + # Compute output dimensions from input image aspect ratio + # Target area of 1024x1024 = ~1M pixels + target_area = 1024 * 1024 + ratio = image_size[0] / image_size[1] + output_width = math.sqrt(target_area * ratio) + output_height = output_width / ratio + output_width = round(output_width / 32) * 32 + output_height = round(output_height / 32) * 32 + + # Ensure multiple of 16 for VAE + vae_scale_factor = 8 + multiple_of = vae_scale_factor * 2 + output_width = output_width // multiple_of * multiple_of + output_height = output_height // multiple_of * multiple_of + + self._image_path = str(image_path) + self._output_width = int(output_width) + self._output_height = int(output_height) + + return self._output_width, self._output_height + + def create_latents(self, seed: int, runtime_config: Config) -> mx.array: + """Create initial noise latents for Kontext. + + Unlike standard img2img which blends noise with encoded input, + Kontext uses pure noise latents. The input image is provided + separately as conditioning. + """ + return FluxLatentCreator.create_noise( + seed=seed, + height=runtime_config.height, + width=runtime_config.width, + ) + + def encode_prompt( + self, prompt: str, negative_prompt: str | None = None + ) -> FluxKontextPromptData: + """Encode prompt and create conditioning from stored input image. + + Must call set_image_dimensions() before this method. + + Args: + prompt: Text prompt for editing + negative_prompt: Ignored (Kontext doesn't use CFG) + + Returns: + FluxKontextPromptData with text embeddings and image conditioning + """ + del negative_prompt # Kontext doesn't support negative prompts or CFG + + if ( + self._image_path is None + or self._output_height is None + or self._output_width is None + ): + raise RuntimeError( + "set_image_dimensions() must be called before encode_prompt() " + "for FluxKontextModelAdapter" + ) + + assert isinstance(self.model.prompt_cache, dict) + assert isinstance(self.model.tokenizers, dict) + + # Encode text prompt + prompt_embeds, pooled_prompt_embeds = PromptEncoder.encode_prompt( + prompt=prompt, + prompt_cache=self.model.prompt_cache, + t5_tokenizer=self.model.tokenizers["t5"], # pyright: ignore[reportAny] + clip_tokenizer=self.model.tokenizers["clip"], # pyright: ignore[reportAny] + t5_text_encoder=self.model.t5_text_encoder, + clip_text_encoder=self.model.clip_text_encoder, + ) + + # Create conditioning latents from input image + conditioning_latents, kontext_image_ids = ( + KontextUtil.create_image_conditioning_latents( + vae=self.model.vae, + height=self._output_height, + width=self._output_width, + image_path=self._image_path, + ) + ) + + return FluxKontextPromptData( + prompt_embeds=prompt_embeds, + pooled_prompt_embeds=pooled_prompt_embeds, + conditioning_latents=conditioning_latents, + kontext_image_ids=kontext_image_ids, + ) + + def compute_embeddings( + self, + hidden_states: mx.array, + prompt_embeds: mx.array, + ) -> tuple[mx.array, mx.array]: + embedded_hidden = self._transformer.x_embedder(hidden_states) + embedded_encoder = self._transformer.context_embedder(prompt_embeds) + return embedded_hidden, embedded_encoder + + def compute_text_embeddings( + self, + t: int, + runtime_config: Config, + pooled_prompt_embeds: mx.array | None = None, + hidden_states: mx.array | None = None, + ) -> mx.array: + if pooled_prompt_embeds is None: + raise ValueError( + "pooled_prompt_embeds is required for Flux Kontext text embeddings" + ) + + return Transformer.compute_text_embeddings( + t, pooled_prompt_embeds, self._transformer.time_text_embed, runtime_config + ) + + def compute_rotary_embeddings( + self, + prompt_embeds: mx.array, + runtime_config: Config, + encoder_hidden_states_mask: mx.array | None = None, + cond_image_grid: tuple[int, int, int] + | list[tuple[int, int, int]] + | None = None, + kontext_image_ids: mx.array | None = None, + ) -> RotaryEmbeddings: + return Transformer.compute_rotary_embeddings( + prompt_embeds, + self._transformer.pos_embed, + runtime_config, + kontext_image_ids, + ) + + def apply_guidance( + self, + noise_positive: mx.array, + noise_negative: mx.array, + guidance_scale: float, + ) -> mx.array: + raise NotImplementedError("Flux Kontext does not use classifier-free guidance") diff --git a/src/exo/worker/engines/image/models/qwen/adapter.py b/src/exo/worker/engines/image/models/qwen/adapter.py index e88d2a75..be1edb0c 100644 --- a/src/exo/worker/engines/image/models/qwen/adapter.py +++ b/src/exo/worker/engines/image/models/qwen/adapter.py @@ -69,6 +69,10 @@ class QwenPromptData(PromptData): def conditioning_latents(self) -> mx.array | None: return None + @property + def kontext_image_ids(self) -> mx.array | None: + return None + def get_batched_cfg_data( self, ) -> tuple[mx.array, mx.array, mx.array | None, mx.array | None] | None: diff --git a/src/exo/worker/engines/image/models/qwen/edit_adapter.py b/src/exo/worker/engines/image/models/qwen/edit_adapter.py index 4a88a4e3..fee79738 100644 --- a/src/exo/worker/engines/image/models/qwen/edit_adapter.py +++ b/src/exo/worker/engines/image/models/qwen/edit_adapter.py @@ -85,6 +85,10 @@ class QwenEditPromptData(PromptData): def qwen_image_ids(self) -> mx.array: return self._qwen_image_ids + @property + def kontext_image_ids(self) -> mx.array | None: + return None + @property def is_edit_mode(self) -> bool: return True diff --git a/src/exo/worker/engines/image/pipeline/runner.py b/src/exo/worker/engines/image/pipeline/runner.py index f7054763..01ad7b05 100644 --- a/src/exo/worker/engines/image/pipeline/runner.py +++ b/src/exo/worker/engines/image/pipeline/runner.py @@ -567,6 +567,7 @@ class DiffusionRunner: | list[tuple[int, int, int]] | None = None, conditioning_latents: mx.array | None = None, + kontext_image_ids: mx.array | None = None, ) -> mx.array: """Run a single forward pass through the transformer. Args: @@ -578,6 +579,7 @@ class DiffusionRunner: encoder_hidden_states_mask: Attention mask for text (Qwen) cond_image_grid: Conditioning image grid dimensions (Qwen edit) conditioning_latents: Conditioning latents for edit mode + kontext_image_ids: Position IDs for Kontext conditioning (Flux Kontext) Returns: Noise prediction tensor @@ -610,6 +612,7 @@ class DiffusionRunner: config, encoder_hidden_states_mask=encoder_hidden_states_mask, cond_image_grid=cond_image_grid, + kontext_image_ids=kontext_image_ids, ) assert self.joint_block_wrappers is not None @@ -681,6 +684,7 @@ class DiffusionRunner: prompt_data: PromptData, ) -> mx.array: cond_image_grid = prompt_data.cond_image_grid + kontext_image_ids = prompt_data.kontext_image_ids results: list[tuple[bool, mx.array]] = [] for branch in self._get_cfg_branches(prompt_data): @@ -700,6 +704,7 @@ class DiffusionRunner: encoder_hidden_states_mask=branch.mask, cond_image_grid=cond_image_grid, conditioning_latents=branch.cond_latents, + kontext_image_ids=kontext_image_ids, ) results.append((branch.positive, noise)) @@ -902,10 +907,10 @@ class DiffusionRunner: config: Config, hidden_states: mx.array, prompt_data: PromptData, - kontext_image_ids: mx.array | None = None, ) -> mx.array: prev_latents = hidden_states cond_image_grid = prompt_data.cond_image_grid + kontext_image_ids = prompt_data.kontext_image_ids scaled_hidden_states = config.scheduler.scale_model_input(hidden_states, t) # pyright: ignore[reportAny] original_latent_tokens: int = scaled_hidden_states.shape[1] # pyright: ignore[reportAny] @@ -979,10 +984,10 @@ class DiffusionRunner: latents: mx.array, prompt_data: PromptData, is_first_async_step: bool, - kontext_image_ids: mx.array | None = None, ) -> mx.array: patch_latents, token_indices = self._create_patches(latents, config) cond_image_grid = prompt_data.cond_image_grid + kontext_image_ids = prompt_data.kontext_image_ids prev_patch_latents = [p for p in patch_latents] diff --git a/tmp/quantize_and_upload.py b/tmp/quantize_and_upload.py new file mode 100755 index 00000000..ee421a14 --- /dev/null +++ b/tmp/quantize_and_upload.py @@ -0,0 +1,377 @@ +#!/usr/bin/env python3 +""" +Download an mflux model, quantize it, and upload to HuggingFace. + +Usage (run from mflux project directory): + cd /path/to/mflux + uv run python /path/to/quantize_and_upload.py --model black-forest-labs/FLUX.1-Kontext-dev + uv run python /path/to/quantize_and_upload.py --model black-forest-labs/FLUX.1-Kontext-dev --skip-base --skip-8bit + uv run python /path/to/quantize_and_upload.py --model black-forest-labs/FLUX.1-Kontext-dev --dry-run + +Requires: + - Must be run from mflux project directory using `uv run` + - huggingface_hub installed (add to mflux deps or install separately) + - HuggingFace authentication: run `huggingface-cli login` or set HF_TOKEN +""" + +from __future__ import annotations + +import argparse +import re +import shutil +import sys +from pathlib import Path +from typing import TYPE_CHECKING + +if TYPE_CHECKING: + from mflux.models.flux.variants.txt2img.flux import Flux1 + + +HF_ORG = "exolabs" + + +def get_model_class(model_name: str) -> type: + """Get the appropriate model class based on model name.""" + from mflux.models.fibo.variants.txt2img.fibo import FIBO + from mflux.models.flux.variants.txt2img.flux import Flux1 + from mflux.models.flux2.variants.txt2img.flux2_klein import Flux2Klein + from mflux.models.qwen.variants.txt2img.qwen_image import QwenImage + from mflux.models.z_image.variants.turbo.z_image_turbo import ZImageTurbo + + model_name_lower = model_name.lower() + if "qwen" in model_name_lower: + return QwenImage + elif "fibo" in model_name_lower: + return FIBO + elif "z-image" in model_name_lower or "zimage" in model_name_lower: + return ZImageTurbo + elif "flux2" in model_name_lower or "flux.2" in model_name_lower: + return Flux2Klein + else: + return Flux1 + + +def get_repo_name(model_name: str, bits: int | None) -> str: + """Get the HuggingFace repo name for a model variant.""" + # Extract repo name from HF path (e.g., "black-forest-labs/FLUX.1-Kontext-dev" -> "FLUX.1-Kontext-dev") + base_name = model_name.split("/")[-1] if "/" in model_name else model_name + suffix = f"-{bits}bit" if bits else "" + return f"{HF_ORG}/{base_name}{suffix}" + + +def get_local_path(output_dir: Path, model_name: str, bits: int | None) -> Path: + """Get the local save path for a model variant.""" + # Extract repo name from HF path (e.g., "black-forest-labs/FLUX.1-Kontext-dev" -> "FLUX.1-Kontext-dev") + base_name = model_name.split("/")[-1] if "/" in model_name else model_name + suffix = f"-{bits}bit" if bits else "" + return output_dir / f"{base_name}{suffix}" + + +def copy_source_repo( + source_repo: str, + local_path: Path, + dry_run: bool = False, +) -> None: + """Copy all files from source repo (replicating original HF structure).""" + print(f"\n{'=' * 60}") + print(f"Copying full repo from source: {source_repo}") + print(f"Output path: {local_path}") + print(f"{'=' * 60}") + + if dry_run: + print("[DRY RUN] Would download all files from source repo") + return + + from huggingface_hub import snapshot_download + + # Download all files to our local path + snapshot_download( + repo_id=source_repo, + local_dir=local_path, + ) + + # Remove root-level safetensors files (flux.1-dev.safetensors, etc.) + # These are redundant with the component directories + for f in local_path.glob("*.safetensors"): + print(f"Removing root-level safetensors: {f.name}") + if not dry_run: + f.unlink() + + print(f"Source repo copied to {local_path}") + + +def load_and_save_quantized_model( + model_name: str, + bits: int, + output_path: Path, + dry_run: bool = False, +) -> None: + """Load a model with quantization and save it in mflux format.""" + print(f"\n{'=' * 60}") + print(f"Loading {model_name} with {bits}-bit quantization...") + print(f"Output path: {output_path}") + print(f"{'=' * 60}") + + if dry_run: + print("[DRY RUN] Would load and save quantized model") + return + + from mflux.models.common.config.model_config import ModelConfig + + model_class = get_model_class(model_name) + model_config = ModelConfig.from_name(model_name=model_name, base_model=None) + + model: Flux1 = model_class( + quantize=bits, + model_config=model_config, + ) + + print(f"Saving model to {output_path}...") + model.save_model(str(output_path)) + print(f"Model saved successfully to {output_path}") + + +def copy_source_metadata( + source_repo: str, + local_path: Path, + dry_run: bool = False, +) -> None: + """Copy metadata files (LICENSE, README, etc.) from source repo, excluding safetensors.""" + print(f"\n{'=' * 60}") + print(f"Copying metadata from source repo: {source_repo}") + print(f"{'=' * 60}") + + if dry_run: + print("[DRY RUN] Would download metadata files (excluding *.safetensors)") + return + + from huggingface_hub import snapshot_download + + # Download all files except safetensors to our local path + snapshot_download( + repo_id=source_repo, + local_dir=local_path, + ignore_patterns=["*.safetensors"], + ) + print(f"Metadata files copied to {local_path}") + + +def upload_to_huggingface( + local_path: Path, + repo_id: str, + dry_run: bool = False, + clean_remote: bool = False, +) -> None: + """Upload a saved model to HuggingFace.""" + print(f"\n{'=' * 60}") + print(f"Uploading to HuggingFace: {repo_id}") + print(f"Local path: {local_path}") + print(f"Clean remote first: {clean_remote}") + print(f"{'=' * 60}") + + if dry_run: + print("[DRY RUN] Would upload to HuggingFace") + return + + from huggingface_hub import HfApi + + api = HfApi() + + # Create the repo if it doesn't exist + print(f"Creating/verifying repo: {repo_id}") + api.create_repo(repo_id=repo_id, repo_type="model", exist_ok=True) + + # Clean remote repo if requested (delete old mflux-format files) + if clean_remote: + print("Cleaning old mflux-format files from remote...") + try: + # Pattern for mflux numbered shards: