From d7be6a09b08e346439c2ba2499eb7ba6ef67f2ce Mon Sep 17 00:00:00 2001 From: ciaranbor Date: Wed, 24 Dec 2025 23:45:46 +0000 Subject: [PATCH] Add BaseModelAdaptor --- .../worker/engines/image/distributed_model.py | 100 ++++---------- src/exo/worker/engines/image/models/base.py | 124 ++++++++++++++++++ .../engines/image/models/flux/adapter.py | 56 +++++++- .../worker/engines/image/pipeline/adapter.py | 31 ++++- 4 files changed, 231 insertions(+), 80 deletions(-) create mode 100644 src/exo/worker/engines/image/models/base.py diff --git a/src/exo/worker/engines/image/distributed_model.py b/src/exo/worker/engines/image/distributed_model.py index 9652baec..772e41a5 100644 --- a/src/exo/worker/engines/image/distributed_model.py +++ b/src/exo/worker/engines/image/distributed_model.py @@ -3,28 +3,24 @@ from typing import TYPE_CHECKING, Any, Literal, Optional import mlx.core as mx from mflux.config.config import Config -from mflux.config.runtime_config import RuntimeConfig -from mflux.models.common.latent_creator.latent_creator import Img2Img, LatentCreator -from mflux.models.flux.latent_creator.flux_latent_creator import FluxLatentCreator -from mflux.models.flux.model.flux_text_encoder.prompt_encoder import PromptEncoder -from mflux.models.flux.variants.txt2img.flux import Flux1 -from mflux.utils.array_util import ArrayUtil -from mflux.utils.image_util import ImageUtil from PIL import Image 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.image.config import ImageModelConfig -from exo.worker.engines.image.models import create_adapter_for_model, get_config_for_model -from exo.worker.engines.image.pipeline import DiffusionRunner, ModelAdapter +from exo.worker.engines.image.models import ( + create_adapter_for_model, + get_config_for_model, +) +from exo.worker.engines.image.models.base import BaseModelAdapter +from exo.worker.engines.image.pipeline import DiffusionRunner from exo.worker.engines.mlx.utils_mlx import mlx_distributed_init, mx_barrier from exo.worker.runner.bootstrap import logger class DistributedImageModel: __slots__ = ( - "_model", "_config", "_adapter", "_group", @@ -32,9 +28,8 @@ class DistributedImageModel: "_runner", ) - _model: Flux1 # Will be generalized to support other model types _config: ImageModelConfig - _adapter: ModelAdapter + _adapter: BaseModelAdapter _group: Optional[mx.distributed.Group] _shard_metadata: PipelineShardMetadata _runner: DiffusionRunner @@ -51,9 +46,6 @@ class DistributedImageModel: config = get_config_for_model(model_id) adapter = create_adapter_for_model(config, model_id, local_path, quantize) - # Get model from adapter - model = adapter.model - # Create diffusion runner (handles both single-node and distributed modes) num_sync_steps = config.get_num_sync_steps("medium") if group else 0 runner = DiffusionRunner( @@ -67,10 +59,10 @@ class DistributedImageModel: if group is not None: logger.info("Initialized distributed diffusion runner") - mx.eval(model.parameters()) + mx.eval(adapter.model.parameters()) # TODO: Do we need this? - mx.eval(model) + mx.eval(adapter.model) # Synchronize processes before generation to avoid timeout mx_barrier(group) @@ -78,7 +70,6 @@ class DistributedImageModel: else: logger.info("Single-node initialization") - object.__setattr__(self, "_model", model) object.__setattr__(self, "_config", config) object.__setattr__(self, "_adapter", adapter) object.__setattr__(self, "_group", group) @@ -114,15 +105,16 @@ class DistributedImageModel: ) @property - def model(self) -> Flux1: - return self._model + def model(self) -> Any: + """Return the underlying mflux model via the adapter.""" + return self._adapter.model @property def config(self) -> ImageModelConfig: return self._config @property - def adapter(self) -> ModelAdapter: + def adapter(self) -> BaseModelAdapter: return self._adapter @property @@ -157,17 +149,16 @@ class DistributedImageModel: def runner(self) -> DiffusionRunner: return self._runner - # Delegate attribute access to the underlying model. + # Delegate attribute access to the underlying model via the adapter. # 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) + return getattr(self._adapter.model, name) def __setattr__(self, name: str, value: Any) -> None: if name in ( - "_model", "_config", "_adapter", "_group", @@ -176,7 +167,7 @@ class DistributedImageModel: ): object.__setattr__(self, name, value) else: - setattr(self._model, name, value) + setattr(self._adapter.model, name, value) def generate( self, @@ -198,61 +189,16 @@ class DistributedImageModel: return image.image def _generate_image(self, settings: Config, prompt: str, seed: int) -> Any: - model = self._model + """Generate image by delegating to the adapter. - # Create runtime config - runtime_config = RuntimeConfig(settings, model.model_config) - - # Create initial latents (all nodes create the same latents with same seed) - latents = LatentCreator.create_for_txt2img_or_img2img( - seed=seed, - height=runtime_config.height, - width=runtime_config.width, - img2img=Img2Img( - vae=model.vae, - latent_creator=FluxLatentCreator, - image_path=runtime_config.image_path, - sigmas=runtime_config.scheduler.sigmas, - init_time_step=runtime_config.init_time_step, - ), - ) - - # Encode the prompt (all nodes encode to get consistent embeddings) - prompt_embeds, pooled_prompt_embeds = PromptEncoder.encode_prompt( + The adapter handles all model-specific logic (latent creation, + prompt encoding, denoising loop, decoding) via the template method pattern. + """ + return self._adapter.generate_image( + settings=settings, prompt=prompt, - prompt_cache=model.prompt_cache, - t5_tokenizer=model.t5_tokenizer, - clip_tokenizer=model.clip_tokenizer, - t5_text_encoder=model.t5_text_encoder, - clip_text_encoder=model.clip_text_encoder, - ) - - # Run the diffusion loop (runner handles callbacks internally) - latents = self._runner.run( - latents=latents, - prompt_embeds=prompt_embeds, - pooled_prompt_embeds=pooled_prompt_embeds, - runtime_config=runtime_config, seed=seed, - prompt=prompt, - ) - - # Decode latents to image (all nodes decode for now) - latents = ArrayUtil.unpack_latents( - latents=latents, height=runtime_config.height, width=runtime_config.width - ) - decoded = model.vae.decode(latents) - return ImageUtil.to_image( - decoded_latents=decoded, - config=runtime_config, - seed=seed, - prompt=prompt, - quantization=model.bits, - lora_paths=model.lora_paths, - lora_scales=model.lora_scales, - image_path=runtime_config.image_path, - image_strength=runtime_config.image_strength, - generation_time=0, # TODO: Track time in runner if needed + runner=self._runner if self.is_distributed else None, ) diff --git a/src/exo/worker/engines/image/models/base.py b/src/exo/worker/engines/image/models/base.py new file mode 100644 index 00000000..1ab481f2 --- /dev/null +++ b/src/exo/worker/engines/image/models/base.py @@ -0,0 +1,124 @@ +from abc import ABC, abstractmethod +from typing import TYPE_CHECKING, Any + +import mlx.core as mx +from mflux.config.config import Config +from mflux.config.runtime_config import RuntimeConfig +from mflux.models.common.latent_creator.latent_creator import Img2Img, LatentCreator +from mflux.utils.array_util import ArrayUtil +from mflux.utils.image_util import ImageUtil + +if TYPE_CHECKING: + from exo.worker.engines.image.pipeline.runner import DiffusionRunner + + +class BaseModelAdapter(ABC): + """Base class for model adapters with shared generation logic. + + Uses the template method pattern to share common generation flow + while allowing subclasses to implement model-specific steps. + """ + + def generate_image( + self, + settings: Config, + prompt: str, + seed: int, + runner: "DiffusionRunner | None" = None, + ) -> Any: + """Generate an image using the template method pattern. + + Args: + settings: Generation config (steps, height, width) + prompt: Text prompt + seed: Random seed + runner: Optional DiffusionRunner for distributed mode + + Returns: + GeneratedImage result + """ + # 1. Create runtime config (shared) + runtime_config = RuntimeConfig(settings, self.model.model_config) + + # 2. Create initial latents (uses model-specific latent creator) + latents = self._create_latents(seed, runtime_config) + + # 3. Encode prompt (model-specific) + prompt_data = self._encode_prompt(prompt) + + # 4. Run denoising loop (model-specific) + latents = self._run_denoising(latents, prompt_data, runtime_config, runner) + + # 5. Decode and return (shared) + return self._decode_latents(latents, runtime_config, seed, prompt) + + def _create_latents(self, seed: int, runtime_config: RuntimeConfig) -> mx.array: + """Create initial latents. Uses model-specific latent creator.""" + return LatentCreator.create_for_txt2img_or_img2img( + seed=seed, + height=runtime_config.height, + width=runtime_config.width, + img2img=Img2Img( + vae=self.model.vae, + latent_creator=self._get_latent_creator(), + sigmas=runtime_config.scheduler.sigmas, + init_time_step=runtime_config.init_time_step, + image_path=runtime_config.image_path, + ), + ) + + def _decode_latents( + self, + latents: mx.array, + runtime_config: RuntimeConfig, + seed: int, + prompt: str, + ) -> Any: + """Decode latents to image. Shared implementation.""" + latents = ArrayUtil.unpack_latents( + latents=latents, + height=runtime_config.height, + width=runtime_config.width, + ) + decoded = self.model.vae.decode(latents) + return ImageUtil.to_image( + decoded_latents=decoded, + config=runtime_config, + seed=seed, + prompt=prompt, + quantization=self.model.bits, + lora_paths=self.model.lora_paths, + lora_scales=self.model.lora_scales, + image_path=runtime_config.image_path, + image_strength=runtime_config.image_strength, + generation_time=0, + ) + + # Abstract methods - subclasses must implement + + @property + @abstractmethod + def model(self) -> Any: + """Return the underlying mflux model.""" + ... + + @abstractmethod + def _get_latent_creator(self) -> type: + """Return the latent creator class for this model.""" + ... + + @abstractmethod + def _encode_prompt(self, prompt: str) -> Any: + """Encode the prompt. Returns model-specific prompt data.""" + ... + + @abstractmethod + def _run_denoising( + self, + latents: mx.array, + prompt_data: Any, + runtime_config: RuntimeConfig, + runner: "DiffusionRunner | None", + ) -> mx.array: + """Run the denoising loop. Model-specific implementation.""" + ... diff --git a/src/exo/worker/engines/image/models/flux/adapter.py b/src/exo/worker/engines/image/models/flux/adapter.py index cf82c453..498c3d40 100644 --- a/src/exo/worker/engines/image/models/flux/adapter.py +++ b/src/exo/worker/engines/image/models/flux/adapter.py @@ -1,9 +1,11 @@ from pathlib import Path -from typing import Any, cast +from typing import TYPE_CHECKING, Any, cast import mlx.core as mx from mflux.config.model_config import ModelConfig from mflux.config.runtime_config import RuntimeConfig +from mflux.models.flux.latent_creator.flux_latent_creator import FluxLatentCreator +from mflux.models.flux.model.flux_text_encoder.prompt_encoder import PromptEncoder from mflux.models.flux.model.flux_transformer.common.attention_utils import ( AttentionUtils, ) @@ -14,6 +16,7 @@ from mflux.models.flux.model.flux_transformer.transformer import Transformer from mflux.models.flux.variants.txt2img.flux import Flux1 from exo.worker.engines.image.config import BlockType, ImageModelConfig +from exo.worker.engines.image.models.base import BaseModelAdapter from exo.worker.engines.image.pipeline.adapter import ( BlockWrapperMode, JointBlockInterface, @@ -21,8 +24,11 @@ from exo.worker.engines.image.pipeline.adapter import ( ) from exo.worker.engines.image.pipeline.kv_cache import ImagePatchKVCache +if TYPE_CHECKING: + from exo.worker.engines.image.pipeline.runner import DiffusionRunner -class FluxModelAdapter: + +class FluxModelAdapter(BaseModelAdapter): def __init__( self, config: ImageModelConfig, @@ -55,6 +61,52 @@ class FluxModelAdapter: def hidden_dim(self) -> int: return self._transformer.x_embedder.weight.shape[0] + # ------------------------------------------------------------------------- + # BaseModelAdapter abstract method implementations + # ------------------------------------------------------------------------- + + def _get_latent_creator(self) -> type: + return FluxLatentCreator + + def _encode_prompt(self, prompt: str) -> tuple[mx.array, mx.array]: + return PromptEncoder.encode_prompt( + prompt=prompt, + prompt_cache=self._model.prompt_cache, + t5_tokenizer=self._model.t5_tokenizer, + clip_tokenizer=self._model.clip_tokenizer, + t5_text_encoder=self._model.t5_text_encoder, + clip_text_encoder=self._model.clip_text_encoder, + ) + + def _run_denoising( + self, + latents: mx.array, + prompt_data: tuple[mx.array, mx.array], + runtime_config: RuntimeConfig, + runner: "DiffusionRunner | None", + ) -> mx.array: + prompt_embeds, pooled_prompt_embeds = prompt_data + if runner: + # Distributed mode - use DiffusionRunner + return runner.run( + latents=latents, + prompt_embeds=prompt_embeds, + pooled_prompt_embeds=pooled_prompt_embeds, + runtime_config=runtime_config, + seed=0, # Not used by runner + prompt="", # Not used by runner + ) + else: + # Single-node mode - use DiffusionRunner with no distribution + # This path shouldn't be hit in practice since we always have a runner + raise NotImplementedError( + "Single-node FLUX generation requires a DiffusionRunner" + ) + + # ------------------------------------------------------------------------- + # ModelAdapter protocol implementations (for distributed inference) + # ------------------------------------------------------------------------- + def compute_embeddings( self, hidden_states: mx.array, diff --git a/src/exo/worker/engines/image/pipeline/adapter.py b/src/exo/worker/engines/image/pipeline/adapter.py index e0386520..2c50093d 100644 --- a/src/exo/worker/engines/image/pipeline/adapter.py +++ b/src/exo/worker/engines/image/pipeline/adapter.py @@ -1,5 +1,5 @@ from enum import Enum -from typing import Any, Protocol +from typing import TYPE_CHECKING, Any, Protocol import mlx.core as mx from mflux.config.runtime_config import RuntimeConfig @@ -7,6 +7,11 @@ from mflux.config.runtime_config import RuntimeConfig from exo.worker.engines.image.config import BlockType, ImageModelConfig from exo.worker.engines.image.pipeline.kv_cache import ImagePatchKVCache +if TYPE_CHECKING: + from mflux.config.config import Config + + from exo.worker.engines.image.pipeline.runner import DiffusionRunner + class AttentionInterface(Protocol): num_heads: int @@ -285,3 +290,27 @@ class ModelAdapter(Protocol): Merged hidden states (default: concatenate [text, image]) """ ... + + def generate_image( + self, + settings: "Config", + prompt: str, + seed: int, + runner: "DiffusionRunner | None" = None, + ) -> Any: + """Generate an image using this model. + + This is the main entry point for image generation. Implementations + should handle the full generation flow: latent creation, prompt + encoding, denoising loop, and decoding. + + Args: + settings: Generation config (steps, height, width) + prompt: Text prompt + seed: Random seed + runner: Optional DiffusionRunner for distributed mode + + Returns: + GeneratedImage result + """ + ...