← back to Exo
Ciaran/image quantization (#1272)
ffba340e7010309b77e53854f984c681e3e0e81d · 2026-01-26 19:25:05 +0000 · ciaranbor
## Motivation
Enable users to select and use quantized variants (8-bit, 4-bit) of
image models
## Changes
Use exolabs HF org for image models
## Why It Works
Quantized versions have been uploaded to exolabs HF org
## Test Plan
Loaded and ran different quantized variants. Confirmed lower memory
usage and different outputs for the same seed. Verified chat completion
still works.
Files touched
M src/exo/shared/models/model_cards.py
Diff
commit ffba340e7010309b77e53854f984c681e3e0e81d
Author: ciaranbor <81697641+ciaranbor@users.noreply.github.com>
Date: Mon Jan 26 19:25:05 2026 +0000
Ciaran/image quantization (#1272)
## Motivation
Enable users to select and use quantized variants (8-bit, 4-bit) of
image models
## Changes
Use exolabs HF org for image models
## Why It Works
Quantized versions have been uploaded to exolabs HF org
## Test Plan
Loaded and ran different quantized variants. Confirmed lower memory
usage and different outputs for the same seed. Verified chat completion
still works.
---
src/exo/shared/models/model_cards.py | 118 ++++++++++++++++++++++++++++++-----
1 file changed, 102 insertions(+), 16 deletions(-)
diff --git a/src/exo/shared/models/model_cards.py b/src/exo/shared/models/model_cards.py
index 1d09293a..bcc5c73f 100644
--- a/src/exo/shared/models/model_cards.py
+++ b/src/exo/shared/models/model_cards.py
@@ -413,9 +413,9 @@ MODEL_CARDS: dict[str, ModelCard] = {
),
}
-_IMAGE_MODEL_CARDS: dict[str, ModelCard] = {
+_IMAGE_BASE_MODEL_CARDS: dict[str, ModelCard] = {
"flux1-schnell": ModelCard(
- model_id=ModelId("black-forest-labs/FLUX.1-schnell"),
+ model_id=ModelId("exolabs/FLUX.1-schnell"),
storage_size=Memory.from_bytes(23782357120 + 9524621312),
n_layers=57,
hidden_size=1,
@@ -428,7 +428,7 @@ _IMAGE_MODEL_CARDS: dict[str, ModelCard] = {
storage_size=Memory.from_kb(0),
n_layers=12,
can_shard=False,
- safetensors_index_filename=None, # Single file
+ safetensors_index_filename=None,
),
ComponentInfo(
component_name="text_encoder_2",
@@ -442,7 +442,7 @@ _IMAGE_MODEL_CARDS: dict[str, ModelCard] = {
component_name="transformer",
component_path="transformer/",
storage_size=Memory.from_bytes(23782357120),
- n_layers=57, # 19 transformer_blocks + 38 single_transformer_blocks
+ n_layers=57,
can_shard=True,
safetensors_index_filename="diffusion_pytorch_model.safetensors.index.json",
),
@@ -457,7 +457,7 @@ _IMAGE_MODEL_CARDS: dict[str, ModelCard] = {
],
),
"flux1-dev": ModelCard(
- model_id=ModelId("black-forest-labs/FLUX.1-dev"),
+ model_id=ModelId("exolabs/FLUX.1-dev"),
storage_size=Memory.from_bytes(23782357120 + 9524621312),
n_layers=57,
hidden_size=1,
@@ -470,7 +470,7 @@ _IMAGE_MODEL_CARDS: dict[str, ModelCard] = {
storage_size=Memory.from_kb(0),
n_layers=12,
can_shard=False,
- safetensors_index_filename=None, # Single file
+ safetensors_index_filename=None,
),
ComponentInfo(
component_name="text_encoder_2",
@@ -484,7 +484,7 @@ _IMAGE_MODEL_CARDS: dict[str, ModelCard] = {
component_name="transformer",
component_path="transformer/",
storage_size=Memory.from_bytes(23802816640),
- n_layers=57, # 19 transformer_blocks + 38 single_transformer_blocks
+ n_layers=57,
can_shard=True,
safetensors_index_filename="diffusion_pytorch_model.safetensors.index.json",
),
@@ -499,7 +499,7 @@ _IMAGE_MODEL_CARDS: dict[str, ModelCard] = {
],
),
"flux1-krea-dev": ModelCard(
- model_id=ModelId("black-forest-labs/FLUX.1-Krea-dev"),
+ model_id=ModelId("exolabs/FLUX.1-Krea-dev"),
storage_size=Memory.from_bytes(23802816640 + 9524621312), # Same as dev
n_layers=57,
hidden_size=1,
@@ -541,9 +541,9 @@ _IMAGE_MODEL_CARDS: dict[str, ModelCard] = {
],
),
"qwen-image": ModelCard(
- model_id=ModelId("Qwen/Qwen-Image"),
+ model_id=ModelId("exolabs/Qwen-Image"),
storage_size=Memory.from_bytes(16584333312 + 40860802176),
- n_layers=60, # Qwen has 60 transformer blocks (all joint-style)
+ n_layers=60,
hidden_size=1,
supports_tensor=False,
tasks=[ModelTask.TextToImage],
@@ -551,10 +551,10 @@ _IMAGE_MODEL_CARDS: dict[str, ModelCard] = {
ComponentInfo(
component_name="text_encoder",
component_path="text_encoder/",
- storage_size=Memory.from_kb(16584333312),
+ storage_size=Memory.from_bytes(16584333312),
n_layers=12,
can_shard=False,
- safetensors_index_filename=None, # Single file
+ safetensors_index_filename=None,
),
ComponentInfo(
component_name="transformer",
@@ -575,9 +575,9 @@ _IMAGE_MODEL_CARDS: dict[str, ModelCard] = {
],
),
"qwen-image-edit-2509": ModelCard(
- model_id=ModelId("Qwen/Qwen-Image-Edit-2509"),
+ model_id=ModelId("exolabs/Qwen-Image-Edit-2509"),
storage_size=Memory.from_bytes(16584333312 + 40860802176),
- n_layers=60, # Qwen has 60 transformer blocks (all joint-style)
+ n_layers=60,
hidden_size=1,
supports_tensor=False,
tasks=[ModelTask.ImageToImage],
@@ -585,10 +585,10 @@ _IMAGE_MODEL_CARDS: dict[str, ModelCard] = {
ComponentInfo(
component_name="text_encoder",
component_path="text_encoder/",
- storage_size=Memory.from_kb(16584333312),
+ storage_size=Memory.from_bytes(16584333312),
n_layers=12,
can_shard=False,
- safetensors_index_filename=None, # Single file
+ safetensors_index_filename=None,
),
ComponentInfo(
component_name="transformer",
@@ -610,6 +610,92 @@ _IMAGE_MODEL_CARDS: dict[str, ModelCard] = {
),
}
+
+def _generate_image_model_quant_variants(
+ base_name: str,
+ base_card: ModelCard,
+) -> dict[str, ModelCard]:
+ """Create quantized variants of an image model card.
+
+ Only the transformer component is quantized; text encoders stay at bf16.
+ Sizes are calculated exactly from the base card's component sizes.
+ """
+ if base_card.components is None:
+ raise ValueError(f"Image model {base_name} must have components defined")
+
+ # quantizations = [8, 6, 5, 4, 3]
+ quantizations = [8, 4]
+
+ num_transformer_bytes = next(
+ c.storage_size.in_bytes
+ for c in base_card.components
+ if c.component_name == "transformer"
+ )
+
+ transformer_bytes = Memory.from_bytes(num_transformer_bytes)
+
+ remaining_bytes = Memory.from_bytes(
+ sum(
+ c.storage_size.in_bytes
+ for c in base_card.components
+ if c.component_name != "transformer"
+ )
+ )
+
+ def with_transformer_size(new_size: Memory) -> list[ComponentInfo]:
+ assert base_card.components is not None
+ return [
+ ComponentInfo(
+ component_name=c.component_name,
+ component_path=c.component_path,
+ storage_size=new_size
+ if c.component_name == "transformer"
+ else c.storage_size,
+ n_layers=c.n_layers,
+ can_shard=c.can_shard,
+ safetensors_index_filename=c.safetensors_index_filename,
+ )
+ for c in base_card.components
+ ]
+
+ variants = {
+ base_name: ModelCard(
+ model_id=base_card.model_id,
+ storage_size=transformer_bytes + remaining_bytes,
+ n_layers=base_card.n_layers,
+ hidden_size=base_card.hidden_size,
+ supports_tensor=base_card.supports_tensor,
+ tasks=base_card.tasks,
+ components=with_transformer_size(transformer_bytes),
+ )
+ }
+
+ for quant in quantizations:
+ quant_transformer_bytes = Memory.from_bytes(
+ (num_transformer_bytes * quant) // 16
+ )
+ total_bytes = remaining_bytes + quant_transformer_bytes
+
+ model_id = ModelId(base_card.model_id + f"-{quant}bit")
+
+ variants[f"{base_name}-{quant}bit"] = ModelCard(
+ model_id=model_id,
+ storage_size=total_bytes,
+ n_layers=base_card.n_layers,
+ hidden_size=base_card.hidden_size,
+ supports_tensor=base_card.supports_tensor,
+ tasks=base_card.tasks,
+ components=with_transformer_size(quant_transformer_bytes),
+ )
+
+ return variants
+
+
+_image_model_cards: dict[str, ModelCard] = {}
+for _base_name, _base_card in _IMAGE_BASE_MODEL_CARDS.items():
+ _image_model_cards |= _generate_image_model_quant_variants(_base_name, _base_card)
+_IMAGE_MODEL_CARDS = _image_model_cards
+
if EXO_ENABLE_IMAGE_MODELS:
MODEL_CARDS.update(_IMAGE_MODEL_CARDS)
← 9968abe8 Leo/fix basic model shard (#1291)
·
back to Exo
·
Only ignore message if actually empty (#1292) 59e991ce →