Further refactor

This commit is contained in:
ciaranbor
2026-01-06 10:51:20 +00:00
parent fb4fae51fa
commit 88996eddcb
3 changed files with 38 additions and 36 deletions
@@ -1,4 +1,5 @@
from typing import Literal
import mlx.core as mx
from mflux.callbacks.callbacks import Callbacks
from mflux.config.config import Config
@@ -13,8 +14,6 @@ from mflux.utils.image_util import ImageUtil
from PIL import Image
from tqdm import tqdm
from exo.shared.types.api import ImageGenerationTaskParams
def _generate_image(model: Flux1, settings: Config, prompt: str, seed: int):
# 0. Create a new runtime config based on the model type and input parameters
+29 -17
View File
@@ -13,6 +13,7 @@ from exo.shared.types.worker.shards import (
ShardMetadata,
TensorShardMetadata,
)
from exo.worker.engines.mlx.utils_mlx import mx_barrier
class _JointBlock(Protocol):
@@ -278,23 +279,9 @@ class FluxSingleSyncBlock(CustomMlxSingleBlock):
return hidden_states
def shard_flux_transformer(
model: Flux1,
group: mx.distributed.Group,
shard_metadata: ShardMetadata,
) -> Flux1:
if isinstance(shard_metadata, TensorShardMetadata):
raise NotImplementedError(
"Tensor parallelism is not yet supported for Flux models. "
"Use pipeline parallelism instead."
)
if not isinstance(shard_metadata, PipelineShardMetadata):
raise ValueError(
f"Unsupported shard metadata type: {type(shard_metadata)}. "
"Expected PipelineShardMetadata."
)
def pipeline_transformer(
model: Flux1, group: mx.distributed.Group, shard_metadata: ShardMetadata
):
transformer: Transformer = model.transformer
# Total = joint blocks + single blocks
@@ -385,3 +372,28 @@ def shard_flux_transformer(
transformer.single_transformer_blocks = assigned_single_blocks
return model
def shard_flux_transformer(
model: Flux1,
group: mx.distributed.Group,
shard_metadata: ShardMetadata,
) -> Flux1:
match shard_metadata:
case TensorShardMetadata():
raise NotImplementedError(
"Tensor parallelism is not yet supported for Flux models. "
"Use pipeline parallelism instead."
)
case PipelineShardMetadata():
model = pipeline_transformer(model, group, shard_metadata)
mx.eval(model.parameters())
# TODO: Do we need this?
mx.eval(model)
# Synchronize processes before generation to avoid timeout
mx_barrier(group)
return model
+8 -17
View File
@@ -12,29 +12,20 @@ def initialize_mflux(bound_instance: BoundInstance) -> Flux1:
model_id = bound_instance.bound_shard.model_meta.model_id
model_path = build_model_path(model_id)
# TODO: generalise
model = Flux1(
model_config=ModelConfig.from_name(model_name=model_id, base_model=None),
local_path=str(model_path),
# quantize=8,
)
is_distributed = len(bound_instance.instance.shard_assignments.node_to_runner) > 1
if not is_distributed:
# Single-node: load full model normally
logger.info(f"Single device used for {bound_instance.instance}")
model = Flux1(
model_config=ModelConfig.from_name(model_name=model_id, base_model=None),
local_path=str(model_path),
# quantize=8,
)
else:
if is_distributed:
# Multi-node: initialize distributed and shard transformer
logger.info("Starting distributed init for Flux")
group = mlx_distributed_init(bound_instance)
logger.info("Loading Flux model for distributed inference")
model = Flux1(
model_config=ModelConfig.from_name(model_name=model_id, base_model=None),
local_path=str(model_path),
# quantize=8,
)
logger.info("Applying pipeline parallelism to Flux transformer")
model = shard_flux_transformer(
model=model,
group=group,