Add DistributedFlux1 class

This commit is contained in:
ciaranbor
2026-01-06 10:51:20 +00:00
parent 6e0c611f37
commit be796e55ac
6 changed files with 127 additions and 12 deletions
@@ -0,0 +1,85 @@
from typing import TYPE_CHECKING, Any
import mlx.core as mx
from mflux.models.flux.variants.txt2img.flux import Flux1
from exo.shared.types.worker.shards import PipelineShardMetadata
class DistributedFlux1:
"""
Wrapper for Flux1 that attaches distributed group and shard metadata.
This wrapper enables the generation runtime to access distributed context
(group, rank, world_size, shard boundaries).
"""
__slots__ = ("_model", "_group", "_shard_metadata")
_model: Flux1
_group: mx.distributed.Group
_shard_metadata: PipelineShardMetadata
def __init__(
self,
model: Flux1,
group: mx.distributed.Group,
shard_metadata: PipelineShardMetadata,
) -> None:
object.__setattr__(self, "_model", model)
object.__setattr__(self, "_group", group)
object.__setattr__(self, "_shard_metadata", shard_metadata)
@property
def model(self) -> Flux1:
"""The underlying Flux1 model."""
return self._model
@property
def group(self) -> mx.distributed.Group:
"""The MLX distributed group for this model."""
return self._group
@property
def shard_metadata(self) -> PipelineShardMetadata:
"""Shard metadata containing layer assignments and device info."""
return self._shard_metadata
@property
def rank(self) -> int:
"""This device's rank in the distributed group."""
return self._shard_metadata.device_rank
@property
def world_size(self) -> int:
"""Total number of devices in the distributed group."""
return self._shard_metadata.world_size
@property
def is_first_stage(self) -> bool:
"""True if this device is the first stage in the pipeline."""
return self._shard_metadata.device_rank == 0
@property
def is_last_stage(self) -> bool:
"""True if this device is the last stage in the pipeline."""
return self._shard_metadata.device_rank == self._shard_metadata.world_size - 1
@property
def is_distributed(self) -> bool:
"""True if running in distributed mode (world_size > 1)."""
return self._shard_metadata.world_size > 1
# Delegate attribute access to the underlying model.
# Guarded with TYPE_CHECKING to prevent type checker complaints
# while still providing full delegation at runtime.
if not TYPE_CHECKING:
def __getattr__(self, name: str) -> Any:
return getattr(self._model, name)
def __setattr__(self, name: str, value: Any) -> None:
if name in ("_model", "_group", "_shard_metadata"):
object.__setattr__(self, name, value)
else:
setattr(self._model, name, value)
@@ -7,6 +7,7 @@ from PIL import Image
from exo.shared.types.api import ImageGenerationTaskParams
from exo.shared.types.worker.runner_response import ImageGenerationResponse
from exo.worker.engines.mflux.distributed_flux import DistributedFlux1
from exo.worker.engines.mflux.generator.flux1 import generate_image
image_generation_stream = mx.new_stream(mx.default_device())
@@ -31,14 +32,21 @@ def parse_size(size_str: str | None) -> tuple[int, int]:
return (1024, 1024)
def warmup_mflux(model: Flux1) -> Image.Image:
def warmup_mflux(model: Flux1 | DistributedFlux1) -> Image.Image:
# Extract underlying model if wrapped
underlying_model = model.model if isinstance(model, DistributedFlux1) else model
return generate_image(
model=model, prompt="Warmup", height=256, width=256, quality="low", seed=2
model=underlying_model,
prompt="Warmup",
height=256,
width=256,
quality="low",
seed=2,
)
def mflux_generate(
model: Flux1,
model: Flux1 | DistributedFlux1,
task: ImageGenerationTaskParams,
) -> Generator[ImageGenerationResponse]:
# Parse parameters
@@ -47,8 +55,12 @@ def mflux_generate(
seed = 2 # TODO: not in OAI API?
# Extract underlying model if wrapped
# TODO: In future, use model.group for async pipeline when distributed
underlying_model = model.model if isinstance(model, DistributedFlux1) else model
image = generate_image(
model=model,
model=underlying_model,
prompt=task.prompt,
height=height,
width=width,
+1 -1
View File
@@ -5,7 +5,7 @@ from exo.shared.types.worker.shards import (
PipelineShardMetadata,
ShardMetadata,
)
from exo.worker.engines.mflux.pipefusion import apply_pipefusion_transformer
from exo.worker.engines.mflux.pipefusion.pipefusion import apply_pipefusion_transformer
from exo.worker.engines.mlx.utils_mlx import mx_barrier
+15 -2
View File
@@ -3,12 +3,14 @@ from mflux.config.model_config import ModelConfig
from mflux.models.flux.variants.txt2img.flux import Flux1
from exo.shared.types.worker.instances import BoundInstance
from exo.shared.types.worker.shards import PipelineShardMetadata
from exo.worker.download.download_utils import build_model_path
from exo.worker.engines.mflux.distributed_flux import DistributedFlux1
from exo.worker.engines.mflux.shard_mflux import shard_flux_transformer
from exo.worker.engines.mlx.utils_mlx import mlx_distributed_init
def initialize_mflux(bound_instance: BoundInstance) -> Flux1:
def initialize_mflux(bound_instance: BoundInstance) -> Flux1 | DistributedFlux1:
model_id = bound_instance.bound_shard.model_meta.model_id
model_path = build_model_path(model_id)
@@ -26,11 +28,22 @@ def initialize_mflux(bound_instance: BoundInstance) -> Flux1:
logger.info("Starting distributed init for Flux")
group = mlx_distributed_init(bound_instance)
shard_metadata = bound_instance.bound_shard
if not isinstance(shard_metadata, PipelineShardMetadata):
raise ValueError("Expected PipelineShardMetadata for distributed Flux")
model = shard_flux_transformer(
model=model,
group=group,
shard_metadata=bound_instance.bound_shard,
shard_metadata=shard_metadata,
)
logger.info(f"Flux transformer sharded for rank {group.rank()}")
# Wrap with distributed context for runtime access
return DistributedFlux1(
model=model,
group=group,
shard_metadata=shard_metadata,
)
return model
+10 -5
View File
@@ -46,6 +46,7 @@ from exo.shared.types.worker.runners import (
RunnerWarmingUp,
)
from exo.utils.channels import ClosedResourceError, MpReceiver, MpSender
from exo.worker.engines.mflux.distributed_flux import DistributedFlux1
from exo.worker.engines.mflux.generator.generate import mflux_generate, warmup_mflux
from exo.worker.engines.mflux.utils_mflux import initialize_mflux
from exo.worker.engines.mlx.generator.generate import mlx_generate, warmup_inference
@@ -158,7 +159,9 @@ def main(
logger.info(f"warming up inference for instance: {instance}")
if ModelTask.TextGeneration in model_tasks:
# assert isinstance(model, Model) TODO: not actually Model
assert model and not isinstance(model, Flux1)
assert model and not isinstance(
model, (Flux1, DistributedFlux1)
)
assert tokenizer
assert sampler
@@ -176,7 +179,7 @@ def main(
ModelTask.TextToImage in model_tasks
or ModelTask.ImageToImage in model_tasks
):
assert isinstance(model, Flux1)
assert isinstance(model, (Flux1, DistributedFlux1))
image = warmup_mflux(model=model)
logger.info(f"warmed up by generating {image.size} image")
@@ -186,7 +189,9 @@ def main(
task_params=task_params, command_id=command_id
) if isinstance(current_status, RunnerReady):
# assert isinstance(model, Model) TODO: not actually Model
assert model and not isinstance(model, Flux1)
assert model and not isinstance(
model, (Flux1, DistributedFlux1)
)
assert tokenizer
assert sampler
logger.info(f"received chat request: {str(task)[:500]}")
@@ -235,7 +240,7 @@ def main(
case ImageGeneration(
task_params=task_params, command_id=command_id
) if isinstance(current_status, RunnerReady):
assert isinstance(model, Flux1)
assert isinstance(model, (Flux1, DistributedFlux1))
logger.info(
f"received image generation request: {str(task)[:500]}"
)
@@ -301,7 +306,7 @@ def main(
case ImageEdits(task_params=task_params, command_id=command_id) if (
isinstance(current_status, RunnerReady)
):
assert isinstance(model, Flux1)
assert isinstance(model, (Flux1, DistributedFlux1))
logger.info(
f"received image generation request: {str(task)[:500]}"
)