Further refactor
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user