[object Object]

← back to Exo

Add FLUX.1-Krea-dev model (#1269)

23fd37fe4d684ec5d90f7a4b5b03b3a829c64fa0 · 2026-01-23 19:48:24 +0000 · ciaranbor

## Why It Works

Same implementation as FLUX.1-dev, just different weights

Files touched

Diff

commit 23fd37fe4d684ec5d90f7a4b5b03b3a829c64fa0
Author: ciaranbor <81697641+ciaranbor@users.noreply.github.com>
Date:   Fri Jan 23 19:48:24 2026 +0000

    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 +++++++++++++++++++++++++
 src/exo/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"

← d229df38 Fix placement filter to use subset matching instead of exact  ·  back to Exo  ·  Prevent conversation collision (#1266) 9967dfa7 →