← back to Exo
Stable stable diffusion mlx
6b28ef03499e699974fdd9c58628ad5cd6b5c859 · 2024-11-19 23:13:22 -0500 · Pranav Veldurthi
Files touched
M exo/api/chatgpt_api.pyM exo/download/hf/hf_helpers.pyM exo/inference/inference_engine.pyA exo/inference/mlx/models/StableDiffusionPipeline.pyA exo/inference/mlx/models/sd_models/clip.pyA exo/inference/mlx/models/sd_models/tokenizer.pyA exo/inference/mlx/models/sd_models/unet.pyA exo/inference/mlx/models/sd_models/vae.pyM exo/inference/mlx/sharded_inference_engine.pyM exo/inference/mlx/sharded_utils.pyM exo/inference/mlx/stateful_model.pyM exo/main.pyM exo/models.pyM exo/networking/grpc/grpc_peer_handle.pyM exo/networking/grpc/grpc_server.pyM exo/networking/grpc/node_service.protoM exo/networking/grpc/node_service_pb2.pyM exo/networking/grpc/node_service_pb2_grpc.pyM exo/orchestration/standard_node.pyM exo/tinychat/index.htmlM exo/tinychat/index.js
Diff
commit 6b28ef03499e699974fdd9c58628ad5cd6b5c859
Author: Pranav Veldurthi <veldurthipranav@gmail.com>
Date: Tue Nov 19 23:13:22 2024 -0500
Stable stable diffusion mlx
---
exo/api/chatgpt_api.py | 79 ++-
exo/download/hf/hf_helpers.py | 4 +
exo/inference/inference_engine.py | 6 +-
.../mlx/models/StableDiffusionPipeline.py | 288 ++++++++++
exo/inference/mlx/models/sd_models/clip.py | 192 +++++++
exo/inference/mlx/models/sd_models/tokenizer.py | 131 +++++
exo/inference/mlx/models/sd_models/unet.py | 629 +++++++++++++++++++++
exo/inference/mlx/models/sd_models/vae.py | 390 +++++++++++++
exo/inference/mlx/sharded_inference_engine.py | 7 +-
exo/inference/mlx/sharded_utils.py | 73 ++-
exo/inference/mlx/stateful_model.py | 21 +-
exo/main.py | 2 +-
exo/models.py | 3 +
exo/networking/grpc/grpc_peer_handle.py | 42 +-
exo/networking/grpc/grpc_server.py | 29 +-
exo/networking/grpc/node_service.proto | 14 +-
exo/networking/grpc/node_service_pb2.py | 94 +--
exo/networking/grpc/node_service_pb2_grpc.py | 88 +--
exo/orchestration/standard_node.py | 84 +--
exo/tinychat/index.html | 10 +-
exo/tinychat/index.js | 130 +++--
21 files changed, 2122 insertions(+), 194 deletions(-)
diff --git a/exo/api/chatgpt_api.py b/exo/api/chatgpt_api.py
index ab074511..d981568b 100644
--- a/exo/api/chatgpt_api.py
+++ b/exo/api/chatgpt_api.py
@@ -15,7 +15,8 @@ from exo.inference.tokenizers import resolve_tokenizer
from exo.orchestration import Node
from exo.models import build_base_shard, model_cards, get_repo, pretty_name, get_supported_models
from typing import Callable, Optional
-
+from PIL import Image
+import numpy as np
class Message:
def __init__(self, role: str, content: Union[str, List[Dict[str, Union[str, Dict[str, str]]]]]):
@@ -169,13 +170,16 @@ class ChatGPTAPI:
cors.add(self.app.router.add_post("/v1/chat/token/encode", self.handle_post_chat_token_encode), {"*": cors_options})
cors.add(self.app.router.add_post("/chat/completions", self.handle_post_chat_completions), {"*": cors_options})
cors.add(self.app.router.add_post("/v1/chat/completions", self.handle_post_chat_completions), {"*": cors_options})
+ cors.add(self.app.router.add_post("/v1/image/generations", self.handle_post_image_generations), {"*": cors_options})
cors.add(self.app.router.add_get("/v1/download/progress", self.handle_get_download_progress), {"*": cors_options})
cors.add(self.app.router.add_get("/modelpool", self.handle_model_support), {"*": cors_options})
cors.add(self.app.router.add_get("/healthcheck", self.handle_healthcheck), {"*": cors_options})
+
self.static_dir = Path(__file__).parent.parent/"tinychat"
self.app.router.add_get("/", self.handle_root)
self.app.router.add_static("/", self.static_dir, name="static")
+ self.app.router.add_static('/images/', self.static_dir / 'images', name='static_images')
self.app.middlewares.append(self.timeout_middleware)
self.app.middlewares.append(self.log_request)
@@ -359,6 +363,79 @@ class ChatGPTAPI:
deregistered_callback = self.node.on_token.deregister(callback_id)
if DEBUG >= 2: print(f"Deregister {callback_id=} {deregistered_callback=}")
+
+ async def handle_post_image_generations(self, request):
+ data = await request.json()
+
+ if DEBUG >= 2: print(f"Handling chat completions request from {request.remote}: {data}")
+ stream = data.get("stream", False)
+ model = data.get("model", "")
+ prompt = data.get("prompt", "")
+ print(f"model: {model}, prompt: {prompt}, stream: {stream}")
+ shard = build_base_shard(model, self.inference_engine_classname)
+ print(f"shard: {shard}")
+ if not shard:
+ return web.json_response({"error": f"Unsupported model: {model} with inference engine {self.inference_engine_classname}"}, status=400)
+
+ request_id = str(uuid.uuid4())
+ callback_id = f"chatgpt-api-wait-response-{request_id}"
+ callback = self.node.on_token.register(callback_id)
+ try:
+ await asyncio.wait_for(asyncio.shield(asyncio.create_task(self.node.process_prompt(shard, prompt, request_id=request_id))), timeout=self.response_timeout)
+
+
+ response = web.StreamResponse(status=200, reason='OK', headers={'Content-Type': 'application/octet-stream',"Cache-Control": "no-cache",})
+ await response.prepare(request)
+
+ def get_progress_bar(current_step, total_steps, bar_length=50):
+ # Calculate the percentage of completion
+ percent = float(current_step) / total_steps
+ # Calculate the number of hashes to display
+ arrow = '-' * int(round(percent * bar_length) - 1) + '>'
+ spaces = ' ' * (bar_length - len(arrow))
+
+ # Create the progress bar string
+ progress_bar = f'Progress: [{arrow}{spaces}] {int(percent * 100)}% ({current_step}/{total_steps})'
+ return progress_bar
+
+ async def stream_image(_request_id: str, result, is_finished: bool):
+ if isinstance(result, list):
+ await response.write(json.dumps({'progress': get_progress_bar((result[0]), (result[1]))}).encode('utf-8') + b'\n')
+
+ elif isinstance(result, np.ndarray):
+ im = Image.fromarray(np.array(result))
+ # Save the image to a file
+ image_filename = f"{_request_id}.png"
+ image_path = self.static_dir / "images" / image_filename
+ im.save(image_path)
+ image_url = request.app.router['static_images'].url_for(filename=image_filename)
+ base_url = f"{request.scheme}://{request.host}"
+ # Construct the full URL correctly
+ full_image_url = base_url + str(image_url)
+
+ await response.write(json.dumps({'images': [{'url': str(full_image_url), 'content_type': 'image/png'}]}).encode('utf-8') + b'\n')
+
+ await response.write_eof()
+
+
+ stream_task = None
+ def on_result(_request_id: str, result, is_finished: bool):
+ nonlocal stream_task
+ stream_task = asyncio.create_task(stream_image(_request_id, result, is_finished))
+ return _request_id == request_id and is_finished
+
+ await callback.wait(on_result, timeout=self.response_timeout*10)
+
+ if stream_task:
+ # Wait for the stream task to complete before returning
+ await stream_task
+
+ return response
+
+ except Exception as e:
+ if DEBUG >= 2: traceback.print_exc()
+ return web.json_response({"detail": f"Error processing prompt (see logs with DEBUG>=2): {str(e)}"}, status=500)
+
async def run(self, host: str = "0.0.0.0", port: int = 52415):
runner = web.AppRunner(self.app)
await runner.setup()
diff --git a/exo/download/hf/hf_helpers.py b/exo/download/hf/hf_helpers.py
index 4729e5f6..a2696f17 100644
--- a/exo/download/hf/hf_helpers.py
+++ b/exo/download/hf/hf_helpers.py
@@ -276,6 +276,10 @@ async def download_repo_files(
await f.write(json.dumps(file_list))
if DEBUG >= 2: print(f"Cached file list at {cached_file_list_path}")
+ model_index_exists = any(file["path"] == "model_index.json" for file in file_list)
+ if model_index_exists:
+ allow_patterns = ["**/*.json", "**/*.txt", "**/*model.safetensors", "*.json"]
+
filtered_file_list = list(filter_repo_objects(file_list, allow_patterns=allow_patterns, ignore_patterns=ignore_patterns, key=lambda x: x["path"]))
total_files = len(filtered_file_list)
total_bytes = sum(file["size"] for file in filtered_file_list)
diff --git a/exo/inference/inference_engine.py b/exo/inference/inference_engine.py
index 0f093591..1e142b16 100644
--- a/exo/inference/inference_engine.py
+++ b/exo/inference/inference_engine.py
@@ -24,10 +24,10 @@ class InferenceEngine(ABC):
async def infer_tensor(self, request_id: str, shard: Shard, input_data: np.ndarray) -> np.ndarray:
pass
- async def infer_prompt(self, request_id: str, shard: Shard, prompt: str) -> np.ndarray:
+ async def infer_prompt(self, request_id: str, shard: Shard, prompt: str, inference_state: Optional[dict] = None) -> np.ndarray:
tokens = await self.encode(shard, prompt)
- output_data = await self.infer_tensor(request_id, shard, tokens)
- return output_data
+ output_data, inference_state = await self.infer_tensor(request_id, shard, tokens, inference_state)
+ return output_data, inference_state
inference_engine_classes = {
"mlx": "MLXDynamicShardInferenceEngine",
diff --git a/exo/inference/mlx/models/StableDiffusionPipeline.py b/exo/inference/mlx/models/StableDiffusionPipeline.py
new file mode 100644
index 00000000..443887a3
--- /dev/null
+++ b/exo/inference/mlx/models/StableDiffusionPipeline.py
@@ -0,0 +1,288 @@
+# Adapted from https://github.com/ml-explore/mlx-examples/blob/main/stable_diffusion/stable_diffusion/__init__.py
+
+import time
+from typing import Optional, Tuple
+import inspect
+
+import mlx.core as mx
+import mlx.nn as nn
+from pathlib import Path
+
+from tqdm import tqdm
+
+from .sd_models.vae import ModelArgs as VAEArgs
+from .sd_models.vae import Autoencoder
+from .sd_models.tokenizer import load_tokenizer
+from .sd_models.clip import CLIPTextModel
+from .sd_models.clip import ModelArgs as CLIPArgs
+from .sd_models.unet import UNetConfig, UNetModel
+
+from dataclasses import dataclass, field
+from exo.inference.shard import Shard
+
+@dataclass
+class DiffusionConfig:
+ beta_schedule: str = "scaled_linear"
+ beta_start: float = 0.00085
+ beta_end: float = 0.012
+ num_train_steps: int = 1000
+
+ @classmethod
+ def from_dict(cls, params):
+ return cls(**{k: v for k, v in params.items() if k in inspect.signature(cls).parameters})
+
+
+#Sampler
+def _linspace(a, b, num):
+ x = mx.arange(0, num) / (num - 1)
+ return (b - a) * x + a
+
+
+def _interp(y, x_new):
+ """Interpolate the function defined by (arange(0, len(y)), y) at positions x_new."""
+ x_low = x_new.astype(mx.int32)
+ x_high = mx.minimum(x_low + 1, len(y) - 1)
+
+ y_low = y[x_low]
+ y_high = y[x_high]
+ delta_x = x_new - x_low
+ y_new = y_low * (1 - delta_x) + delta_x * y_high
+
+ return y_new
+
+class SimpleEulerSampler:
+ """A simple Euler integrator that can be used to sample from our diffusion models.
+
+ The method ``step()`` performs one Euler step from x_t to x_t_prev.
+ """
+
+ def __init__(self, config: DiffusionConfig):
+ # Compute the noise schedule
+ if config.beta_schedule == "linear":
+ betas = _linspace(
+ config.beta_start, config.beta_end, config.num_train_steps
+ )
+ elif config.beta_schedule == "scaled_linear":
+ betas = _linspace(
+ config.beta_start**0.5, config.beta_end**0.5, config.num_train_steps
+ ).square()
+ else:
+ raise NotImplementedError(f"{config.beta_schedule} is not implemented.")
+
+ alphas = 1 - betas
+ alphas_cumprod = mx.cumprod(alphas)
+
+ self._sigmas = mx.concatenate(
+ [mx.zeros(1), ((1 - alphas_cumprod) / alphas_cumprod).sqrt()]
+ )
+
+ @property
+ def max_time(self):
+ return len(self._sigmas) - 1
+
+ def sample_prior(self, shape, dtype=mx.float32, key=None):
+ noise = mx.random.normal(shape, key=key)
+ return (
+ noise * self._sigmas[-1] * (self._sigmas[-1].square() + 1).rsqrt()
+ ).astype(dtype)
+
+ def add_noise(self, x, t, key=None):
+ noise = mx.random.normal(x.shape, key=key)
+ s = self.sigmas(t)
+ return (x + noise * s) * (s.square() + 1).rsqrt()
+
+ def sigmas(self, t):
+ return _interp(self._sigmas, t)
+
+ def timesteps(self, num_steps: int, start_time=None, dtype=mx.float32):
+ start_time = start_time or (len(self._sigmas) - 1)
+ assert 0 < start_time <= (len(self._sigmas) - 1)
+ steps = _linspace(start_time, 0, num_steps + 1).astype(dtype)
+ return list(zip(steps, steps[1:]))
+
+ def current_timestep(self, step, total_steps, start_time=None):
+ if step < total_steps:
+ steps = self.timesteps(total_steps, start_time)
+ return steps[step]
+ else:
+ return mx.array(0),mx.array(0)
+
+ def step(self, eps_pred, x_t, t, t_prev):
+ sigma = self.sigmas(t).astype(eps_pred.dtype)
+ sigma_prev = self.sigmas(t_prev).astype(eps_pred.dtype)
+
+ dt = sigma_prev - sigma
+ x_t_prev = (sigma.square() + 1).sqrt() * x_t + eps_pred * dt
+
+ x_t_prev = x_t_prev * (sigma_prev.square() + 1).rsqrt()
+
+ return x_t_prev
+
+@dataclass
+class ShardConfig:
+ model_id:str
+ start_layer:int
+ end_layer:int
+ n_layers:int
+
+@dataclass
+class StableDiffusionConfig:
+ model_type:str
+ vae:VAEArgs
+ text_encoder:CLIPArgs
+ scheduler:DiffusionConfig
+ unet:UNetConfig
+ shard:ShardConfig
+
+ @classmethod
+ def from_dict(cls, params):
+ return cls(**{k: v for k, v in params.items() if k in inspect.signature(cls).parameters})
+
+@dataclass
+class ModelArgs(StableDiffusionConfig):
+ shard:Shard = field(default_factory=lambda: Shard("", 0, 0, 0))
+
+ def __post_init__(self):
+ if isinstance(self.shard, dict):
+ self.shard = Shard(**self.shard)
+
+ if not isinstance(self.shard, Shard):
+ raise TypeError(f"Expected shard to be a Shard instance or a dict, got {type(self.shard)} instead")
+
+
+class Model(nn.Module):
+ def __init__(self, config):
+ super().__init__()
+ self.model_type = config.model_type
+ self.config = config
+ self.model_path = config.vae['path'].split('/vae')[0]
+ self.shard = config.shard
+ self.shard_clip, self.shard_unet, self.shard_vae = model_shards(config.shard)
+ self.config_clip=CLIPArgs.from_dict(config.text_encoder['config'])
+ if self.shard_clip.start_layer != -1:
+ self.text_encoder = CLIPTextModel(self.config_clip, shard=self.shard_clip)
+ else:
+ self.text_encoder = nn.Identity()
+ self.tokenizer = load_tokenizer(Path(self.model_path), "vocab.json", "merges.txt")
+ self.diffusion_config = DiffusionConfig.from_dict(config.scheduler['config'])
+ self.sampler = SimpleEulerSampler(self.diffusion_config)
+ if self.shard_unet.start_layer!=-1:
+ self.config_unet = UNetConfig.from_dict(config.unet['config'])
+ self.unet = UNetModel(self.config_unet, self.shard_unet)
+ else:
+ self.unet = nn.Identity()
+ self.config_vae=VAEArgs.from_dict(config.vae['config'])
+ if self.shard_vae.start_layer != -1:
+ self.first_stage_model=Autoencoder(self.config_vae, self.shard_vae)
+ else:
+ self.first_stage_model = nn.Identity()
+
+ def __call__(self,x, step= 0, cfg_weight: float = 7.5,total_steps=50,conditioning=None,mask=None,residual=None,x_t_prev=None,is_finished=False,is_step_finished=False):
+ if self.shard.is_first_layer():
+ x = x.squeeze(0)
+ t, t_prev = self.sampler.current_timestep(step=step, total_steps=total_steps)
+ is_finished = False
+ is_step_finished = False
+ if t.item()==1000:
+ if self.shard_clip.start_layer == 0:
+ conditioning = x
+ if self.shard_clip.start_layer != -1:
+
+ conditioning, mask= self.text_encoder(conditioning,mask)
+ seed = int(time.time())
+ mx.random.seed(seed)
+ if self.shard_unet.is_first_layer():
+ x = self.sampler.sample_prior((1, *(64, 64), self.config_vae.latent_channels_in), dtype=mx.float32)
+ x_t_prev=x
+ # Perform the denoising loop
+ if self.shard_unet.start_layer != -1:
+ with tqdm(total=total_steps,initial=step+1) as pbar:
+ if step<total_steps:
+ if self.shard_unet.is_first_layer():
+ x_t_unet = mx.concatenate([x] * 2, axis=0) if cfg_weight> 1 else x
+ else:
+ x_t_unet = x
+ t_unet = mx.broadcast_to(t, [len(x_t_unet)])
+ x, residual= self.unet(x_t_unet, t_unet, encoder_x=conditioning, residuals=residual)
+ if self.shard_unet.is_last_layer():
+ if cfg_weight > 1:
+ eps_text, eps_neg = x.split(2)
+ eps_pred = eps_neg + cfg_weight * (eps_text - eps_neg)
+ x = self.sampler.step(eps_pred, x_t_prev, t, t_prev)
+ x_t_prev=x
+ mx.eval(x)
+
+ if self.shard_vae.is_last_layer():
+ is_step_finished=True
+ if t_prev.item() ==0:
+ if self.shard_vae.start_layer != -1:
+ x=self.first_stage_model.decode(x)
+ if self.shard_vae.is_last_layer():
+ x = mx.clip(x / 2 + 0.5, 0, 1)
+ x = mx.pad(x, [(0, 0), (8, 8), (8, 8), (0, 0)])
+ B, H, W, C = x.shape
+ x = x.reshape(1, B // 1, H, W, C).transpose(0, 2, 1, 3, 4)
+ x = x.reshape(1 * H, B // 1 * W, C)
+ x = (x * 255).astype(mx.uint8)
+ is_finished=True
+
+ return x, {'conditioning':conditioning, 'mask':mask,'residual':residual,'x_t_prev':x_t_prev,'is_finished':is_finished,'is_step_finished':is_step_finished, 'step':step, 'total_steps':total_steps}
+
+
+ def load(self):
+ if self.shard_vae.start_layer != -1:
+ vae_weights = mx.load(self.config_vae.weight_files[0])
+ vae_weights = self.first_stage_model.sanitize(vae_weights)
+ self.first_stage_model.load_weights(list(vae_weights.items()), strict=True)
+ if self.shard_clip.start_layer != -1:
+ clip_weights = mx.load(self.config_clip.weight_files[0])
+ clip_weights = self.text_encoder.sanitize(clip_weights)
+ self.text_encoder.load_weights(list(clip_weights.items()), strict=True)
+ if self.shard_unet.start_layer !=-1:
+ unet_weights = mx.load(self.config_unet.weight_files[0])
+ unet_weights = self.unet.sanitize(unet_weights)
+ self.unet.load_weights(list(unet_weights.items()), strict=True)
+
+
+def model_shards(shard:ShardConfig):
+ def create_shard(shard, model_ranges):
+ start_layer = shard.start_layer
+ end_layer = shard.end_layer
+
+ shards = {}
+
+ for model_name, (range_start, range_end) in model_ranges.items():
+ if start_layer < range_end and end_layer >= range_start:
+ # Calculate the overlap with the model range
+ overlap_start = max(start_layer, range_start)
+ overlap_end = min(end_layer, range_end - 1)
+
+ # Adjust the layers relative to the model's range
+ relative_start = overlap_start - range_start
+ relative_end = overlap_end - range_start
+ shards[model_name] = Shard(model_name, relative_start, relative_end, range_end - range_start)
+ else:
+ # If no overlap, create a zero-layer shard
+ shards[model_name] = Shard(model_name, -1, -1, range_end - range_start)
+
+ return shards
+
+ # Define the ranges for different models
+ model_ranges = {
+ 'clip': (0, 23),
+ 'unet':(23,32),
+ 'vae': (32, 37) # Example range for unet
+ }
+
+ # Call the function and get the shards for all models
+ shards = create_shard(shard, model_ranges)
+
+ # Access individual shards
+ shard_clip = shards['clip']
+ shard_unet = shards['unet']
+ shard_vae = shards['vae']
+
+ return shard_clip, shard_unet, shard_vae
+
+
+
diff --git a/exo/inference/mlx/models/sd_models/clip.py b/exo/inference/mlx/models/sd_models/clip.py
new file mode 100644
index 00000000..78d95321
--- /dev/null
+++ b/exo/inference/mlx/models/sd_models/clip.py
@@ -0,0 +1,192 @@
+# Adapted from https://github.com/ml-explore/mlx-examples/blob/main/stable_diffusion/stable_diffusion/clip.py
+
+from dataclasses import dataclass
+from typing import List, Optional
+
+import mlx.core as mx
+import mlx.nn as nn
+from dataclasses import field, dataclass
+from exo.inference.shard import Shard
+from exo.inference.mlx.models.base import IdentityBlock
+
+_ACTIVATIONS = {"quick_gelu": nn.gelu_fast_approx, "gelu": nn.gelu}
+
+
+
+@dataclass
+class CLIPTextModelConfig:
+ num_layers: int = 23
+ model_dims: int = 1024
+ num_heads: int = 16
+ max_length: int = 77
+ vocab_size: int = 49408
+ projection_dim: Optional[int] = None
+ hidden_act: str = "quick_gelu"
+
+ @classmethod
+ def from_dict(cls, config):
+ return ModelArgs(
+ num_layers=config["num_hidden_layers"],
+ model_dims=config["hidden_size"],
+ num_heads=config["num_attention_heads"],
+ max_length=config["max_position_embeddings"],
+ vocab_size=config["vocab_size"],
+ projection_dim=config["projection_dim"] if "WithProjection" in config['architectures'][0] else None,
+ hidden_act=config.get("hidden_act", "quick_gelu"),
+ weight_files=config.get("weight_files", [])
+ )
+
+@dataclass
+class ModelArgs(CLIPTextModelConfig):
+ shard: Shard = field(default_factory=lambda: Shard("", 0, 0, 0))
+ weight_files: List[str] = field(default_factory=lambda: [])
+ def __post_init__(self):
+ if isinstance(self.shard, dict):
+ self.shard = Shard(**self.shard)
+
+ if not isinstance(self.shard, Shard):
+ raise TypeError(f"Expected shard to be a Shard instance or a dict, got {type(self.shard)} instead")
+
+ if not self.shard.is_first_layer():
+ self.vision_config = None
+
+
+@dataclass
+class CLIPOutput:
+ pooled_output: Optional[mx.array] = None
+ last_hidden_state: Optional[mx.array] = None
+ hidden_states: Optional[List[mx.array]] = None
+
+
+class CLIPEncoderLayer(nn.Module):
+ """The transformer encoder layer from CLIP."""
+
+ def __init__(self, model_dims: int, num_heads: int, activation: str):
+ super().__init__()
+
+ self.layer_norm1 = nn.LayerNorm(model_dims)
+ self.layer_norm2 = nn.LayerNorm(model_dims)
+
+ self.attention = nn.MultiHeadAttention(model_dims, num_heads)
+ self.attention.query_proj.bias = mx.zeros(model_dims)
+ self.attention.key_proj.bias = mx.zeros(model_dims)
+ self.attention.value_proj.bias = mx.zeros(model_dims)
+ self.attention.out_proj.bias = mx.zeros(model_dims)
+
+ self.linear1 = nn.Linear(model_dims, 4 * model_dims)
+ self.linear2 = nn.Linear(4 * model_dims, model_dims)
+
+ self.act = _ACTIVATIONS[activation]
+
+ def __call__(self, x, attn_mask=None):
+
+ y = self.layer_norm1(x)
+ y = self.attention(y, y, y, attn_mask)
+ x = y + x
+
+ y = self.layer_norm2(x)
+ y = self.linear1(y)
+ y = self.act(y)
+ y = self.linear2(y)
+ x = y + x
+ return x
+
+
+class CLIPTextModel(nn.Module):
+ """Implements the text encoder transformer from CLIP."""
+
+ def __init__(self, config: CLIPTextModelConfig, shard: Shard):
+ super().__init__()
+
+ self.shard = shard
+
+ if self.shard.is_first_layer():
+ self.token_embedding = nn.Embedding(config.vocab_size, config.model_dims)
+ self.position_embedding = nn.Embedding(config.max_length, config.model_dims)
+ self.layers = []
+ for i in range(config.num_layers):
+ if self.shard.start_layer <= i <= self.shard.end_layer:
+ self.layers.append(CLIPEncoderLayer(config.model_dims, config.num_heads, config.hidden_act))
+ else:
+ self.layers.append(IdentityBlock())
+ if self.shard.is_last_layer():
+ self.final_layer_norm = nn.LayerNorm(config.model_dims)
+
+ if config.projection_dim is not None:
+ self.text_projection = nn.Linear(
+ config.model_dims, config.projection_dim, bias=False
+ )
+
+ def _get_mask(self, N, dtype):
+ indices = mx.arange(N)
+ mask = indices[:, None] < indices[None]
+ mask = mask.astype(dtype) * (-6e4 if dtype == mx.float16 else -1e9)
+ return mask
+
+ def __call__(self, x, mask=None):
+ # Extract some shapes
+ if self.shard.is_first_layer():
+ B, N = x.shape
+ eos_tokens = x.argmax(-1)
+
+ # Compute the embeddings
+ x = self.token_embedding(x)
+
+ x = x + self.position_embedding.weight[:N]
+ # Compute the features from the transformer
+ mask = self._get_mask(N, x.dtype)
+
+ hidden_states = []
+ for l in self.layers:
+ x = l(x, mask)
+ hidden_states.append(x)
+ # Apply the final layernorm and return
+
+ if self.shard.is_last_layer():
+ x = self.final_layer_norm(x)
+ last_hidden_state = x
+
+
+
+ return x, mask
+ def sanitize(self, weights):
+ sanitized_weights = {}
+
+ for key, value in weights.items():
+ if "position_ids" in key:
+ continue
+ if key.startswith("text_model."):
+ key = key[11:]
+ if key.startswith("embeddings."):
+ key = key[11:]
+ if key.startswith("encoder."):
+ key = key[8:]
+
+ # Map attention layers
+ if "self_attn." in key:
+ key = key.replace("self_attn.", "attention.")
+ if "q_proj." in key:
+ key = key.replace("q_proj.", "query_proj.")
+ if "k_proj." in key:
+ key = key.replace("k_proj.", "key_proj.")
+ if "v_proj." in key:
+ key = key.replace("v_proj.", "value_proj.")
+
+ # Map ffn layers
+ if "mlp.fc1" in key:
+ key = key.replace("mlp.fc1", "linear1")
+ if "mlp.fc2" in key:
+ key = key.replace("mlp.fc2", "linear2")
+
+ if key.startswith("layers."):
+ layer_num = int(key.split(".")[1])
+ if layer_num < self.shard.start_layer or layer_num > self.shard.end_layer:
+ continue
+ if not self.shard.start_layer == 0 and "embedding" in key:
+ continue
+ if not self.shard.end_layer == 22 and key.startswith("final_layer_norm"):
+ continue
+ if not self.shard.end_layer == 22 and key.startswith("text_projection"):
+ continue
+ sanitized_weights[key] = value
+ return sanitized_weights
diff --git a/exo/inference/mlx/models/sd_models/tokenizer.py b/exo/inference/mlx/models/sd_models/tokenizer.py
new file mode 100644
index 00000000..4987bb90
--- /dev/null
+++ b/exo/inference/mlx/models/sd_models/tokenizer.py
@@ -0,0 +1,131 @@
+# adapted from https://github.com/ml-explore/mlx-examples/blob/main/stable_diffusion/stable_diffusion/tokenizer.py
+
+import regex
+import json
+import glob
+
+
+class Tokenizer:
+ """A simple port of CLIPTokenizer from https://github.com/huggingface/transformers/ ."""
+
+ def __init__(self, bpe_ranks, vocab):
+ self.bpe_ranks = bpe_ranks
+ self.vocab = vocab
+ self.pat = regex.compile(
+ r"""<\|startoftext\|>|<\|endoftext\|>|'s|'t|'re|'ve|'m|'ll|'d|[\p{L}]+|[\p{N}]|[^\s\p{L}\p{N}]+""",
+ regex.IGNORECASE,
+ )
+
+ self._cache = {self.bos: self.bos, self.eos: self.eos}
+
+ @property
+ def bos(self):
+ return "<|startoftext|>"
+
+ @property
+ def bos_token(self):
+ return self.vocab[self.bos]
+
+ @property
+ def eos(self):
+ return "<|endoftext|>"
+
+ @property
+ def eos_token(self):
+ return self.vocab[self.eos]
+
+ def bpe(self, text):
+ if text in self._cache:
+ return self._cache[text]
+
+ unigrams = list(text[:-1]) + [text[-1] + "</w>"]
+ unique_bigrams = set(zip(unigrams, unigrams[1:]))
+
+ if not unique_bigrams:
+ return unigrams
+
+ # In every iteration try to merge the two most likely bigrams. If none
+ # was merged we are done.
+ #
+ # Ported from https://github.com/huggingface/transformers/blob/main/src/transformers/models/clip/tokenization_clip.py
+ while unique_bigrams:
+ bigram = min(
+ unique_bigrams, key=lambda pair: self.bpe_ranks.get(pair, float("inf"))
+ )
+ if bigram not in self.bpe_ranks:
+ break
+
+ new_unigrams = []
+ skip = False
+ for a, b in zip(unigrams, unigrams[1:]):
+ if skip:
+ skip = False
+ continue
+
+ if (a, b) == bigram:
+ new_unigrams.append(a + b)
+ skip = True
+
+ else:
+ new_unigrams.append(a)
+
+ if not skip:
+ new_unigrams.append(b)
+
+ unigrams = new_unigrams
+ unique_bigrams = set(zip(unigrams, unigrams[1:]))
+
+ self._cache[text] = unigrams
+
+ return unigrams
+
+ def tokenize(self, text, prepend_bos=True, append_eos=True):
+ if isinstance(text, list):
+ return [self.tokenize(t, prepend_bos, append_eos) for t in text]
+
+ # Lower case cleanup and split according to self.pat. Hugging Face does
+ # a much more thorough job here but this should suffice for 95% of
+ # cases.
+ clean_text = regex.sub(r"\s+", " ", text.lower())
+ tokens = regex.findall(self.pat, clean_text)
+
+ # Split the tokens according to the byte-pair merge file
+ bpe_tokens = [ti for t in tokens for ti in self.bpe(t)]
+
+ # Map to token ids and return
+ tokens = [self.vocab[t] for t in bpe_tokens]
+ if prepend_bos:
+ tokens = [self.bos_token] + tokens
+ if append_eos:
+ tokens.append(self.eos_token)
+
+ return tokens
+
+ def encode(self, prompt):
+ tokens = [self.tokenize(prompt)]
+ negative_text = ""
+ if negative_text is not None:
+ tokens += [self.tokenize(negative_text)]
+ lengths = [len(t) for t in tokens]
+ N = max(lengths)
+ tokens = [t + [0] * (N - len(t)) for t in tokens]
+ return tokens
+
+def load_tokenizer(
+ model_path: str,
+ vocab_key: str = "tokenizer_vocab",
+ merges_key: str = "tokenizer_merges",
+):
+
+ vocab_file = glob.glob(str(model_path/"tokenizer"/vocab_key))[0]
+ with open(vocab_file, encoding="utf-8") as f:
+ vocab = json.load(f)
+
+ merges_file = glob.glob(str(model_path/"tokenizer"/merges_key))[0]
+ with open(merges_file, encoding="utf-8") as f:
+ bpe_merges = f.read().strip().split("\n")[1 : 49152 - 256 - 2 + 1]
+ bpe_merges = [tuple(m.split()) for m in bpe_merges]
+ bpe_ranks = dict(map(reversed, enumerate(bpe_merges)))
+
+ return Tokenizer(bpe_ranks, vocab)
+
diff --git a/exo/inference/mlx/models/sd_models/unet.py b/exo/inference/mlx/models/sd_models/unet.py
new file mode 100644
index 00000000..3fe44b86
--- /dev/null
+++ b/exo/inference/mlx/models/sd_models/unet.py
@@ -0,0 +1,629 @@
+# Adapted from https://github.com/ml-explore/mlx-examples/blob/main/stable_diffusion/stable_diffusion/unet.py
+
+import math
+from typing import Optional
+
+import mlx.core as mx
+import mlx.nn as nn
+
+from dataclasses import dataclass, field
+from typing import Tuple, Optional, List
+from exo.inference.shard import Shard
+
+@dataclass
+class UNetConfig:
+ in_channels: int = 4
+ out_channels: int = 4
+ conv_in_kernel: int = 3
+ conv_out_kernel: int = 3
+ block_out_channels: Tuple[int] = (320, 640, 1280, 1280)
+ layers_per_block: Tuple[int] = (2, 2, 2, 2)
+ mid_block_layers: int = 2
+ transformer_layers_per_block: Tuple[int] = (1, 1, 1, 1)
+ num_attention_heads: Tuple[int] = (5, 10, 20, 20)
+ cross_attention_dim: Tuple[int] = (1024,) * 4
+ norm_num_groups: int = 32
+ down_block_types: Tuple[str] = (
+ "CrossAttnDownBlock2D",
+ "CrossAttnDownBlock2D",
+ "CrossAttnDownBlock2D",
+ "DownBlock2D",
+ )
+ up_block_types: Tuple[str] = (
+ "UpBlock2D",
+ "CrossAttnUpBlock2D",
+ "CrossAttnUpBlock2D",
+ "CrossAttnUpBlock2D",
+ )
+ addition_embed_type: Optional[str] = None
+ addition_time_embed_dim: Optional[int] = None
+ projection_class_embeddings_input_dim: Optional[int] = None
+ weight_files: List[str] = field(default_factory=lambda: [])
+
+
+
+ @classmethod
+ def from_dict(cls,config):
+ n_blocks = len(config['block_out_channels'])
+ return UNetConfig(
+ in_channels=config["in_channels"],
+ out_channels=config["out_channels"],
+ block_out_channels=config["block_out_channels"],
+ layers_per_block=[config["layers_per_block"]] * n_blocks,
+ transformer_layers_per_block=config.get(
+ "transformer_layers_per_block", (1,) * 4
+ ),
+ num_attention_heads=(
+ [config["attention_head_dim"]] * n_blocks
+ if isinstance(config["attention_head_dim"], int)
+ else config["attention_head_dim"]
+ ),
+ cross_attention_dim=[config["cross_attention_dim"]] * n_blocks,
+ norm_num_groups=config["norm_num_groups"],
+ down_block_types=config["down_block_types"],
+ up_block_types=config["up_block_types"][::-1],
+ addition_embed_type=config.get("addition_embed_type", None),
+ addition_time_embed_dim=config.get("addition_time_embed_dim", None),
+ projection_class_embeddings_input_dim=config.get(
+ "projection_class_embeddings_input_dim", None
+ ),
+ weight_files=config.get("weight_files", [])
+
+ )
+
+
+def upsample_nearest(x, scale: int = 2):
+ B, H, W, C = x.shape
+ x = mx.broadcast_to(x[:, :, None, :, None, :], (B, H, scale, W, scale, C))
+ x = x.reshape(B, H * scale, W * scale, C)
+
+ return x
+
+
+class TimestepEmbedding(nn.Module):
+ def __init__(self, in_channels: int, time_embed_dim: int):
+ super().__init__()
+
+ self.linear_1 = nn.Linear(in_channels, time_embed_dim)
+ self.linear_2 = nn.Linear(time_embed_dim, time_embed_dim)
+
+ def __call__(self, x):
+ x = self.linear_1(x)
+ x = nn.silu(x)
+ x = self.linear_2(x)
+
+ return x
+
+
+class TransformerBlock(nn.Module):
+ def __init__(
+ self,
+ model_dims: int,
+ num_heads: int,
+ hidden_dims: Optional[int] = None,
+ memory_dims: Optional[int] = None,
+ ):
+ super().__init__()
+
+ self.norm1 = nn.LayerNorm(model_dims)
+ self.attn1 = nn.MultiHeadAttention(model_dims, num_heads)
+ self.attn1.out_proj.bias = mx.zeros(model_dims)
+
+ memory_dims = memory_dims or model_dims
+ self.norm2 = nn.LayerNorm(model_dims)
+ self.attn2 = nn.MultiHeadAttention(
+ model_dims, num_heads, key_input_dims=memory_dims
+ )
+ self.attn2.out_proj.bias = mx.zeros(model_dims)
+
+ hidden_dims = hidden_dims or 4 * model_dims
+ self.norm3 = nn.LayerNorm(model_dims)
+ self.linear1 = nn.Linear(model_dims, hidden_dims)
+ self.linear2 = nn.Linear(model_dims, hidden_dims)
+ self.linear3 = nn.Linear(hidden_dims, model_dims)
+
+ def __call__(self, x, memory, attn_mask, memory_mask):
+ # Self attention
+ y = self.norm1(x)
+ y = self.attn1(y, y, y, attn_mask)
+ x = x + y
+
+ # Cross attention
+ y = self.norm2(x)
+ y = self.attn2(y, memory, memory, memory_mask)
+ x = x + y
+
+ # FFN
+ y = self.norm3(x)
+ y_a = self.linear1(y)
+ y_b = self.linear2(y)
+ y = y_a * nn.gelu(y_b)
+ y = self.linear3(y)
+ x = x + y
+
+ return x
+
+
+class Transformer2D(nn.Module):
+ """A transformer model for inputs with 2 spatial dimensions."""
+
+ def __init__(
+ self,
+ in_channels: int,
+ model_dims: int,
+ encoder_dims: int,
+ num_heads: int,
+ num_layers: int = 1,
+ norm_num_groups: int = 32,
+ ):
+ super().__init__()
+
+ self.norm = nn.GroupNorm(norm_num_groups, in_channels, pytorch_compatible=True)
+ self.proj_in = nn.Linear(in_channels, model_dims)
+ self.transformer_blocks = [
+ TransformerBlock(model_dims, num_heads, memory_dims=encoder_dims)
+ for i in range(num_layers)
+ ]
+ self.proj_out = nn.Linear(model_dims, in_channels)
+
+ def __call__(self, x, encoder_x, attn_mask, encoder_attn_mask):
+ # Save the input to add to the output
+ input_x = x
+ dtype = x.dtype
+
+ # Perform the input norm and projection
+ B, H, W, C = x.shape
+ x = self.norm(x.astype(mx.float32)).astype(dtype).reshape(B, -1, C)
+ x = self.proj_in(x)
+
+ # Apply the transformer
+ for block in self.transformer_blocks:
+ x = block(x, encoder_x, attn_mask, encoder_attn_mask)
+
+ # Apply the output projection and reshape
+ x = self.proj_out(x)
+ x = x.reshape(B, H, W, C)
+
+ return x + input_x
+
+
+class ResnetBlock2D(nn.Module):
+ def __init__(
+ self,
+ in_channels: int,
+ out_channels: Optional[int] = None,
+ groups: int = 32,
+ temb_channels: Optional[int] = None,
+ ):
+ super().__init__()
+
+ out_channels = out_channels or in_channels
+
+ self.norm1 = nn.GroupNorm(groups, in_channels, pytorch_compatible=True)
+ self.conv1 = nn.Conv2d(
+ in_channels, out_channels, kernel_size=3, stride=1, padding=1
+ )
+ if temb_channels is not None:
+ self.time_emb_proj = nn.Linear(temb_channels, out_channels)
+ self.norm2 = nn.GroupNorm(groups, out_channels, pytorch_compatible=True)
+ self.conv2 = nn.Conv2d(
+ out_channels, out_channels, kernel_size=3, stride=1, padding=1
+ )
+
+ if in_channels != out_channels:
+ self.conv_shortcut = nn.Linear(in_channels, out_channels)
+
+ def __call__(self, x, temb=None):
+ dtype = x.dtype
+
+ if temb is not None:
+ temb = self.time_emb_proj(nn.silu(temb))
+ y = self.norm1(x.astype(mx.float32)).astype(dtype)
+
+ y = nn.silu(y)
+
+ y = self.conv1(y)
+
+
+ if temb is not None:
+ y = y + temb[:, None, None, :]
+ y = self.norm2(y.astype(mx.float32)).astype(dtype)
+ y = nn.silu(y)
+ y = self.conv2(y)
+
+ x = y + (x if "conv_shortcut" not in self else self.conv_shortcut(x))
+ return x
+
+
+class UNetBlock2D(nn.Module):
+ def __init__(
+ self,
+ in_channels: int,
+ out_channels: int,
+ temb_channels: int,
+ prev_out_channels: Optional[int] = None,
+ num_layers: int = 1,
+ transformer_layers_per_block: int = 1,
+ num_attention_heads: int = 8,
+ cross_attention_dim=1280,
+ resnet_groups: int = 32,
+ add_downsample=True,
+ add_upsample=True,
+ add_cross_attention=True,
+ ):
+ super().__init__()
+
+ # Prepare the in channels list for the resnets
+ if prev_out_channels is None:
+ in_channels_list = [in_channels] + [out_channels] * (num_layers - 1)
+ else:
+ in_channels_list = [prev_out_channels] + [out_channels] * (num_layers - 1)
+ res_channels_list = [out_channels] * (num_layers - 1) + [in_channels]
+ in_channels_list = [
+ a + b for a, b in zip(in_channels_list, res_channels_list)
+ ]
+
+ # Add resnet blocks that also process the time embedding
+ self.resnets = [
+ ResnetBlock2D(
+ in_channels=ic,
+ out_channels=out_channels,
+ temb_channels=temb_channels,
+ groups=resnet_groups,
+ )
+ for ic in in_channels_list
+ ]
+
+ # Add optional cross attention layers
+ if add_cross_attention:
+ self.attentions = [
+ Transformer2D(
+ in_channels=out_channels,
+ model_dims=out_channels,
+ num_heads=num_attention_heads,
+ num_layers=transformer_layers_per_block,
+ encoder_dims=cross_attention_dim,
+ )
+ for i in range(num_layers)
+ ]
+
+ # Add an optional downsampling layer
+ if add_downsample:
+ self.downsample = nn.Conv2d(
+ out_channels, out_channels, kernel_size=3, stride=2, padding=1
+ )
+
+ # or upsampling layer
+ if add_upsample:
+ self.upsample = nn.Conv2d(
+ out_channels, out_channels, kernel_size=3, stride=1, padding=1
+ )
+
+ def __call__(
+ self,
+ x,
+ encoder_x=None,
+ temb=None,
+ attn_mask=None,
+ encoder_attn_mask=None,
+ residual_hidden_states=None,
+ ):
+ output_states = []
+
+ for i in range(len(self.resnets)):
+ if residual_hidden_states is not None:
+ x = mx.concatenate([x, residual_hidden_states.pop()], axis=-1)
+
+ x = self.resnets[i](x, temb)
+
+ if "attentions" in self:
+ x = self.attentions[i](x, encoder_x, attn_mask, encoder_attn_mask)
+
+ output_states.append(x)
+
+ if "downsample" in self:
+ x = self.downsample(x)
+ output_states.append(x)
+
+ if "upsample" in self:
+ x = self.upsample(upsample_nearest(x))
+ output_states.append(x)
+
+ return x, output_states
+
+
+class UNetModel(nn.Module):
+ """The conditional 2D UNet model that actually performs the denoising."""
+
+ def __init__(self, config: UNetConfig, shard: Shard):
+ super().__init__()
+ self.shard = shard
+ self.start_layer = shard.start_layer
+ self.end_layer = shard.end_layer
+ self.layers_range = list(range(self.start_layer, self.end_layer+1))
+ if shard.is_first_layer():
+ self.conv_in = nn.Conv2d(
+ config.in_channels,
+ config.block_out_channels[0],
+ config.conv_in_kernel,
+ padding=(config.conv_in_kernel - 1) // 2,
+ )
+
+ self.timesteps = nn.SinusoidalPositionalEncoding(
+ config.block_out_channels[0],
+ max_freq=1,
+ min_freq=math.exp(
+ -math.log(10000) + 2 * math.log(10000) / config.block_out_channels[0]
+ ),
+ scale=1.0,
+ cos_first=True,
+ full_turns=False,
+ )
+ self.time_embedding = TimestepEmbedding(
+ config.block_out_channels[0],
+ config.block_out_channels[0] * 4,
+ )
+
+ if config.addition_embed_type == "text_time":
+ self.add_time_proj = nn.SinusoidalPositionalEncoding(
+ config.addition_time_embed_dim,
+ max_freq=1,
+ min_freq=math.exp(
+ -math.log(10000)
+ + 2 * math.log(10000) / config.addition_time_embed_dim
+ ),
+ scale=1.0,
+ cos_first=True,
+ full_turns=False,
+ )
+ self.add_embedding = TimestepEmbedding(
+ config.projection_class_embeddings_input_dim,
+ config.block_out_channels[0] * 4,
+ )
+
+ # Make the downsampling blocks
+ block_channels = [config.block_out_channels[0]] + list(
+ config.block_out_channels
+ )
+ self.down_blocks = []
+
+ for i, (in_channels, out_channels) in enumerate(zip(block_channels, block_channels[1:])):
+ if i in self.layers_range:
+ self.down_blocks.append(
+ UNetBlock2D(
+ in_channels=in_channels,
+ out_channels=out_channels,
+ temb_channels=config.block_out_channels[0] * 4,
+ num_layers=config.layers_per_block[i],
+ transformer_layers_per_block=config.transformer_layers_per_block[i],
+ num_attention_heads=config.num_attention_heads[i],
+ cross_attention_dim=config.cross_attention_dim[i],
+ resnet_groups=config.norm_num_groups,
+ add_downsample=(i < len(config.block_out_channels) - 1),
+ add_upsample=False,
+ add_cross_attention="CrossAttn" in config.down_block_types[i],
+ )
+ )
+ else:
+ self.down_blocks.append(nn.Identity())
+
+
+ # Make the middle block
+ if 4 in self.layers_range:
+ self.mid_blocks = [
+ ResnetBlock2D(
+ in_channels=config.block_out_channels[-1],
+ out_channels=config.block_out_channels[-1],
+ temb_channels=config.block_out_channels[0] * 4,
+ groups=config.norm_num_groups,
+ ),
+ Transformer2D(
+ in_channels=config.block_out_channels[-1],
+ model_dims=config.block_out_channels[-1],
+ num_heads=config.num_attention_heads[-1],
+ num_layers=config.transformer_layers_per_block[-1],
+ encoder_dims=config.cross_attention_dim[-1],
+ ),
+ ResnetBlock2D(
+ in_channels=config.block_out_channels[-1],
+ out_channels=config.block_out_channels[-1],
+ temb_channels=config.block_out_channels[0] * 4,
+ groups=config.norm_num_groups,
+ ),
+ ]
+
+ # Make the upsampling blocks
+ block_channels = (
+ [config.block_out_channels[0]]
+ + list(config.block_out_channels)
+ + [config.block_out_channels[-1]]
+ )
+
+ total_items = len(block_channels) - 3
+ reversed_channels = list(reversed(list(zip(block_channels, block_channels[1:], block_channels[2:]))))
+
+ self.up_blocks = []
+ for rev_i, (in_channels, out_channels, prev_out_channels) in enumerate(reversed_channels):
+ i = total_items - rev_i
+ if rev_i+5 in self.layers_range:
+ self.up_blocks.append(
+ UNetBlock2D(
+ in_channels=in_channels,
+ out_channels=out_channels,
+ temb_channels=config.block_out_channels[0] * 4,
+ prev_out_channels=prev_out_channels,
+ num_layers=config.layers_per_block[i] + 1,
+ transformer_layers_per_block=config.transformer_layers_per_block[i],
+ num_attention_heads=config.num_attention_heads[i],
+ cross_attention_dim=config.cross_attention_dim[i],
+ resnet_groups=config.norm_num_groups,
+ add_downsample=False,
+ add_upsample=(i > 0),
+ add_cross_attention="CrossAttn" in config.up_block_types[i],
+ )
+ )
+ else:
+ self.up_blocks.append(nn.Identity())
+
+
+ if shard.is_last_layer():
+ self.conv_norm_out = nn.GroupNorm(
+ config.norm_num_groups,
+ config.block_out_channels[0],
+ pytorch_compatible=True,
+ )
+ self.conv_out = nn.Conv2d(
+ config.block_out_channels[0],
+ config.out_channels,
+ config.conv_out_kernel,
+ padding=(config.conv_out_kernel - 1) // 2,
+ )
+
+ def __call__(
+ self,
+ x,
+ timestep,
+ encoder_x,
+ attn_mask=None,
+ encoder_attn_mask=None,
+ text_time=None,
+ residuals=None,
+ ):
+ # Compute the time embeddings
+
+ temb = self.timesteps(timestep).astype(x.dtype)
+ temb = self.time_embedding(temb)
+
+ # Add the extra text_time conditioning
+ if text_time is not None:
+ text_emb, time_ids = text_time
+ emb = self.add_time_proj(time_ids).flatten(1).astype(x.dtype)
+ emb = mx.concatenate([text_emb, emb], axis=-1)
+ emb = self.add_embedding(emb)
+ temb = temb + emb
+
+ if self.shard.is_first_layer():
+ # Preprocess the input
+ x = self.conv_in(x)
+ residuals = [x]
+ # Run the downsampling part of the unet
+
+ for i in range(len(self.down_blocks)):
+ if i in self.layers_range:
+ x, res = self.down_blocks[i](
+ x,
+ encoder_x=encoder_x,
+ temb=temb,
+ attn_mask=attn_mask,
+ encoder_attn_mask=encoder_attn_mask,
+ )
+ residuals.extend(res)
+ else:
+ x= self.down_blocks[i](x)
+
+ if 4 in self.layers_range:
+ # Run the middle part of the unet
+ x = self.mid_blocks[0](x, temb)
+ x = self.mid_blocks[1](x, encoder_x, attn_mask, encoder_attn_mask)
+ x = self.mid_blocks[2](x, temb)
+
+ # Run the upsampling part of the unet
+ for i in range(len(self.up_blocks)):
+ if i+5 in self.layers_range:
+ x, _ = self.up_blocks[i](
+ x,
+ encoder_x=encoder_x,
+ temb=temb,
+ attn_mask=attn_mask,
+ encoder_attn_mask=encoder_attn_mask,
+ residual_hidden_states=residuals,
+ )
+ else:
+ x= self.up_blocks[i](x)
+
+ # Postprocess the output
+ if self.shard.is_last_layer():
+ dtype = x.dtype
+ x = self.conv_norm_out(x.astype(mx.float32)).astype(dtype)
+ x = nn.silu(x)
+ x = self.conv_out(x)
+
+ return x, residuals
+ def sanitize(self, weights):
+ sanitized_weights = {}
+ for key, value in weights.items():
+ k1=""
+ k2=""
+ if "downsamplers" in key:
+ key = key.replace("downsamplers.0.conv", "downsample")
+ if "upsamplers" in key:
+ key = key.replace("upsamplers.0.conv", "upsample")
+
+ # Map the mid block
+ if "mid_block.resnets.0" in key:
+ key = key.replace("mid_block.resnets.0", "mid_blocks.0")
+ if "mid_block.attentions.0" in key:
+ key = key.replace("mid_block.attentions.0", "mid_blocks.1")
+ if "mid_block.resnets.1" in key:
+ key = key.replace("mid_block.resnets.1", "mid_blocks.2")
+
+ # Map attention layers
+ if "to_k" in key:
+ key = key.replace("to_k", "key_proj")
+ if "to_out.0" in key:
+ key = key.replace("to_out.0", "out_proj")
+ if "to_q" in key:
+ key = key.replace("to_q", "query_proj")
+ if "to_v" in key:
+ key = key.replace("to_v", "value_proj")
+
+ # Map transformer ffn
+ if "ff.net.2" in key:
+ key = key.replace("ff.net.2", "linear3")
+ if "ff.net.0" in key:
+ k1 = key.replace("ff.net.0.proj", "linear1")
+ k2 = key.replace("ff.net.0.proj", "linear2")
+ v1, v2 = mx.split(value, 2)
+
+
+ if "conv_shortcut.weight" in key:
+ value = value.squeeze()
+
+ # Transform the weights from 1x1 convs to linear
+ if len(value.shape) == 4 and ("proj_in" in key or "proj_out" in key):
+ value = value.squeeze()
+
+ if len(value.shape) == 4:
+ value = value.transpose(0, 2, 3, 1)
+ value = value.reshape(-1).reshape(value.shape)
+
+ if key.startswith("conv_in") :
+ if 0 not in self.layers_range:
+ continue
+
+ if key.startswith("down_blocks"):
+ layer_num = int(key.split(".")[1])
+ if layer_num not in self.layers_range:
+ continue
+
+ if key.startswith("mid_block"):
+ if 4 not in self.layers_range:
+ continue
+
+ if key.startswith("up_blocks"):
+ layer_num = int(key.split(".")[1])
+ if (layer_num+5) not in self.layers_range:
+ continue
+
+ if key.startswith("conv_out") or key.startswith("conv_norm_out"):
+ if 8 not in self.layers_range:
+ continue
+
+ if len(k1)>0:
+ sanitized_weights[k1] = v1
+ sanitized_weights[k2] = v2
+ else:
+ sanitized_weights[key] = value
+
+
+ return sanitized_weights
diff --git a/exo/inference/mlx/models/sd_models/vae.py b/exo/inference/mlx/models/sd_models/vae.py
new file mode 100644
index 00000000..39037bdf
--- /dev/null
+++ b/exo/inference/mlx/models/sd_models/vae.py
@@ -0,0 +1,390 @@
+# Adapted from https://github.com/ml-explore/mlx-examples/blob/main/stable_diffusion/stable_diffusion/vae.py
+
+import math
+from typing import List
+
+import mlx.core as mx
+import mlx.nn as nn
+
+from .unet import ResnetBlock2D, upsample_nearest
+from dataclasses import dataclass, field
+from exo.inference.shard import Shard
+from typing import Tuple
+import inspect
+from ..base import IdentityBlock
+
+@dataclass
+class AutoencoderConfig:
+ in_channels: int = 3
+ out_channels: int = 3
+ latent_channels_out: int = 8
+ latent_channels_in: int = 4
+ block_out_channels: Tuple[int] = (128, 256, 512, 512)
+ layers_per_block: int = 2
+ norm_num_groups: int = 32
+ scaling_factor: float = 0.18215
+ weight_files: List[str] = field(default_factory=lambda: [])
+ @classmethod
+ def from_dict(cls, params):
+ return cls(**{k: v for k, v in params.items() if k in inspect.signature(cls).parameters})
+
+
+@dataclass
+class ModelArgs(AutoencoderConfig):
+ shard: Shard = field(default_factory=lambda: Shard("", 0, 0, 0))
+
+ def __post_init__(self):
+ if isinstance(self.shard, dict):
+ self.shard = Shard(**self.shard)
+
+ if not isinstance(self.shard, Shard):
+ raise TypeError(f"Expected shard to be a Shard instance or a dict, got {type(self.shard)} instead")
+
+ if not self.shard.is_first_layer():
+ self.vision_config = None
+
+
+class Attention(nn.Module):
+ """A single head unmasked attention for use with the VAE."""
+
+ def __init__(self, dims: int, norm_groups: int = 32):
+ super().__init__()
+
+ self.group_norm = nn.GroupNorm(norm_groups, dims, pytorch_compatible=True)
+ self.query_proj = nn.Linear(dims, dims)
+ self.key_proj = nn.Linear(dims, dims)
+ self.value_proj = nn.Linear(dims, dims)
+ self.out_proj = nn.Linear(dims, dims)
+
+ def __call__(self, x):
+ B, H, W, C = x.shape
+
+ y = self.group_norm(x)
+
+ queries = self.query_proj(y).reshape(B, H * W, C)
+ keys = self.key_proj(y).reshape(B, H * W, C)
+ values = self.value_proj(y).reshape(B, H * W, C)
+
+ scale = 1 / math.sqrt(queries.shape[-1])
+ scores = (queries * scale) @ keys.transpose(0, 2, 1)
+ attn = mx.softmax(scores, axis=-1)
+ y = (attn @ values).reshape(B, H, W, C)
+
+ y = self.out_proj(y)
+ x = x + y
+
+ return x
+
+
+class EncoderDecoderBlock2D(nn.Module):
+ def __init__(
+ self,
+ in_channels: int,
+ out_channels: int,
+ num_layers: int = 1,
+ resnet_groups: int = 32,
+ add_downsample=True,
+ add_upsample=True,
+ ):
+ super().__init__()
+
+ # Add the resnet blocks
+ self.resnets = [
+ ResnetBlock2D(
+ in_channels=in_channels if i == 0 else out_channels,
+ out_channels=out_channels,
+ groups=resnet_groups,
+ )
+ for i in range(num_layers)
+ ]
+
+ # Add an optional downsampling layer
+ if add_downsample:
+ self.downsample = nn.Conv2d(
+ out_channels, out_channels, kernel_size=3, stride=2, padding=0
+ )
+
+ # or upsampling layer
+ if add_upsample:
+ self.upsample = nn.Conv2d(
+ out_channels, out_channels, kernel_size=3, stride=1, padding=1
+ )
+
+ def __call__(self, x):
+ for resnet in self.resnets:
+ x = resnet(x)
+ if "downsample" in self:
+ x = mx.pad(x, [(0, 0), (0, 1), (0, 1), (0, 0)])
+ x = self.downsample(x)
+
+ if "upsample" in self:
+ x = self.upsample(upsample_nearest(x))
+ return x
+
+
+class Encoder(nn.Module):
+ """Implements the encoder side of the Autoencoder."""
+
+ def __init__(
+ self,
+ in_channels: int,
+ out_channels: int,
+ block_out_channels: List[int] = [64],
+ layers_per_block: int = 2,
+ resnet_groups: int = 32,
+ ):
+ super().__init__()
+
+ self.conv_in = nn.Conv2d(
+ in_channels, block_out_channels[0], kernel_size=3, stride=1, padding=1
+ )
+
+ channels = [block_out_channels[0]] + list(block_out_channels)
+ self.down_blocks = [
+ EncoderDecoderBlock2D(
+ in_channels,
+ out_channels,
+ num_layers=layers_per_block,
+ resnet_groups=resnet_groups,
+ add_downsample=i < len(block_out_channels) - 1,
+ add_upsample=False,
+ )
+ for i, (in_channels, out_channels) in enumerate(zip(channels, channels[1:]))
+ ]
+
+ self.mid_blocks = [
+ ResnetBlock2D(
+ in_channels=block_out_channels[-1],
+ out_channels=block_out_channels[-1],
+ groups=resnet_groups,
+ ),
+ Attention(block_out_channels[-1], resnet_groups),
+ ResnetBlock2D(
+ in_channels=block_out_channels[-1],
+ out_channels=block_out_channels[-1],
+ groups=resnet_groups,
+ ),
+ ]
+
+ self.conv_norm_out = nn.GroupNorm(
+ resnet_groups, block_out_channels[-1], pytorch_compatible=True
+ )
+ self.conv_out = nn.Conv2d(block_out_channels[-1], out_channels, 3, padding=1)
+
+ def __call__(self, x):
+ x = self.conv_in(x)
+
+ for l in self.down_blocks:
+ x = l(x)
+
+ x = self.mid_blocks[0](x)
+ x = self.mid_blocks[1](x)
+ x = self.mid_blocks[2](x)
+
+ x = self.conv_norm_out(x)
+ x = nn.silu(x)
+ x = self.conv_out(x)
+
+ return x
+
+
+class Decoder(nn.Module):
+ """Implements the decoder side of the Autoencoder."""
+
+ def __init__(
+ self,
+ in_channels: int,
+ out_channels: int,
+ shard: Shard,
+ layer_range: List[int],
+ block_out_channels: List[int] = [64],
+ layers_per_block: int = 2,
+ resnet_groups: int = 32,
+ ):
+ super().__init__()
+ self.out_channels = out_channels
+ self.layers_range = layer_range
+ if 0 in layer_range:
+ self.conv_in = nn.Conv2d(
+ in_channels, block_out_channels[-1], kernel_size=3, stride=1, padding=1
+ )
+
+ if 0 in layer_range:
+ self.mid_blocks = [
+ ResnetBlock2D(
+ in_channels=block_out_channels[-1],
+ out_channels=block_out_channels[-1],
+ groups=resnet_groups,
+ ),
+ Attention(block_out_channels[-1], resnet_groups),
+ ResnetBlock2D(
+ in_channels=block_out_channels[-1],
+ out_channels=block_out_channels[-1],
+ groups=resnet_groups,
+ ),
+ ]
+
+ channels = list(reversed(block_out_channels))
+ channels = [channels[0]] + channels
+
+ self.up_blocks = []
+ current_layer = 1
+
+ for i, (in_channels, out_channels) in enumerate(zip(channels, channels[1:])):
+ if current_layer in layer_range:
+ self.up_blocks.append(
+ EncoderDecoderBlock2D(
+ in_channels,
+ out_channels,
+ num_layers=layers_per_block,
+ resnet_groups=resnet_groups,
+ add_downsample=False,
+ add_upsample=i < len(block_out_channels) - 1,
+ )
+ )
+ else:
+ self.up_blocks.append(IdentityBlock())
+ current_layer += 1
+ if 4 in layer_range:
+ self.conv_norm_out = nn.GroupNorm(
+ resnet_groups, block_out_channels[0], pytorch_compatible=True
+ )
+ self.conv_out = nn.Conv2d(block_out_channels[0], self.out_channels, 3, padding=1)
+
+
+ def __call__(self, x):
+ if 0 in self.layers_range:
+ x = self.conv_in(x)
+ x = self.mid_blocks[0](x)
+ x = self.mid_blocks[1](x)
+ x = self.mid_blocks[2](x)
+
+ for l in self.up_blocks:
+ x = l(x)
+ if 4 in self.layers_range:
+ x = self.conv_norm_out(x)
+ x = nn.silu(x)
+ x = self.conv_out(x)
+ return x
+
+
+class Autoencoder(nn.Module):
+ """The autoencoder that allows us to perform diffusion in the latent space."""
+
+ def __init__(self, config: AutoencoderConfig, shard: Shard):
+ super().__init__()
+ self.shard = shard
+ self.start_layer = shard.start_layer
+ self.end_layer = shard.end_layer
+ self.layers_range = list(range(self.start_layer, self.end_layer+1))
+ self.latent_channels = config.latent_channels_in
+ self.scaling_factor = config.scaling_factor
+ self.decoder_only = True # stable diffusion text to speech only uses decoder from the autoencoder
+ if not self.decoder_only:
+ self.encoder = Encoder(
+ config.in_channels,
+ config.latent_channels_out,
+ config.block_out_channels,
+ config.layers_per_block,
+ resnet_groups=config.norm_num_groups,
+ )
+ self.quant_proj = nn.Linear(
+ config.latent_channels_out, config.latent_channels_out
+ )
+ self.decoder = Decoder(
+ config.latent_channels_in,
+ config.out_channels,
+ shard,
+ self.layers_range,
+ config.block_out_channels,
+ config.layers_per_block + 1,
+ resnet_groups=config.norm_num_groups,
+ )
+ if 0 in self.layers_range:
+ self.post_quant_proj = nn.Linear(
+ config.latent_channels_in, config.latent_channels_in
+ )
+
+ def decode(self, z):
+ if 0 in self.layers_range:
+ z = z / self.scaling_factor
+ z=self.post_quant_proj(z)
+ return self.decoder(z)
+
+ def encode(self, x):
+ x = self.encoder(x)
+ x = self.quant_proj(x)
+ mean, logvar = x.split(2, axis=-1)
+ mean = mean * self.scaling_factor
+ logvar = logvar + 2 * math.log(self.scaling_factor)
+
+ return mean, logvar
+
+ def __call__(self, x, key=None):
+ mean, logvar = self.encode(x)
+ z = mx.random.normal(mean.shape, key=key) * mx.exp(0.5 * logvar) + mean
+ x_hat = self.decode(z)
+
+ return dict(x_hat=x_hat, z=z, mean=mean, logvar=logvar)
+
+ def sanitize(self, weights):
+ layers = self.layers_range
+ sanitized_weights = {}
+ for key, value in weights.items():
+ if 'decoder' in key and self.decoder_only:
+ if "downsamplers" in key:
+ key = key.replace("downsamplers.0.conv", "downsample")
+ if "upsamplers" in key:
+ key = key.replace("upsamplers.0.conv", "upsample")
+
+ # Map attention layers
+ if "key" in key:
+ key = key.replace("key", "key_proj")
+ if "proj_attn" in key:
+ key = key.replace("proj_attn", "out_proj")
+ if "query" in key:
+ key = key.replace("query", "query_proj")
+ if "value" in key:
+ key = key.replace("value", "value_proj")
+
+ # Map the mid block
+ if "mid_block.resnets.0" in key:
+ key = key.replace("mid_block.resnets.0", "mid_blocks.0")
+ if "mid_block.attentions.0" in key:
+ key = key.replace("mid_block.attentions.0", "mid_blocks.1")
+ if "mid_block.resnets.1" in key:
+ key = key.replace("mid_block.resnets.1", "mid_blocks.2")
+
+ # Map the quant/post_quant layers
+ if "quant_conv" in key:
+ key = key.replace("quant_conv", "quant_proj")
+ value = value.squeeze()
+
+ # Map the conv_shortcut to linear
+ if "conv_shortcut.weight" in key:
+ value = value.squeeze()
+
+ if len(value.shape) == 4:
+ value = value.transpose(0, 2, 3, 1)
+ value = value.reshape(-1).reshape(value.shape)
+
+ if key.startswith("decoder.mid_blocks."):
+ if 0 in layers:
+ sanitized_weights[key] = value
+ if "conv_in" in key and 0 in layers:
+ sanitized_weights[key] = value
+ if key.startswith("decoder.up_blocks."):
+ layer_num = int(key.split(".")[2])+1
+ if layer_num in layers:
+ sanitized_weights[key] = value
+ if key.startswith("decoder.conv_norm_out") and 4 in layers:
+ sanitized_weights[key] = value
+ if key.startswith("decoder.conv_out") and 4 in layers:
+ sanitized_weights[key] = value
+
+ if "post_quant_conv" in key and 0 in layers:
+ key = key.replace("quant_conv", "quant_proj")
+ value = value.squeeze()
+ sanitized_weights[key] = value
+ return sanitized_weights
+
diff --git a/exo/inference/mlx/sharded_inference_engine.py b/exo/inference/mlx/sharded_inference_engine.py
index e43a65da..dc50c1d6 100644
--- a/exo/inference/mlx/sharded_inference_engine.py
+++ b/exo/inference/mlx/sharded_inference_engine.py
@@ -53,10 +53,11 @@ class MLXDynamicShardInferenceEngine(InferenceEngine):
tokens = await asyncio.get_running_loop().run_in_executor(self.executor, self.tokenizer.decode, tokens)
return tokens
- async def infer_tensor(self, request_id: str, shard: Shard, input_data: np.ndarray) -> np.ndarray:
+ async def infer_tensor(self, request_id: str, shard: Shard, input_data: np.ndarray, inference_state: Optional[dict] = None) -> np.ndarray:
await self.ensure_shard(shard)
- output_data: np.ndarray = np.array(await asyncio.get_running_loop().run_in_executor(self.executor, self.model, mx.array(input_data), request_id))
- return output_data
+ output_data, inference_state = await asyncio.get_running_loop().run_in_executor(self.executor, self.model, mx.array(input_data), request_id, inference_state)
+ output_data = np.array(output_data)
+ return output_data, inference_state
async def ensure_shard(self, shard: Shard):
if self.shard == shard:
diff --git a/exo/inference/mlx/sharded_utils.py b/exo/inference/mlx/sharded_utils.py
index 1a6343e9..3a23badc 100644
--- a/exo/inference/mlx/sharded_utils.py
+++ b/exo/inference/mlx/sharded_utils.py
@@ -61,8 +61,16 @@ def _get_classes(config: dict):
def load_config(model_path: Path) -> dict:
try:
- with open(model_path/"config.json", "r") as f:
- config = json.load(f)
+ config_path = model_path / "config.json"
+ if config_path.exists():
+ with open(config_path, "r") as f:
+ config = json.load(f)
+ return config
+
+ model_index_path = model_path / "model_index.json"
+ if model_index_path.exists():
+ config = load_model_index(model_path, model_index_path)
+ return config
except FileNotFoundError:
logging.error(f"Config file not found in {model_path}")
raise
@@ -109,6 +117,24 @@ def load_model_shard(
# Try weight for back-compat
weight_files = glob.glob(str(model_path/"weight*.safetensors"))
+ model_class, model_args_class = _get_classes(config=config)
+
+ class ShardedModel(model_class):
+ def __init__(self, args):
+ super().__init__(args)
+ self.shard = Shard(args.shard.model_id, args.shard.start_layer, args.shard.end_layer, args.shard.n_layers)
+
+ def __call__(self, x, *args, **kwargs):
+ y = super().__call__(x[None] if self.shard.is_first_layer() else x, *args, **kwargs)
+ return y
+
+ model_args = model_args_class.from_dict(config)
+ model = ShardedModel(model_args)
+
+ if config.get("model_index", False):
+ model.load()
+ return model
+
if not weight_files:
logging.error(f"No safetensors found in {model_path}")
raise FileNotFoundError(f"No safetensors found in {model_path}")
@@ -128,19 +154,7 @@ def load_model_shard(
weights.update(mx.load(wf))
- model_class, model_args_class = _get_classes(config=config)
-
- class ShardedModel(model_class):
- def __init__(self, args):
- super().__init__(args)
- self.shard = Shard(args.shard.model_id, args.shard.start_layer, args.shard.end_layer, args.shard.n_layers)
-
- def __call__(self, x, *args, **kwargs):
- y = super().__call__(x[None] if self.shard.is_first_layer() else x, *args, **kwargs)
- return y
-
- model_args = model_args_class.from_dict(config)
- model = ShardedModel(model_args)
+
if hasattr(model, "sanitize"):
weights = model.sanitize(weights)
@@ -182,6 +196,9 @@ async def load_shard(
processor.eos_token_id = processor.tokenizer.eos_token_id
processor.encode = processor.tokenizer.encode
return model, processor
+ elif hasattr(model, "tokenizer"):
+ tokenizer = model.tokenizer
+ return model, tokenizer
else:
tokenizer = load_tokenizer(model_path, tokenizer_config)
return model, tokenizer
@@ -210,3 +227,29 @@ async def get_image_from_str(_image_str: str):
return img
else:
raise ValueError("Invalid image_str format. Must be a URL or a base64 encoded image.")
+
+# loading a combined config for all models in the index
+def load_model_index(model_path: Path, model_index_path: Path):
+ models_config = {}
+ with open(model_index_path, "r") as f:
+ model_index = json.load(f)
+ models_config["model_index"] = True
+ models_config["model_type"] = model_index["_class_name"]
+ models_config["models"] = {}
+ for model in model_index.keys():
+ model_config_path = glob.glob(str(model_path / model / "*config.json"))
+ if len(model_config_path)>0:
+ with open(model_config_path[0], "r") as f:
+ model_config = { }
+ model_config["model_type"] = model
+ model_config["config"] = json.load(f)
+ model_config["path"] = model_path / model
+ if model_config["path"]/"*model.safetensors":
+ model_config["config"].update({"weight_files": list(glob.glob(str(model_config["path"]/"*model.safetensors")))})
+ model_config["path"] = str(model_path / model)
+ m = {}
+ m[model] = model_config
+ models_config.update(m)
+ models_config = json.dumps(models_config)
+ models_config = json.loads(models_config)
+ return models_config
diff --git a/exo/inference/mlx/stateful_model.py b/exo/inference/mlx/stateful_model.py
index ff213ace..79e9baeb 100644
--- a/exo/inference/mlx/stateful_model.py
+++ b/exo/inference/mlx/stateful_model.py
@@ -1,4 +1,4 @@
-from typing import Dict, Tuple
+from typing import Dict, Tuple, Optional
from collections import OrderedDict
import mlx.core as mx
@@ -29,14 +29,17 @@ class StatefulModel(nn.Module):
self.caches[request_id] = cache
- def __call__(self, x, request_id: str):
- if request_id not in self.caches:
- self.init_cache(request_id)
- else:
- self.caches.move_to_end(request_id)
+ def __call__(self, x, request_id: str, inference_state: Optional[dict] = None):
+ if self.model.model_type !='StableDiffusionPipeline':
+ if request_id not in self.caches:
+ self.init_cache(request_id)
+ else:
+ self.caches.move_to_end(request_id)
- cache = self.caches[request_id]
+ cache = self.caches[request_id]
- y = self.model(x, cache=cache)
- return y
+ y = self.model(x, cache=cache)
+ else:
+ y, inference_state = self.model(x, **inference_state)
+ return y, inference_state
diff --git a/exo/main.py b/exo/main.py
index b6701c6d..03de33d0 100644
--- a/exo/main.py
+++ b/exo/main.py
@@ -126,7 +126,7 @@ api = ChatGPTAPI(
default_model=args.default_model
)
node.on_token.register("update_topology_viz").on_next(
- lambda req_id, tokens, __: topology_viz.update_prompt_output(req_id, inference_engine.tokenizer.decode(tokens)) if topology_viz and hasattr(inference_engine, "tokenizer") else None
+ lambda req_id, tokens, __: topology_viz.update_prompt_output(req_id, inference_engine.tokenizer.decode(tokens)) if topology_viz and hasattr(inference_engine, "tokenizer") and inference_engine.shard.model_id != 'stable-diffusion-2-1-base' else None
)
diff --git a/exo/models.py b/exo/models.py
index 1fb567a6..8262251a 100644
--- a/exo/models.py
+++ b/exo/models.py
@@ -80,6 +80,8 @@ model_cards = {
# gemma
"gemma2-9b": { "layers": 42, "repo": { "MLXDynamicShardInferenceEngine": "mlx-community/gemma-2-9b-it-4bit", }, },
"gemma2-27b": { "layers": 46, "repo": { "MLXDynamicShardInferenceEngine": "mlx-community/gemma-2-27b-it-4bit", }, },
+ # stable diffusion
+ "stable-diffusion-2-1-base": { "layers": 37, "repo": { "MLXDynamicShardInferenceEngine": "stabilityai/stable-diffusion-2-1-base" } },
# dummy
"dummy": { "layers": 8, "repo": { "DummyInferenceEngine": "dummy", }, },
}
@@ -113,6 +115,7 @@ pretty_name = {
"qwen-2.5-math-72b": "Qwen 2.5 72B (Math)",
"llama-3-8b": "Llama 3 8B",
"llama-3-70b": "Llama 3 70B",
+ "stable-diffusion-2-1-base": "Stable Diffusion 2.1",
}
def get_repo(model_id: str, inference_engine_classname: str) -> Optional[str]:
diff --git a/exo/networking/grpc/grpc_peer_handle.py b/exo/networking/grpc/grpc_peer_handle.py
index a788b87a..98220e33 100644
--- a/exo/networking/grpc/grpc_peer_handle.py
+++ b/exo/networking/grpc/grpc_peer_handle.py
@@ -11,7 +11,8 @@ from exo.inference.shard import Shard
from exo.topology.topology import Topology
from exo.topology.device_capabilities import DeviceCapabilities, DeviceFlops
from exo.helpers import DEBUG
-
+import json
+import mlx.core as mx
class GRPCPeerHandle(PeerHandle):
def __init__(self, _id: str, address: str, device_capabilities: DeviceCapabilities):
@@ -85,7 +86,7 @@ class GRPCPeerHandle(PeerHandle):
return np.frombuffer(response.tensor_data, dtype=np.dtype(response.dtype)).reshape(response.shape)
- async def send_tensor(self, shard: Shard, tensor: np.ndarray, request_id: Optional[str] = None) -> Optional[np.array]:
+ async def send_tensor(self, shard: Shard, tensor: np.ndarray, inference_state: Optional[dict] = None, request_id: Optional[str] = None) -> Optional[np.array]:
request = node_service_pb2.TensorRequest(
shard=node_service_pb2.Shard(
model_id=shard.model_id,
@@ -95,6 +96,7 @@ class GRPCPeerHandle(PeerHandle):
),
tensor=node_service_pb2.Tensor(tensor_data=tensor.tobytes(), shape=tensor.shape, dtype=str(tensor.dtype)),
request_id=request_id,
+ inference_state=self.serialize_inference_state(inference_state)
)
response = await self.stub.SendTensor(request)
@@ -128,9 +130,43 @@ class GRPCPeerHandle(PeerHandle):
return topology
async def send_result(self, request_id: str, result: List[int], is_finished: bool) -> None:
- request = node_service_pb2.SendResultRequest(request_id=request_id, result=result, is_finished=is_finished)
+ tensor = None
+ if isinstance(result, np.ndarray):
+ tensor = node_service_pb2.Tensor(tensor_data=result.tobytes(), shape=result.shape, dtype=str(result.dtype))
+ result = []
+ request = node_service_pb2.SendResultRequest(request_id=request_id, result=result, tensor=tensor, is_finished=is_finished)
await self.stub.SendResult(request)
async def send_opaque_status(self, request_id: str, status: str) -> None:
request = node_service_pb2.SendOpaqueStatusRequest(request_id=request_id, status=status)
await self.stub.SendOpaqueStatus(request)
+
+ def serialize_inference_state(self, inference_state: dict) -> node_service_pb2.InferenceState:
+ proto_inference_state = node_service_pb2.InferenceState()
+ other_data = {}
+ for k, v in inference_state.items():
+ if isinstance(v, mx.array):
+ np_array = np.array(v)
+ tensor_data = node_service_pb2.Tensor(
+ tensor_data=np_array.tobytes(),
+ shape=list(np_array.shape),
+ dtype=str(np_array.dtype)
+ )
+ proto_inference_state.tensor_data[k].CopyFrom(tensor_data)
+ elif isinstance(v, list) and all(isinstance(item, mx.array) for item in v):
+ tensor_list = node_service_pb2.TensorList()
+ for tensor in v:
+ np_array = np.array(tensor)
+ tensor_data = node_service_pb2.Tensor(
+ tensor_data=np_array.tobytes(),
+ shape=list(np_array.shape),
+ dtype=str(np_array.dtype)
+ )
+ tensor_list.tensors.append(tensor_data)
+ proto_inference_state.tensor_list_data[k].CopyFrom(tensor_list)
+ else:
+ # For non-tensor data, we'll still use JSON
+ other_data[k] = v
+ if other_data:
+ proto_inference_state.other_data_json = json.dumps(other_data)
+ return proto_inference_state
diff --git a/exo/networking/grpc/grpc_server.py b/exo/networking/grpc/grpc_server.py
index db489475..c03ba507 100644
--- a/exo/networking/grpc/grpc_server.py
+++ b/exo/networking/grpc/grpc_server.py
@@ -8,6 +8,8 @@ from . import node_service_pb2_grpc
from exo import DEBUG
from exo.inference.shard import Shard
from exo.orchestration import Node
+import json
+import mlx.core as mx
class GRPCServer(node_service_pb2_grpc.NodeServiceServicer):
@@ -65,7 +67,9 @@ class GRPCServer(node_service_pb2_grpc.NodeServiceServicer):
tensor = np.frombuffer(request.tensor.tensor_data, dtype=np.dtype(request.tensor.dtype)).reshape(request.tensor.shape)
request_id = request.request_id
- result = await self.node.process_tensor(shard, tensor, request_id)
+ inference_state = self.deserialize_inference_state(request.inference_state)
+
+ result = await self.node.process_tensor(shard, tensor, request_id, inference_state)
if DEBUG >= 5: print(f"SendTensor tensor {shard=} {tensor=} {request_id=} result: {result}")
tensor_data = result.tobytes() if result is not None else None
return node_service_pb2.Tensor(tensor_data=tensor_data, shape=result.shape, dtype=str(result.dtype)) if result is not None else node_service_pb2.Tensor()
@@ -104,7 +108,11 @@ class GRPCServer(node_service_pb2_grpc.NodeServiceServicer):
request_id = request.request_id
result = request.result
is_finished = request.is_finished
+ img = request.tensor
if DEBUG >= 5: print(f"Received SendResult request: {request_id=} {result=} {is_finished=}")
+ result = list(result)
+ if len(img.tensor_data) > 0:
+ result=np.frombuffer(img.tensor_data, dtype=np.dtype(img.dtype)).reshape(img.shape)
self.node.on_token.trigger_all(request_id, result, is_finished)
return node_service_pb2.Empty()
@@ -117,3 +125,22 @@ class GRPCServer(node_service_pb2_grpc.NodeServiceServicer):
async def HealthCheck(self, request, context):
return node_service_pb2.HealthCheckResponse(is_healthy=True)
+
+ def deserialize_inference_state(self,inference_state_proto: node_service_pb2.InferenceState) -> dict:
+ inference_state = {}
+
+ for k, tensor_data in inference_state_proto.tensor_data.items():
+ np_array = np.frombuffer(tensor_data.tensor_data, dtype=tensor_data.dtype).reshape(tensor_data.shape)
+ inference_state[k] = mx.array(np_array)
+
+ for k, tensor_list in inference_state_proto.tensor_list_data.items():
+ inference_state[k] = [
+ mx.array(np.frombuffer(tensor.tensor_data, dtype=tensor.dtype).reshape(tensor.shape))
+ for tensor in tensor_list.tensors
+ ]
+
+ if inference_state_proto.other_data_json:
+ other_data = json.loads(inference_state_proto.other_data_json)
+ inference_state.update(other_data)
+
+ return inference_state
diff --git a/exo/networking/grpc/node_service.proto b/exo/networking/grpc/node_service.proto
index 146976b6..a04c997b 100644
--- a/exo/networking/grpc/node_service.proto
+++ b/exo/networking/grpc/node_service.proto
@@ -29,6 +29,7 @@ message TensorRequest {
Shard shard = 1;
Tensor tensor = 2;
optional string request_id = 3;
+ optional InferenceState inference_state = 4;
}
message GetInferenceResultRequest {
@@ -46,6 +47,16 @@ message Tensor {
string dtype = 3;
}
+message TensorList {
+ repeated Tensor tensors = 1;
+}
+
+message InferenceState {
+ map<string, Tensor> tensor_data = 1;
+ map<string, TensorList> tensor_list_data = 2;
+ string other_data_json = 3;
+}
+
message CollectTopologyRequest {
repeated string visited = 1;
int32 max_depth = 2;
@@ -76,7 +87,8 @@ message DeviceCapabilities {
message SendResultRequest {
string request_id = 1;
repeated int32 result = 2;
- bool is_finished = 3;
+ optional Tensor tensor = 3;
+ bool is_finished = 4;
}
message SendOpaqueStatusRequest {
diff --git a/exo/networking/grpc/node_service_pb2.py b/exo/networking/grpc/node_service_pb2.py
index ab9e6bac..78f4a75b 100644
--- a/exo/networking/grpc/node_service_pb2.py
+++ b/exo/networking/grpc/node_service_pb2.py
@@ -1,6 +1,6 @@
# -*- coding: utf-8 -*-
# Generated by the protocol buffer compiler. DO NOT EDIT!
-# source: node_service.proto
+# source: exo/networking/grpc/node_service.proto
# Protobuf Python Version: 5.26.1
"""Generated protocol buffer code."""
from google.protobuf import descriptor as _descriptor
@@ -14,53 +14,65 @@ _sym_db = _symbol_database.Default()
-DESCRIPTOR = _descriptor_pool.Default().AddSerializedFile(b'\n\x12node_service.proto\x12\x0cnode_service\"S\n\x05Shard\x12\x10\n\x08model_id\x18\x01 \x01(\t\x12\x13\n\x0bstart_layer\x18\x02 \x01(\x05\x12\x11\n\tend_layer\x18\x03 \x01(\x05\x12\x10\n\x08n_layers\x18\x04 \x01(\x05\"k\n\rPromptRequest\x12\"\n\x05shard\x18\x01 \x01(\x0b\x32\x13.node_service.Shard\x12\x0e\n\x06prompt\x18\x02 \x01(\t\x12\x17\n\nrequest_id\x18\x03 \x01(\tH\x00\x88\x01\x01\x42\r\n\x0b_request_id\"\x81\x01\n\rTensorRequest\x12\"\n\x05shard\x18\x01 \x01(\x0b\x32\x13.node_service.Shard\x12$\n\x06tensor\x18\x02 \x01(\x0b\x32\x14.node_service.Tensor\x12\x17\n\nrequest_id\x18\x03 \x01(\tH\x00\x88\x01\x01\x42\r\n\x0b_request_id\"/\n\x19GetInferenceResultRequest\x12\x12\n\nrequest_id\x18\x01 \x01(\t\"\\\n\x0fInferenceResult\x12)\n\x06tensor\x18\x01 \x01(\x0b\x32\x14.node_service.TensorH\x00\x88\x01\x01\x12\x13\n\x0bis_finished\x18\x02 \x01(\x08\x42\t\n\x07_tensor\";\n\x06Tensor\x12\x13\n\x0btensor_data\x18\x01 \x01(\x0c\x12\r\n\x05shape\x18\x02 \x03(\x05\x12\r\n\x05\x64type\x18\x03 \x01(\t\"<\n\x16\x43ollectTopologyRequest\x12\x0f\n\x07visited\x18\x01 \x03(\t\x12\x11\n\tmax_depth\x18\x02 \x01(\x05\"\x8e\x02\n\x08Topology\x12\x30\n\x05nodes\x18\x01 \x03(\x0b\x32!.node_service.Topology.NodesEntry\x12\x39\n\npeer_graph\x18\x02 \x03(\x0b\x32%.node_service.Topology.PeerGraphEntry\x1aN\n\nNodesEntry\x12\x0b\n\x03key\x18\x01 \x01(\t\x12/\n\x05value\x18\x02 \x01(\x0b\x32 .node_service.DeviceCapabilities:\x02\x38\x01\x1a\x45\n\x0ePeerGraphEntry\x12\x0b\n\x03key\x18\x01 \x01(\t\x12\"\n\x05value\x18\x02 \x01(\x0b\x32\x13.node_service.Peers:\x02\x38\x01\"\x19\n\x05Peers\x12\x10\n\x08peer_ids\x18\x01 \x03(\t\"7\n\x0b\x44\x65viceFlops\x12\x0c\n\x04\x66p32\x18\x01 \x01(\x02\x12\x0c\n\x04\x66p16\x18\x02 \x01(\x02\x12\x0c\n\x04int8\x18\x03 \x01(\x02\"k\n\x12\x44\x65viceCapabilities\x12\r\n\x05model\x18\x01 \x01(\t\x12\x0c\n\x04\x63hip\x18\x02 \x01(\t\x12\x0e\n\x06memory\x18\x03 \x01(\x05\x12(\n\x05\x66lops\x18\x04 \x01(\x0b\x32\x19.node_service.DeviceFlops\"L\n\x11SendResultRequest\x12\x12\n\nrequest_id\x18\x01 \x01(\t\x12\x0e\n\x06result\x18\x02 \x03(\x05\x12\x13\n\x0bis_finished\x18\x03 \x01(\x08\"=\n\x17SendOpaqueStatusRequest\x12\x12\n\nrequest_id\x18\x01 \x01(\t\x12\x0e\n\x06status\x18\x02 \x01(\t\"\x14\n\x12HealthCheckRequest\")\n\x13HealthCheckResponse\x12\x12\n\nis_healthy\x18\x01 \x01(\x08\"\x07\n\x05\x45mpty2\xb4\x04\n\x0bNodeService\x12\x41\n\nSendPrompt\x12\x1b.node_service.PromptRequest\x1a\x14.node_service.Tensor\"\x00\x12\x41\n\nSendTensor\x12\x1b.node_service.TensorRequest\x1a\x14.node_service.Tensor\"\x00\x12^\n\x12GetInferenceResult\x12\'.node_service.GetInferenceResultRequest\x1a\x1d.node_service.InferenceResult\"\x00\x12Q\n\x0f\x43ollectTopology\x12$.node_service.CollectTopologyRequest\x1a\x16.node_service.Topology\"\x00\x12\x44\n\nSendResult\x12\x1f.node_service.SendResultRequest\x1a\x13.node_service.Empty\"\x00\x12P\n\x10SendOpaqueStatus\x12%.node_service.SendOpaqueStatusRequest\x1a\x13.node_service.Empty\"\x00\x12T\n\x0bHealthCheck\x12 .node_service.HealthCheckRequest\x1a!.node_service.HealthCheckResponse\"\x00\x62\x06proto3')
+DESCRIPTOR = _descriptor_pool.Default().AddSerializedFile(b'\n&exo/networking/grpc/node_service.proto\x12\x0cnode_service\"S\n\x05Shard\x12\x10\n\x08model_id\x18\x01 \x01(\t\x12\x13\n\x0bstart_layer\x18\x02 \x01(\x05\x12\x11\n\tend_layer\x18\x03 \x01(\x05\x12\x10\n\x08n_layers\x18\x04 \x01(\x05\"k\n\rPromptRequest\x12\"\n\x05shard\x18\x01 \x01(\x0b\x32\x13.node_service.Shard\x12\x0e\n\x06prompt\x18\x02 \x01(\t\x12\x17\n\nrequest_id\x18\x03 \x01(\tH\x00\x88\x01\x01\x42\r\n\x0b_request_id\"\xd1\x01\n\rTensorRequest\x12\"\n\x05shard\x18\x01 \x01(\x0b\x32\x13.node_service.Shard\x12$\n\x06tensor\x18\x02 \x01(\x0b\x32\x14.node_service.Tensor\x12\x17\n\nrequest_id\x18\x03 \x01(\tH\x00\x88\x01\x01\x12:\n\x0finference_state\x18\x04 \x01(\x0b\x32\x1c.node_service.InferenceStateH\x01\x88\x01\x01\x42\r\n\x0b_request_idB\x12\n\x10_inference_state\"/\n\x19GetInferenceResultRequest\x12\x12\n\nrequest_id\x18\x01 \x01(\t\"\\\n\x0fInferenceResult\x12)\n\x06tensor\x18\x01 \x01(\x0b\x32\x14.node_service.TensorH\x00\x88\x01\x01\x12\x13\n\x0bis_finished\x18\x02 \x01(\x08\x42\t\n\x07_tensor\";\n\x06Tensor\x12\x13\n\x0btensor_data\x18\x01 \x01(\x0c\x12\r\n\x05shape\x18\x02 \x03(\x05\x12\r\n\x05\x64type\x18\x03 \x01(\t\"3\n\nTensorList\x12%\n\x07tensors\x18\x01 \x03(\x0b\x32\x14.node_service.Tensor\"\xd2\x02\n\x0eInferenceState\x12\x41\n\x0btensor_data\x18\x01 \x03(\x0b\x32,.node_service.InferenceState.TensorDataEntry\x12J\n\x10tensor_list_data\x18\x02 \x03(\x0b\x32\x30.node_service.InferenceState.TensorListDataEntry\x12\x17\n\x0fother_data_json\x18\x03 \x01(\t\x1aG\n\x0fTensorDataEntry\x12\x0b\n\x03key\x18\x01 \x01(\t\x12#\n\x05value\x18\x02 \x01(\x0b\x32\x14.node_service.Tensor:\x02\x38\x01\x1aO\n\x13TensorListDataEntry\x12\x0b\n\x03key\x18\x01 \x01(\t\x12\'\n\x05value\x18\x02 \x01(\x0b\x32\x18.node_service.TensorList:\x02\x38\x01\"<\n\x16\x43ollectTopologyRequest\x12\x0f\n\x07visited\x18\x01 \x03(\t\x12\x11\n\tmax_depth\x18\x02 \x01(\x05\"\x8e\x02\n\x08Topology\x12\x30\n\x05nodes\x18\x01 \x03(\x0b\x32!.node_service.Topology.NodesEntry\x12\x39\n\npeer_graph\x18\x02 \x03(\x0b\x32%.node_service.Topology.PeerGraphEntry\x1aN\n\nNodesEntry\x12\x0b\n\x03key\x18\x01 \x01(\t\x12/\n\x05value\x18\x02 \x01(\x0b\x32 .node_service.DeviceCapabilities:\x02\x38\x01\x1a\x45\n\x0ePeerGraphEntry\x12\x0b\n\x03key\x18\x01 \x01(\t\x12\"\n\x05value\x18\x02 \x01(\x0b\x32\x13.node_service.Peers:\x02\x38\x01\"\x19\n\x05Peers\x12\x10\n\x08peer_ids\x18\x01 \x03(\t\"7\n\x0b\x44\x65viceFlops\x12\x0c\n\x04\x66p32\x18\x01 \x01(\x02\x12\x0c\n\x04\x66p16\x18\x02 \x01(\x02\x12\x0c\n\x04int8\x18\x03 \x01(\x02\"k\n\x12\x44\x65viceCapabilities\x12\r\n\x05model\x18\x01 \x01(\t\x12\x0c\n\x04\x63hip\x18\x02 \x01(\t\x12\x0e\n\x06memory\x18\x03 \x01(\x05\x12(\n\x05\x66lops\x18\x04 \x01(\x0b\x32\x19.node_service.DeviceFlops\"\x82\x01\n\x11SendResultRequest\x12\x12\n\nrequest_id\x18\x01 \x01(\t\x12\x0e\n\x06result\x18\x02 \x03(\x05\x12)\n\x06tensor\x18\x03 \x01(\x0b\x32\x14.node_service.TensorH\x00\x88\x01\x01\x12\x13\n\x0bis_finished\x18\x04 \x01(\x08\x42\t\n\x07_tensor\"=\n\x17SendOpaqueStatusRequest\x12\x12\n\nrequest_id\x18\x01 \x01(\t\x12\x0e\n\x06status\x18\x02 \x01(\t\"\x14\n\x12HealthCheckRequest\")\n\x13HealthCheckResponse\x12\x12\n\nis_healthy\x18\x01 \x01(\x08\"\x07\n\x05\x45mpty2\xb4\x04\n\x0bNodeService\x12\x41\n\nSendPrompt\x12\x1b.node_service.PromptRequest\x1a\x14.node_service.Tensor\"\x00\x12\x41\n\nSendTensor\x12\x1b.node_service.TensorRequest\x1a\x14.node_service.Tensor\"\x00\x12^\n\x12GetInferenceResult\x12\'.node_service.GetInferenceResultRequest\x1a\x1d.node_service.InferenceResult\"\x00\x12Q\n\x0f\x43ollectTopology\x12$.node_service.CollectTopologyRequest\x1a\x16.node_service.Topology\"\x00\x12\x44\n\nSendResult\x12\x1f.node_service.SendResultRequest\x1a\x13.node_service.Empty\"\x00\x12P\n\x10SendOpaqueStatus\x12%.node_service.SendOpaqueStatusRequest\x1a\x13.node_service.Empty\"\x00\x12T\n\x0bHealthCheck\x12 .node_service.HealthCheckRequest\x1a!.node_service.HealthCheckResponse\"\x00\x62\x06proto3')
_globals = globals()
_builder.BuildMessageAndEnumDescriptors(DESCRIPTOR, _globals)
-_builder.BuildTopDescriptorsAndMessages(DESCRIPTOR, 'node_service_pb2', _globals)
+_builder.BuildTopDescriptorsAndMessages(DESCRIPTOR, 'exo.networking.grpc.node_service_pb2', _globals)
if not _descriptor._USE_C_DESCRIPTORS:
DESCRIPTOR._loaded_options = None
+ _globals['_INFERENCESTATE_TENSORDATAENTRY']._loaded_options = None
+ _globals['_INFERENCESTATE_TENSORDATAENTRY']._serialized_options = b'8\001'
+ _globals['_INFERENCESTATE_TENSORLISTDATAENTRY']._loaded_options = None
+ _globals['_INFERENCESTATE_TENSORLISTDATAENTRY']._serialized_options = b'8\001'
_globals['_TOPOLOGY_NODESENTRY']._loaded_options = None
_globals['_TOPOLOGY_NODESENTRY']._serialized_options = b'8\001'
_globals['_TOPOLOGY_PEERGRAPHENTRY']._loaded_options = None
_globals['_TOPOLOGY_PEERGRAPHENTRY']._serialized_options = b'8\001'
- _globals['_SHARD']._serialized_start=36
- _globals['_SHARD']._serialized_end=119
- _globals['_PROMPTREQUEST']._serialized_start=121
- _globals['_PROMPTREQUEST']._serialized_end=228
- _globals['_TENSORREQUEST']._serialized_start=231
- _globals['_TENSORREQUEST']._serialized_end=360
- _globals['_GETINFERENCERESULTREQUEST']._serialized_start=362
- _globals['_GETINFERENCERESULTREQUEST']._serialized_end=409
- _globals['_INFERENCERESULT']._serialized_start=411
- _globals['_INFERENCERESULT']._serialized_end=503
- _globals['_TENSOR']._serialized_start=505
- _globals['_TENSOR']._serialized_end=564
- _globals['_COLLECTTOPOLOGYREQUEST']._serialized_start=566
- _globals['_COLLECTTOPOLOGYREQUEST']._serialized_end=626
- _globals['_TOPOLOGY']._serialized_start=629
- _globals['_TOPOLOGY']._serialized_end=899
- _globals['_TOPOLOGY_NODESENTRY']._serialized_start=750
- _globals['_TOPOLOGY_NODESENTRY']._serialized_end=828
- _globals['_TOPOLOGY_PEERGRAPHENTRY']._serialized_start=830
- _globals['_TOPOLOGY_PEERGRAPHENTRY']._serialized_end=899
- _globals['_PEERS']._serialized_start=901
- _globals['_PEERS']._serialized_end=926
- _globals['_DEVICEFLOPS']._serialized_start=928
- _globals['_DEVICEFLOPS']._serialized_end=983
- _globals['_DEVICECAPABILITIES']._serialized_start=985
- _globals['_DEVICECAPABILITIES']._serialized_end=1092
- _globals['_SENDRESULTREQUEST']._serialized_start=1094
- _globals['_SENDRESULTREQUEST']._serialized_end=1170
- _globals['_SENDOPAQUESTATUSREQUEST']._serialized_start=1172
- _globals['_SENDOPAQUESTATUSREQUEST']._serialized_end=1233
- _globals['_HEALTHCHECKREQUEST']._serialized_start=1235
- _globals['_HEALTHCHECKREQUEST']._serialized_end=1255
- _globals['_HEALTHCHECKRESPONSE']._serialized_start=1257
- _globals['_HEALTHCHECKRESPONSE']._serialized_end=1298
- _globals['_EMPTY']._serialized_start=1300
- _globals['_EMPTY']._serialized_end=1307
- _globals['_NODESERVICE']._serialized_start=1310
- _globals['_NODESERVICE']._serialized_end=1874
+ _globals['_SHARD']._serialized_start=56
+ _globals['_SHARD']._serialized_end=139
+ _globals['_PROMPTREQUEST']._serialized_start=141
+ _globals['_PROMPTREQUEST']._serialized_end=248
+ _globals['_TENSORREQUEST']._serialized_start=251
+ _globals['_TENSORREQUEST']._serialized_end=460
+ _globals['_GETINFERENCERESULTREQUEST']._serialized_start=462
+ _globals['_GETINFERENCERESULTREQUEST']._serialized_end=509
+ _globals['_INFERENCERESULT']._serialized_start=511
+ _globals['_INFERENCERESULT']._serialized_end=603
+ _globals['_TENSOR']._serialized_start=605
+ _globals['_TENSOR']._serialized_end=664
+ _globals['_TENSORLIST']._serialized_start=666
+ _globals['_TENSORLIST']._serialized_end=717
+ _globals['_INFERENCESTATE']._serialized_start=720
+ _globals['_INFERENCESTATE']._serialized_end=1058
+ _globals['_INFERENCESTATE_TENSORDATAENTRY']._serialized_start=906
+ _globals['_INFERENCESTATE_TENSORDATAENTRY']._serialized_end=977
+ _globals['_INFERENCESTATE_TENSORLISTDATAENTRY']._serialized_start=979
+ _globals['_INFERENCESTATE_TENSORLISTDATAENTRY']._serialized_end=1058
+ _globals['_COLLECTTOPOLOGYREQUEST']._serialized_start=1060
+ _globals['_COLLECTTOPOLOGYREQUEST']._serialized_end=1120
+ _globals['_TOPOLOGY']._serialized_start=1123
+ _globals['_TOPOLOGY']._serialized_end=1393
+ _globals['_TOPOLOGY_NODESENTRY']._serialized_start=1244
+ _globals['_TOPOLOGY_NODESENTRY']._serialized_end=1322
+ _globals['_TOPOLOGY_PEERGRAPHENTRY']._serialized_start=1324
+ _globals['_TOPOLOGY_PEERGRAPHENTRY']._serialized_end=1393
+ _globals['_PEERS']._serialized_start=1395
+ _globals['_PEERS']._serialized_end=1420
+ _globals['_DEVICEFLOPS']._serialized_start=1422
+ _globals['_DEVICEFLOPS']._serialized_end=1477
+ _globals['_DEVICECAPABILITIES']._serialized_start=1479
+ _globals['_DEVICECAPABILITIES']._serialized_end=1586
+ _globals['_SENDRESULTREQUEST']._serialized_start=1589
+ _globals['_SENDRESULTREQUEST']._serialized_end=1719
+ _globals['_SENDOPAQUESTATUSREQUEST']._serialized_start=1721
+ _globals['_SENDOPAQUESTATUSREQUEST']._serialized_end=1782
+ _globals['_HEALTHCHECKREQUEST']._serialized_start=1784
+ _globals['_HEALTHCHECKREQUEST']._serialized_end=1804
+ _globals['_HEALTHCHECKRESPONSE']._serialized_start=1806
+ _globals['_HEALTHCHECKRESPONSE']._serialized_end=1847
+ _globals['_EMPTY']._serialized_start=1849
+ _globals['_EMPTY']._serialized_end=1856
+ _globals['_NODESERVICE']._serialized_start=1859
+ _globals['_NODESERVICE']._serialized_end=2423
# @@protoc_insertion_point(module_scope)
diff --git a/exo/networking/grpc/node_service_pb2_grpc.py b/exo/networking/grpc/node_service_pb2_grpc.py
index dd166ca9..aa8d8993 100644
--- a/exo/networking/grpc/node_service_pb2_grpc.py
+++ b/exo/networking/grpc/node_service_pb2_grpc.py
@@ -3,7 +3,7 @@
import grpc
import warnings
-from . import node_service_pb2 as node__service__pb2
+from exo.networking.grpc import node_service_pb2 as exo_dot_networking_dot_grpc_dot_node__service__pb2
GRPC_GENERATED_VERSION = '1.64.1'
GRPC_VERSION = grpc.__version__
@@ -20,7 +20,7 @@ except ImportError:
if _version_not_supported:
warnings.warn(
f'The grpc package installed is at version {GRPC_VERSION},'
- + f' but the generated code in node_service_pb2_grpc.py depends on'
+ + f' but the generated code in exo/networking/grpc/node_service_pb2_grpc.py depends on'
+ f' grpcio>={GRPC_GENERATED_VERSION}.'
+ f' Please upgrade your grpc module to grpcio>={GRPC_GENERATED_VERSION}'
+ f' or downgrade your generated code using grpcio-tools<={GRPC_VERSION}.'
@@ -41,38 +41,38 @@ class NodeServiceStub(object):
"""
self.SendPrompt = channel.unary_unary(
'/node_service.NodeService/SendPrompt',
- request_serializer=node__service__pb2.PromptRequest.SerializeToString,
- response_deserializer=node__service__pb2.Tensor.FromString,
+ request_serializer=exo_dot_networking_dot_grpc_dot_node__service__pb2.PromptRequest.SerializeToString,
+ response_deserializer=exo_dot_networking_dot_grpc_dot_node__service__pb2.Tensor.FromString,
_registered_method=True)
self.SendTensor = channel.unary_unary(
'/node_service.NodeService/SendTensor',
- request_serializer=node__service__pb2.TensorRequest.SerializeToString,
- response_deserializer=node__service__pb2.Tensor.FromString,
+ request_serializer=exo_dot_networking_dot_grpc_dot_node__service__pb2.TensorRequest.SerializeToString,
+ response_deserializer=exo_dot_networking_dot_grpc_dot_node__service__pb2.Tensor.FromString,
_registered_method=True)
self.GetInferenceResult = channel.unary_unary(
'/node_service.NodeService/GetInferenceResult',
- request_serializer=node__service__pb2.GetInferenceResultRequest.SerializeToString,
- response_deserializer=node__service__pb2.InferenceResult.FromString,
+ request_serializer=exo_dot_networking_dot_grpc_dot_node__service__pb2.GetInferenceResultRequest.SerializeToString,
+ response_deserializer=exo_dot_networking_dot_grpc_dot_node__service__pb2.InferenceResult.FromString,
_registered_method=True)
self.CollectTopology = channel.unary_unary(
'/node_service.NodeService/CollectTopology',
- request_serializer=node__service__pb2.CollectTopologyRequest.SerializeToString,
- response_deserializer=node__service__pb2.Topology.FromString,
+ request_serializer=exo_dot_networking_dot_grpc_dot_node__service__pb2.CollectTopologyRequest.SerializeToString,
+ response_deserializer=exo_dot_networking_dot_grpc_dot_node__service__pb2.Topology.FromString,
_registered_method=True)
self.SendResult = channel.unary_unary(
'/node_service.NodeService/SendResult',
- request_serializer=node__service__pb2.SendResultRequest.SerializeToString,
- response_deserializer=node__service__pb2.Empty.FromString,
+ request_serializer=exo_dot_networking_dot_grpc_dot_node__service__pb2.SendResultRequest.SerializeToString,
+ response_deserializer=exo_dot_networking_dot_grpc_dot_node__service__pb2.Empty.FromString,
_registered_method=True)
self.SendOpaqueStatus = channel.unary_unary(
'/node_service.NodeService/SendOpaqueStatus',
- request_serializer=node__service__pb2.SendOpaqueStatusRequest.SerializeToString,
- response_deserializer=node__service__pb2.Empty.FromString,
+ request_serializer=exo_dot_networking_dot_grpc_dot_node__service__pb2.SendOpaqueStatusRequest.SerializeToString,
+ response_deserializer=exo_dot_networking_dot_grpc_dot_node__service__pb2.Empty.FromString,
_registered_method=True)
self.HealthCheck = channel.unary_unary(
'/node_service.NodeService/HealthCheck',
- request_serializer=node__service__pb2.HealthCheckRequest.SerializeToString,
- response_deserializer=node__service__pb2.HealthCheckResponse.FromString,
+ request_serializer=exo_dot_networking_dot_grpc_dot_node__service__pb2.HealthCheckRequest.SerializeToString,
+ response_deserializer=exo_dot_networking_dot_grpc_dot_node__service__pb2.HealthCheckResponse.FromString,
_registered_method=True)
@@ -126,38 +126,38 @@ def add_NodeServiceServicer_to_server(servicer, server):
rpc_method_handlers = {
'SendPrompt': grpc.unary_unary_rpc_method_handler(
servicer.SendPrompt,
- request_deserializer=node__service__pb2.PromptRequest.FromString,
- response_serializer=node__service__pb2.Tensor.SerializeToString,
+ request_deserializer=exo_dot_networking_dot_grpc_dot_node__service__pb2.PromptRequest.FromString,
+ response_serializer=exo_dot_networking_dot_grpc_dot_node__service__pb2.Tensor.SerializeToString,
),
'SendTensor': grpc.unary_unary_rpc_method_handler(
servicer.SendTensor,
- request_deserializer=node__service__pb2.TensorRequest.FromString,
- response_serializer=node__service__pb2.Tensor.SerializeToString,
+ request_deserializer=exo_dot_networking_dot_grpc_dot_node__service__pb2.TensorRequest.FromString,
+ response_serializer=exo_dot_networking_dot_grpc_dot_node__service__pb2.Tensor.SerializeToString,
),
'GetInferenceResult': grpc.unary_unary_rpc_method_handler(
servicer.GetInferenceResult,
- request_deserializer=node__service__pb2.GetInferenceResultRequest.FromString,
- response_serializer=node__service__pb2.InferenceResult.SerializeToString,
+ request_deserializer=exo_dot_networking_dot_grpc_dot_node__service__pb2.GetInferenceResultRequest.FromString,
+ response_serializer=exo_dot_networking_dot_grpc_dot_node__service__pb2.InferenceResult.SerializeToString,
),
'CollectTopology': grpc.unary_unary_rpc_method_handler(
servicer.CollectTopology,
- request_deserializer=node__service__pb2.CollectTopologyRequest.FromString,
- response_serializer=node__service__pb2.Topology.SerializeToString,
+ request_deserializer=exo_dot_networking_dot_grpc_dot_node__service__pb2.CollectTopologyRequest.FromString,
+ response_serializer=exo_dot_networking_dot_grpc_dot_node__service__pb2.Topology.SerializeToString,
),
'SendResult': grpc.unary_unary_rpc_method_handler(
servicer.SendResult,
- request_deserializer=node__service__pb2.SendResultRequest.FromString,
- response_serializer=node__service__pb2.Empty.SerializeToString,
+ request_deserializer=exo_dot_networking_dot_grpc_dot_node__service__pb2.SendResultRequest.FromString,
+ response_serializer=exo_dot_networking_dot_grpc_dot_node__service__pb2.Empty.SerializeToString,
),
'SendOpaqueStatus': grpc.unary_unary_rpc_method_handler(
servicer.SendOpaqueStatus,
- request_deserializer=node__service__pb2.SendOpaqueStatusRequest.FromString,
- response_serializer=node__service__pb2.Empty.SerializeToString,
+ request_deserializer=exo_dot_networking_dot_grpc_dot_node__service__pb2.SendOpaqueStatusRequest.FromString,
+ response_serializer=exo_dot_networking_dot_grpc_dot_node__service__pb2.Empty.SerializeToString,
),
'HealthCheck': grpc.unary_unary_rpc_method_handler(
servicer.HealthCheck,
- request_deserializer=node__service__pb2.HealthCheckRequest.FromString,
- response_serializer=node__service__pb2.HealthCheckResponse.SerializeToString,
+ request_deserializer=exo_dot_networking_dot_grpc_dot_node__service__pb2.HealthCheckRequest.FromString,
+ response_serializer=exo_dot_networking_dot_grpc_dot_node__service__pb2.HealthCheckResponse.SerializeToString,
),
}
generic_handler = grpc.method_handlers_generic_handler(
@@ -185,8 +185,8 @@ class NodeService(object):
request,
target,
'/node_service.NodeService/SendPrompt',
- node__service__pb2.PromptRequest.SerializeToString,
- node__service__pb2.Tensor.FromString,
+ exo_dot_networking_dot_grpc_dot_node__service__pb2.PromptRequest.SerializeToString,
+ exo_dot_networking_dot_grpc_dot_node__service__pb2.Tensor.FromString,
options,
channel_credentials,
insecure,
@@ -212,8 +212,8 @@ class NodeService(object):
request,
target,
'/node_service.NodeService/SendTensor',
- node__service__pb2.TensorRequest.SerializeToString,
- node__service__pb2.Tensor.FromString,
+ exo_dot_networking_dot_grpc_dot_node__service__pb2.TensorRequest.SerializeToString,
+ exo_dot_networking_dot_grpc_dot_node__service__pb2.Tensor.FromString,
options,
channel_credentials,
insecure,
@@ -239,8 +239,8 @@ class NodeService(object):
request,
target,
'/node_service.NodeService/GetInferenceResult',
- node__service__pb2.GetInferenceResultRequest.SerializeToString,
- node__service__pb2.InferenceResult.FromString,
+ exo_dot_networking_dot_grpc_dot_node__service__pb2.GetInferenceResultRequest.SerializeToString,
+ exo_dot_networking_dot_grpc_dot_node__service__pb2.InferenceResult.FromString,
options,
channel_credentials,
insecure,
@@ -266,8 +266,8 @@ class NodeService(object):
request,
target,
'/node_service.NodeService/CollectTopology',
- node__service__pb2.CollectTopologyRequest.SerializeToString,
- node__service__pb2.Topology.FromString,
+ exo_dot_networking_dot_grpc_dot_node__service__pb2.CollectTopologyRequest.SerializeToString,
+ exo_dot_networking_dot_grpc_dot_node__service__pb2.Topology.FromString,
options,
channel_credentials,
insecure,
@@ -293,8 +293,8 @@ class NodeService(object):
request,
target,
'/node_service.NodeService/SendResult',
- node__service__pb2.SendResultRequest.SerializeToString,
- node__service__pb2.Empty.FromString,
+ exo_dot_networking_dot_grpc_dot_node__service__pb2.SendResultRequest.SerializeToString,
+ exo_dot_networking_dot_grpc_dot_node__service__pb2.Empty.FromString,
options,
channel_credentials,
insecure,
@@ -320,8 +320,8 @@ class NodeService(object):
request,
target,
'/node_service.NodeService/SendOpaqueStatus',
- node__service__pb2.SendOpaqueStatusRequest.SerializeToString,
- node__service__pb2.Empty.FromString,
+ exo_dot_networking_dot_grpc_dot_node__service__pb2.SendOpaqueStatusRequest.SerializeToString,
+ exo_dot_networking_dot_grpc_dot_node__service__pb2.Empty.FromString,
options,
channel_credentials,
insecure,
@@ -347,8 +347,8 @@ class NodeService(object):
request,
target,
'/node_service.NodeService/HealthCheck',
- node__service__pb2.HealthCheckRequest.SerializeToString,
- node__service__pb2.HealthCheckResponse.FromString,
+ exo_dot_networking_dot_grpc_dot_node__service__pb2.HealthCheckRequest.SerializeToString,
+ exo_dot_networking_dot_grpc_dot_node__service__pb2.HealthCheckResponse.FromString,
options,
channel_credentials,
insecure,
diff --git a/exo/orchestration/standard_node.py b/exo/orchestration/standard_node.py
index 17c84a5a..055d3fdb 100644
--- a/exo/orchestration/standard_node.py
+++ b/exo/orchestration/standard_node.py
@@ -111,41 +111,50 @@ class StandardNode(Node):
shard,
result: np.ndarray,
request_id: Optional[str] = None,
+ inference_state: Optional[dict] = None,
):
- if request_id not in self.buffered_token_output:
- self.buffered_token_output[request_id] = ([], False)
-
- if request_id not in self.buffered_logits:
- self.buffered_logits[request_id] = []
-
- self.buffered_logits[request_id] += [i for i in np.reshape(result, (-1, 1, result.shape[-1]))]
-
- if shard.is_last_layer():
- result = await self.inference_engine.sample(result)
+ if shard.model_id != 'stable-diffusion-2-1-base':
+ if request_id not in self.buffered_token_output:
+ self.buffered_token_output[request_id] = ([], False)
+
+ if request_id not in self.buffered_logits:
+ self.buffered_logits[request_id] = []
+
+ self.buffered_logits[request_id] += [i for i in np.reshape(result, (-1, 1, result.shape[-1]))]
+ intermediate_result = self.buffered_token_output[request_id][0]
+
+ if shard.is_last_layer():
+ result = await self.inference_engine.sample(result)
+ else:
+ intermediate_result, inference_state = self.handle_stable_diffusion(inference_state, result)
await self.inference_engine.ensure_shard(shard)
- is_finished = result.size == 1 and result.item() == self.inference_engine.tokenizer.eos_token_id or len(self.buffered_token_output[request_id][0]) >= self.max_generate_tokens
-
- asyncio.create_task(self.broadcast_result(request_id, self.buffered_token_output[request_id][0], is_finished)) # TODO: this is n^2 communication complexity
-
- if result.size == 1: # we got a new token out
- self.buffered_token_output[request_id][0].append(result.item())
- self.trigger_on_token_callbacks(request_id, self.buffered_token_output[request_id][0], is_finished)
+ is_finished = inference_state.get('is_finished', False) if inference_state else (result.size == 1 and result.item() == self.inference_engine.tokenizer.eos_token_id) or len(self.buffered_token_output[request_id][0]) >= self.max_generate_tokens
- if DEBUG >= 2: print(f"[{request_id}] result size: {result.size}, is finished: {is_finished}, buffered tokens: {len(self.buffered_token_output[request_id][0])}")
+ asyncio.create_task(self.broadcast_result(request_id, intermediate_result, is_finished)) # TODO: this is n^2 communication complexity
+ if shard.model_id != 'stable-diffusion-2-1-base':
+ if result.size == 1: # we got a new token out
+ self.buffered_token_output[request_id][0].append(result.item())
+ self.trigger_on_token_callbacks(request_id, self.buffered_token_output[request_id][0], is_finished)
+ else:
+ self.trigger_on_token_callbacks(request_id, intermediate_result, is_finished)
+
+ if DEBUG >= 2: print(f"[{request_id}] result size: {result.size}, is finished: {is_finished}, buffered tokens: {len(intermediate_result)}")
if is_finished:
- self.buffered_token_output[request_id] = (self.buffered_token_output[request_id][0], True)
+ if shard.model_id != 'stable-diffusion-2-1-base':
+ self.buffered_token_output[request_id] = (self.buffered_token_output[request_id][0], True)
else:
- asyncio.create_task(self.forward_to_next_shard(shard, result, request_id))
+ asyncio.create_task(self.forward_to_next_shard(shard, result, request_id, inference_state))
- return np.array(self.buffered_token_output[request_id][0]) if len(self.buffered_token_output[request_id][0]) > 0 else None
+ return np.array(intermediate_result) if len(intermediate_result) > 0 else None
async def process_prompt(
self,
base_shard: Shard,
prompt: str,
request_id: Optional[str] = None,
+ inference_state: Optional[dict] = {},
) -> Optional[np.ndarray]:
shard = self.get_current_shard(base_shard)
asyncio.create_task(
@@ -163,7 +172,7 @@ class StandardNode(Node):
)
)
start_time = time.perf_counter_ns()
- resp = await self._process_prompt(base_shard, prompt, request_id)
+ resp = await self._process_prompt(base_shard, prompt, request_id, inference_state)
end_time = time.perf_counter_ns()
elapsed_time_ns = end_time - start_time
asyncio.create_task(
@@ -184,19 +193,18 @@ class StandardNode(Node):
)
return resp
- async def _process_prompt(self, base_shard: Shard, prompt: str, request_id: Optional[str] = None) -> Optional[np.ndarray]:
+ async def _process_prompt(self, base_shard: Shard, prompt: str, request_id: Optional[str] = None, inference_state: Optional[dict] = None) -> Optional[np.ndarray]:
if request_id is None:
request_id = str(uuid.uuid4())
shard = self.get_current_shard(base_shard)
-
if DEBUG >= 2: print(f"[{request_id}] process prompt: {base_shard=} {shard=} {prompt=}")
if shard.start_layer != 0:
if DEBUG >= 2: print(f"[{request_id}] forwarding to next shard: {base_shard=} {shard=} {prompt=}")
await self.forward_to_next_shard(shard, prompt, request_id)
return None
else:
- result = await self.inference_engine.infer_prompt(request_id, shard, prompt)
- ret = await self.process_result(shard, result, request_id)
+ result,inference_state = await self.inference_engine.infer_prompt(request_id, shard, prompt, inference_state)
+ ret = await self.process_result(shard, result, request_id, inference_state)
return result
async def process_tensor(
@@ -204,6 +212,7 @@ class StandardNode(Node):
base_shard: Shard,
tensor: np.ndarray,
request_id: Optional[str] = None,
+ inference_state: Optional[dict] = None,
) -> Optional[np.ndarray]:
shard = self.get_current_shard(base_shard)
asyncio.create_task(
@@ -222,7 +231,7 @@ class StandardNode(Node):
)
)
start_time = time.perf_counter_ns()
- resp = await self._process_tensor(shard, tensor, request_id)
+ resp = await self._process_tensor(shard, tensor, request_id, inference_state)
end_time = time.perf_counter_ns()
elapsed_time_ns = end_time - start_time
asyncio.create_task(
@@ -247,6 +256,7 @@ class StandardNode(Node):
base_shard: Shard,
tensor: np.ndarray,
request_id: Optional[str] = None,
+ inference_state: Optional[dict] = None,
) -> Optional[np.ndarray]:
if request_id is None:
request_id = str(uuid.uuid4())
@@ -254,8 +264,8 @@ class StandardNode(Node):
if DEBUG >= 1: print(f"[{request_id}] process_tensor: {tensor.size=} {tensor.shape=}")
try:
- result = await self.inference_engine.infer_tensor(request_id, shard, tensor)
- ret = await self.process_result(shard, result, request_id)
+ result, inference_state = await self.inference_engine.infer_tensor(request_id, shard, tensor, inference_state)
+ ret = await self.process_result(shard, result, request_id, inference_state)
return ret
except Exception as e:
print(f"Error processing tensor for shard {shard}: {e}")
@@ -267,6 +277,7 @@ class StandardNode(Node):
base_shard: Shard,
tensor_or_prompt: Union[np.ndarray, str],
request_id: str,
+ inference_state: Optional[dict] = None,
) -> None:
if not self.partitioning_strategy:
if DEBUG >= 1: print("No partitioning strategy found. Skipping forward.")
@@ -281,16 +292,16 @@ class StandardNode(Node):
is_tensor = isinstance(tensor_or_prompt, np.ndarray)
if target_id == self.id:
if is_tensor:
- await self.process_tensor(next_shard, tensor_or_prompt, request_id)
+ await self.process_tensor(next_shard, tensor_or_prompt, request_id, inference_state)
else:
- await self.process_prompt(next_shard, tensor_or_prompt, request_id)
+ await self.process_prompt(next_shard, tensor_or_prompt, request_id, inference_state)
else:
target_peer = next((p for p in self.peers if p.id() == target_id), None)
if not target_peer:
raise ValueError(f"Peer for {next_partition_index} not found")
if is_tensor:
if DEBUG >= 1: print(f"Sending tensor to {target_peer.id()}: {tensor_or_prompt}")
- await target_peer.send_tensor(next_shard, tensor_or_prompt, request_id=request_id)
+ await target_peer.send_tensor(next_shard, tensor_or_prompt, inference_state, request_id=request_id)
else:
await target_peer.send_prompt(next_shard, tensor_or_prompt, request_id=request_id)
@@ -461,3 +472,12 @@ class StandardNode(Node):
@property
def current_topology(self) -> Topology:
return self.topology
+
+ def handle_stable_diffusion(self, inference_state, result):
+ if inference_state['is_step_finished']:
+ inference_state['step']+=1
+ progress = [inference_state['step'],inference_state['total_steps']]
+ intermediate_result = progress
+ if progress[0] == progress[1]:
+ intermediate_result = result
+ return intermediate_result, inference_state
diff --git a/exo/tinychat/index.html b/exo/tinychat/index.html
index 44fbb99b..898baef9 100644
--- a/exo/tinychat/index.html
+++ b/exo/tinychat/index.html
@@ -115,7 +115,15 @@
const div = document.createElement('div');
div.className = `message message-role-${role}`;
try {
- div.innerHTML = DOMPurify.sanitize(marked.parse(content));
+ if (content.includes('![Generated Image]')) {
+ const imageUrl = content.match(/\((.*?)\)/)[1];
+ const img = document.createElement('img');
+ img.src = imageUrl;
+ img.alt = 'Generated Image';
+ div.appendChild(img);
+ } else {
+ div.innerHTML = DOMPurify.sanitize(marked.parse(content));
+ }
} catch (e) {
console.log(content);
console.error(e);
diff --git a/exo/tinychat/index.js b/exo/tinychat/index.js
index f13bd93b..5e1fea9c 100644
--- a/exo/tinychat/index.js
+++ b/exo/tinychat/index.js
@@ -229,53 +229,105 @@ document.addEventListener("alpine:init", () => {
};
}
});
- const containsImage = apiMessages.some(msg => Array.isArray(msg.content) && msg.content.some(item => item.type === 'image_url'));
- if (containsImage) {
- // Map all messages with string content to object with type text
- apiMessages = apiMessages.map(msg => {
- if (typeof msg.content === 'string') {
- return {
- ...msg,
- content: [
- {
- type: "text",
- text: msg.content
- }
- ]
- };
- }
- return msg;
+
+ if (this.cstate.selectedModel === "stable-diffusion-2-1-base") {
+ // Send a request to the image generation endpoint
+ console.log(apiMessages[apiMessages.length - 1].content)
+ console.log(this.cstate.selectedModel)
+ console.log(this.endpoint)
+ const response = await fetch(`${this.endpoint}/image/generations`, {
+ method: "POST",
+ headers: {
+ "Content-Type": "application/json",
+ },
+ body: JSON.stringify({
+ "model": 'stable-diffusion-2-1-base',
+ "prompt": apiMessages[apiMessages.length - 1].content,
+ }),
});
+
+ if (!response.ok) {
+ throw new Error("Failed to fetch");
+ }
+ const reader = response.body.getReader();
+ let done = false;
+ let gottenFirstChunk = false;
+
+ while (!done) {
+ const { value, done: readerDone } = await reader.read();
+ done = readerDone;
+ const decoder = new TextDecoder();
+
+ if (value) {
+ // Assume non-binary data (text) comes first
+ const chunk = decoder.decode(value, { stream: true });
+ const parsed = JSON.parse(chunk);
+ console.log(parsed)
+
+ if (parsed.progress) {
+ if (!gottenFirstChunk) {
+ this.cstate.messages.push({ role: "assistant", content: "" });
+ gottenFirstChunk = true;
+ }
+ this.cstate.messages[this.cstate.messages.length - 1].content = parsed.progress;
+ }
+ else if (parsed.images) {
+ const imageUrl = parsed.images[0].url;
+ console.log(imageUrl)
+ this.cstate.messages[this.cstate.messages.length - 1].content = ``;
+ }
+ }
+ }
}
-
-
- // start receiving server sent events
- let gottenFirstChunk = false;
- for await (
- const chunk of this.openaiChatCompletion(this.cstate.selectedModel, apiMessages)
- ) {
- if (!gottenFirstChunk) {
- this.cstate.messages.push({ role: "assistant", content: "" });
- gottenFirstChunk = true;
+
+ else{
+ const containsImage = apiMessages.some(msg => Array.isArray(msg.content) && msg.content.some(item => item.type === 'image_url'));
+ if (containsImage) {
+ // Map all messages with string content to object with type text
+ apiMessages = apiMessages.map(msg => {
+ if (typeof msg.content === 'string') {
+ return {
+ ...msg,
+ content: [
+ {
+ type: "text",
+ text: msg.content
+ }
+ ]
+ };
+ }
+ return msg;
+ });
}
- // add chunk to the last message
- this.cstate.messages[this.cstate.messages.length - 1].content += chunk;
+ console.log(apiMessages)
+ //start receiving server sent events
+ let gottenFirstChunk = false;
+ for await (
+ const chunk of this.openaiChatCompletion(this.cstate.selectedModel, apiMessages)
+ ) {
+ if (!gottenFirstChunk) {
+ this.cstate.messages.push({ role: "assistant", content: "" });
+ gottenFirstChunk = true;
+ }
- // calculate performance tracking
- tokens += 1;
- this.total_tokens += 1;
- if (start_time === 0) {
- start_time = Date.now();
- this.time_till_first = start_time - prefill_start;
- } else {
- const diff = Date.now() - start_time;
- if (diff > 0) {
- this.tokens_per_second = tokens / (diff / 1000);
+ // add chunk to the last message
+ this.cstate.messages[this.cstate.messages.length - 1].content += chunk;
+
+ // calculate performance tracking
+ tokens += 1;
+ this.total_tokens += 1;
+ if (start_time === 0) {
+ start_time = Date.now();
+ this.time_till_first = start_time - prefill_start;
+ } else {
+ const diff = Date.now() - start_time;
+ if (diff > 0) {
+ this.tokens_per_second = tokens / (diff / 1000);
+ }
}
}
}
-
// Clean the cstate before adding it to histories
const cleanedCstate = JSON.parse(JSON.stringify(this.cstate));
cleanedCstate.messages = cleanedCstate.messages.map(msg => {
← f337e781 pr fixes
·
back to Exo
·
build error fix 44118252 →