From 23fd37fe4d684ec5d90f7a4b5b03b3a829c64fa0 Mon Sep 17 00:00:00 2001 From: ciaranbor <81697641+ciaranbor@users.noreply.github.com> Date: Fri, 23 Jan 2026 19:48:24 +0000 Subject: [PATCH] Add FLUX.1-Krea-dev model (#1269) ## Why It Works Same implementation as FLUX.1-dev, just different weights --- pyproject.toml | 2 +- src/exo/download/download_utils.py | 15 +++++++ src/exo/shared/models/model_cards.py | 42 +++++++++++++++++++ .../worker/engines/image/models/__init__.py | 1 + uv.lock | 27 ++++++------ 5 files changed, 72 insertions(+), 15 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index 2dc2453a..702e198d 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.14.2", + "mflux==0.15.4", "python-multipart>=0.0.21", ] diff --git a/src/exo/download/download_utils.py b/src/exo/download/download_utils.py index f08aa0ed..a721390f 100644 --- a/src/exo/download/download_utils.py +++ b/src/exo/download/download_utils.py @@ -32,6 +32,7 @@ from exo.download.huggingface_utils import ( get_hf_token, ) from exo.shared.constants import EXO_MODELS_DIR +from exo.shared.models.model_cards import ModelTask from exo.shared.types.common import ModelId from exo.shared.types.memory import Memory from exo.shared.types.worker.downloads import ( @@ -481,6 +482,11 @@ async def resolve_allow_patterns(shard: ShardMetadata) -> list[str]: return ["*"] +def is_image_model(shard: ShardMetadata) -> bool: + tasks = shard.model_card.tasks + return ModelTask.TextToImage in tasks or ModelTask.ImageToImage in tasks + + async def get_downloaded_size(path: Path) -> int: partial_path = path.with_suffix(path.suffix + ".partial") if await aios.path.exists(path): @@ -522,6 +528,15 @@ async def download_shard( file_list, allow_patterns=allow_patterns, key=lambda x: x.path ) ) + + # For image models, skip root-level safetensors files since weights + # are stored in component subdirectories (e.g., transformer/, vae/) + if is_image_model(shard): + filtered_file_list = [ + f + for f in filtered_file_list + if "/" in f.path or not f.path.endswith(".safetensors") + ] file_progress: dict[str, RepoFileDownloadProgress] = {} async def on_progress_wrapper( diff --git a/src/exo/shared/models/model_cards.py b/src/exo/shared/models/model_cards.py index 35e077ba..1d09293a 100644 --- a/src/exo/shared/models/model_cards.py +++ b/src/exo/shared/models/model_cards.py @@ -498,6 +498,48 @@ _IMAGE_MODEL_CARDS: dict[str, ModelCard] = { ), ], ), + "flux1-krea-dev": ModelCard( + model_id=ModelId("black-forest-labs/FLUX.1-Krea-dev"), + storage_size=Memory.from_bytes(23802816640 + 9524621312), # Same as dev + n_layers=57, + hidden_size=1, + supports_tensor=False, + tasks=[ModelTask.TextToImage], + components=[ + ComponentInfo( + component_name="text_encoder", + component_path="text_encoder/", + storage_size=Memory.from_kb(0), + n_layers=12, + can_shard=False, + safetensors_index_filename=None, + ), + ComponentInfo( + component_name="text_encoder_2", + component_path="text_encoder_2/", + storage_size=Memory.from_bytes(9524621312), + n_layers=24, + can_shard=False, + safetensors_index_filename="model.safetensors.index.json", + ), + ComponentInfo( + component_name="transformer", + component_path="transformer/", + storage_size=Memory.from_bytes(23802816640), + n_layers=57, + can_shard=True, + safetensors_index_filename="diffusion_pytorch_model.safetensors.index.json", + ), + ComponentInfo( + component_name="vae", + component_path="vae/", + storage_size=Memory.from_kb(0), + n_layers=None, + can_shard=False, + safetensors_index_filename=None, + ), + ], + ), "qwen-image": ModelCard( model_id=ModelId("Qwen/Qwen-Image"), storage_size=Memory.from_bytes(16584333312 + 40860802176), diff --git a/src/exo/worker/engines/image/models/__init__.py b/src/exo/worker/engines/image/models/__init__.py index b205af60..dc0a9d8c 100644 --- a/src/exo/worker/engines/image/models/__init__.py +++ b/src/exo/worker/engines/image/models/__init__.py @@ -33,6 +33,7 @@ _ADAPTER_REGISTRY: dict[str, AdapterFactory] = { # Config registry: maps model ID patterns to configs _CONFIG_REGISTRY: dict[str, ImageModelConfig] = { "flux.1-schnell": FLUX_SCHNELL_CONFIG, + "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 "qwen-image": QWEN_IMAGE_CONFIG, diff --git a/uv.lock b/uv.lock index ac4bf40a..c81ef825 100644 --- a/uv.lock +++ b/uv.lock @@ -412,7 +412,7 @@ requires-dist = [ { name = "huggingface-hub", specifier = ">=0.33.4" }, { name = "hypercorn", specifier = ">=0.18.0" }, { name = "loguru", specifier = ">=0.7.3" }, - { name = "mflux", specifier = ">=0.14.2" }, + { name = "mflux", specifier = "==0.15.4" }, { name = "mlx", marker = "sys_platform == 'darwin'", specifier = "==0.30.3" }, { name = "mlx", extras = ["cpu"], marker = "sys_platform == 'linux'", specifier = "==0.30.3" }, { name = "mlx-lm", git = "https://github.com/AlexCheema/mlx-lm.git?rev=fix-transformers-5.0.0rc2" }, @@ -458,16 +458,6 @@ dev = [ { name = "pytest-asyncio", specifier = ">=1.0.0" }, ] -[[package]] -name = "tomlkit" -version = "0.14.0" -source = { registry = "https://pypi.org/simple" } -sdist = { url = "https://files.pythonhosted.org/packages/c3/af/14b24e41977adb296d6bd1fb59402cf7d60ce364f90c890bd2ec65c43b5a/tomlkit-0.14.0.tar.gz", hash = "sha256:cf00efca415dbd57575befb1f6634c4f42d2d87dbba376128adb42c121b87064", size = 187167 } -wheels = [ - { url = "https://files.pythonhosted.org/packages/b5/11/87d6d29fb5d237229d67973a6c9e06e048f01cf4994dee194ab0ea841814/tomlkit-0.14.0-py3-none-any.whl", hash = "sha256:592064ed85b40fa213469f81ac584f67a4f2992509a7c3ea2d632208623a3680", size = 39310 }, -] - - [[package]] name = "fastapi" version = "0.128.0" @@ -997,7 +987,7 @@ wheels = [ [[package]] name = "mflux" -version = "0.15.3" +version = "0.15.4" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "filelock", marker = "sys_platform == 'darwin' or sys_platform == 'linux'" }, @@ -1023,9 +1013,9 @@ dependencies = [ { name = "twine", marker = "sys_platform == 'darwin' or sys_platform == 'linux'" }, { name = "urllib3", marker = "sys_platform == 'darwin' or sys_platform == 'linux'" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/23/c5/dd12e16714702255d89b7ccc6f217c405a9fdcf2af950a2236892c50a219/mflux-0.15.3.tar.gz", hash = "sha256:e32ea66a81aad4f77eea2415b17c27fc3d9ce662a842565c62871ff570f4ef2f", size = 740701, upload-time = "2026-01-19T22:54:59.066Z" } +sdist = { url = "https://files.pythonhosted.org/packages/a6/f8/95322db7a865e4df6bad108b1c99aa7fbe211aac3f298f3ad696c2744a39/mflux-0.15.4.tar.gz", hash = "sha256:138e1aedae86e13eafeb8faec017945fcdcca42c3234daabcd81a83c9a202ace", size = 741228, upload-time = "2026-01-20T15:39:26.807Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/cf/9f/a673ee12877a0943a4059c51b5beb6cf909c92f25384365cf8beeb475159/mflux-0.15.3-py3-none-any.whl", hash = "sha256:631cfcc038f27e9bd0ff76c25c2bc7373562b8f64cf0ce961fc268a246fa699e", size = 987270, upload-time = "2026-01-19T22:54:57.155Z" }, + { url = "https://files.pythonhosted.org/packages/8e/be/81cf4ce2d1933b9b210c028a05ac95e958008c0d43e377a5f2757b7f2d4d/mflux-0.15.4-py3-none-any.whl", hash = "sha256:f04d9b1d7c5cd67880f483ab29fb2097648a25459eef9c5ee6480fad46de5e82", size = 987644, upload-time = "2026-01-20T15:39:24.817Z" }, ] [[package]] @@ -2227,6 +2217,15 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/44/6f/7120676b6d73228c96e17f1f794d8ab046fc910d781c8d151120c3f1569e/toml-0.10.2-py2.py3-none-any.whl", hash = "sha256:806143ae5bfb6a3c6e736a764057db0e6a0e05e338b5630894a5f779cabb4f9b", size = 16588, upload-time = "2020-11-01T01:40:20.672Z" }, ] +[[package]] +name = "tomlkit" +version = "0.14.0" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/c3/af/14b24e41977adb296d6bd1fb59402cf7d60ce364f90c890bd2ec65c43b5a/tomlkit-0.14.0.tar.gz", hash = "sha256:cf00efca415dbd57575befb1f6634c4f42d2d87dbba376128adb42c121b87064", size = 187167, upload-time = "2026-01-13T01:14:53.304Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/b5/11/87d6d29fb5d237229d67973a6c9e06e048f01cf4994dee194ab0ea841814/tomlkit-0.14.0-py3-none-any.whl", hash = "sha256:592064ed85b40fa213469f81ac584f67a4f2992509a7c3ea2d632208623a3680", size = 39310, upload-time = "2026-01-13T01:14:51.965Z" }, +] + [[package]] name = "torch" version = "2.9.1"