diff --git a/src/exo/worker/engines/mflux/distributed_model.py b/src/exo/worker/engines/mflux/distributed_model.py index 5c90f50f..c46d3241 100644 --- a/src/exo/worker/engines/mflux/distributed_model.py +++ b/src/exo/worker/engines/mflux/distributed_model.py @@ -20,7 +20,7 @@ from exo.shared.types.worker.shards import PipelineShardMetadata from exo.worker.download.download_utils import build_model_path from exo.worker.engines.mflux.config import get_config_for_model from exo.worker.engines.mflux.config.model_config import ImageModelConfig -from exo.worker.engines.mflux.pipefusion import create_model, get_adapter_for_model +from exo.worker.engines.mflux.pipefusion import get_adapter_for_model from exo.worker.engines.mflux.pipefusion.adapter import ModelAdapter from exo.worker.engines.mflux.pipefusion.distributed_denoising import ( DistributedDenoising, @@ -68,8 +68,8 @@ class DistributedImageModel: config = get_config_for_model(model_id) adapter = get_adapter_for_model(config) - # Create the model using the factory registry - model = create_model(config, model_id, local_path, quantize) + # Create the model using the adapter + model = adapter.create_model(model_id, local_path, quantize) if group is not None: # Apply pipeline parallelism by wrapping the transformer diff --git a/src/exo/worker/engines/mflux/pipefusion/__init__.py b/src/exo/worker/engines/mflux/pipefusion/__init__.py index 53e6cbdf..f8ddfabd 100644 --- a/src/exo/worker/engines/mflux/pipefusion/__init__.py +++ b/src/exo/worker/engines/mflux/pipefusion/__init__.py @@ -1,16 +1,14 @@ """ -Adapter and model factory registries. +Adapter registry for model-specific operations. -This module provides registry patterns for managing model adapters and -model factories, allowing new model families to be added without modifying -core code. +This module provides a registry pattern for managing model adapters, +allowing new model families to be added without modifying core code. + +Each adapter is responsible for both model creation and model-specific +distributed inference operations. """ -from pathlib import Path -from typing import Any, Callable - -from mflux.config.model_config import ModelConfig -from mflux.models.flux.variants.txt2img.flux import Flux1 +from typing import Callable from exo.worker.engines.mflux.config.model_config import ImageModelConfig from exo.worker.engines.mflux.pipefusion.adapter import ModelAdapter @@ -51,63 +49,3 @@ def register_adapter(model_family: str, factory: AdapterFactory) -> None: factory: A callable that takes an ImageModelConfig and returns a ModelAdapter """ _ADAPTER_REGISTRY[model_family] = factory - - -# ============================================================================= -# Model Factory Registry -# ============================================================================= - -# Type alias for model factory functions -# Takes (model_id, local_path, quantize) and returns a model instance -ModelFactory = Callable[[str, Path, int | None], Any] - - -def _create_flux_model(model_id: str, local_path: Path, quantize: int | None) -> Flux1: - """Create a Flux1 model instance.""" - return Flux1( - model_config=ModelConfig.from_name(model_name=model_id, base_model=None), - local_path=str(local_path), - quantize=quantize, - ) - - -# Registry maps model_family string to model factory -_MODEL_REGISTRY: dict[str, ModelFactory] = { - "flux": _create_flux_model, -} - - -def create_model( - config: ImageModelConfig, - model_id: str, - local_path: Path, - quantize: int | None = None, -) -> Any: - """Create a model instance for a model configuration. - - Args: - config: The model configuration - model_id: The model identifier (e.g., "black-forest-labs/FLUX.1-schnell") - local_path: Path to the local model weights - quantize: Optional quantization bit width - - Returns: - A model instance for the model family - - Raises: - ValueError: If no factory is registered for the model family - """ - factory = _MODEL_REGISTRY.get(config.model_family) - if factory is None: - raise ValueError(f"No model factory found for model family: {config.model_family}") - return factory(model_id, local_path, quantize) - - -def register_model_factory(model_family: str, factory: ModelFactory) -> None: - """Register a new model factory for a model family. - - Args: - model_family: The model family identifier (e.g., "flux", "fibo", "qwen") - factory: A callable that takes (model_id, local_path, quantize) and returns a model - """ - _MODEL_REGISTRY[model_family] = factory diff --git a/src/exo/worker/engines/mflux/pipefusion/adapter.py b/src/exo/worker/engines/mflux/pipefusion/adapter.py index 0841601f..10bc12a8 100644 --- a/src/exo/worker/engines/mflux/pipefusion/adapter.py +++ b/src/exo/worker/engines/mflux/pipefusion/adapter.py @@ -1,4 +1,5 @@ from enum import Enum +from pathlib import Path from typing import Any, Protocol import mlx.core as mx @@ -29,6 +30,24 @@ class ModelAdapter(Protocol): """Return the model configuration.""" ... + def create_model( + self, + model_id: str, + local_path: Path, + quantize: int | None = None, + ) -> Any: + """Create the underlying mflux model instance. + + Args: + model_id: The model identifier (e.g., "black-forest-labs/FLUX.1-schnell") + local_path: Path to the local model weights + quantize: Optional quantization bit width + + Returns: + The mflux model instance (e.g., Flux1, Fibo, Qwen) + """ + ... + def compute_embeddings( self, hidden_states: mx.array, diff --git a/src/exo/worker/engines/mflux/pipefusion/flux_adapter.py b/src/exo/worker/engines/mflux/pipefusion/flux_adapter.py index 90f394bc..d4cc5db3 100644 --- a/src/exo/worker/engines/mflux/pipefusion/flux_adapter.py +++ b/src/exo/worker/engines/mflux/pipefusion/flux_adapter.py @@ -1,6 +1,8 @@ +from pathlib import Path from typing import Any import mlx.core as mx +from mflux.config.model_config import ModelConfig from mflux.config.runtime_config import RuntimeConfig from mflux.models.flux.model.flux_transformer.common.attention_utils import ( AttentionUtils, @@ -9,6 +11,7 @@ from mflux.models.flux.model.flux_transformer.joint_transformer_block import ( JointTransformerBlock, ) from mflux.models.flux.model.flux_transformer.transformer import Transformer +from mflux.models.flux.variants.txt2img.flux import Flux1 from exo.worker.engines.mflux.config.model_config import ImageModelConfig from exo.worker.engines.mflux.pipefusion.adapter import BlockWrapperMode @@ -31,6 +34,19 @@ class FluxModelAdapter: def config(self) -> ImageModelConfig: return self._config + def create_model( + self, + model_id: str, + local_path: Path, + quantize: int | None = None, + ) -> Flux1: + """Create a Flux1 model instance.""" + return Flux1( + model_config=ModelConfig.from_name(model_name=model_id, base_model=None), + local_path=str(local_path), + quantize=quantize, + ) + def compute_embeddings( self, hidden_states: mx.array,