Combine model factory and adaptor
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user