This commit is contained in:
Evan
2026-03-24 11:54:04 +00:00
committed by rltakashige
parent 6355f3d8fb
commit a9c7b1c68a
70 changed files with 235 additions and 221 deletions
@@ -0,0 +1,475 @@
import itertools
import time
from collections import deque
from collections.abc import Callable, Generator, Iterable
from dataclasses import dataclass, field
from typing import TYPE_CHECKING
import mlx.core as mx
from exo_core.constants import EXO_MAX_CONCURRENT_REQUESTS
from exo_core.types.chunks import ErrorChunk, PrefillProgressChunk
from exo_core.types.common import CommandId, ModelId
from exo_core.types.runner_response import GenerationResponse, ToolCallResponse
from exo_core.types.tasks import CANCEL_ALL_TASKS, TaskId, TextGeneration
from exo_core.types.text_generation import TextGenerationTaskParams
from exo_core.utils.channels import MpReceiver, MpSender
from mlx_lm.tokenizer_utils import TokenizerWrapper
from mlx_engine.cache import KVPrefixCache
from mlx_engine.generator.batch_generate import ExoBatchGenerator
if TYPE_CHECKING:
from vllm_engine.vllm_generator import VllmBatchEngine
from exo_core.engine import Cancelled, Engine, Finished
from exo_core.tokenizers.model_output_parsers import apply_all_parsers
from exo_core.tokenizers.tool_parsers import ToolParser
from loguru import logger
from mlx_engine.generator.generate import (
PrefillCancelled,
)
from mlx_engine.utils_mlx import (
apply_chat_template,
mx_all_gather_tasks,
mx_any,
)
class GeneratorQueue[T]:
def __init__(self):
self._q = deque[T]()
def push(self, t: T):
self._q.append(t)
def gen(self) -> Generator[T | None]:
while True:
if len(self._q) == 0:
yield None
else:
yield self._q.popleft()
EXO_RUNNER_MUST_FAIL = "EXO RUNNER MUST FAIL"
EXO_RUNNER_MUST_OOM = "EXO RUNNER MUST OOM"
EXO_RUNNER_MUST_TIMEOUT = "EXO RUNNER MUST TIMEOUT"
def _check_for_debug_prompts(task_params: TextGenerationTaskParams) -> None:
"""Check for debug prompt triggers in the input."""
from mlx_engine.utils_mlx import mlx_force_oom
if len(task_params.input) == 0:
return
prompt = task_params.input[0].content
if not prompt:
return
if EXO_RUNNER_MUST_FAIL in prompt:
raise Exception("Artificial runner exception - for testing purposes only.")
if EXO_RUNNER_MUST_OOM in prompt:
mlx_force_oom()
if EXO_RUNNER_MUST_TIMEOUT in prompt:
time.sleep(100)
@dataclass(eq=False)
class SequentialGenerator(
Engine[TextGeneration, GenerationResponse | ToolCallResponse]
):
tokenizer: TokenizerWrapper
group: mx.distributed.Group | None
kv_prefix_cache: KVPrefixCache | None
tool_parser: ToolParser | None
model_id: ModelId
device_rank: int
cancel_receiver: MpReceiver[TaskId]
event_sender: MpSender[tuple[CommandId, ErrorChunk | PrefillProgressChunk]]
_generate_fn: Callable[..., Generator[GenerationResponse]]
_warmup_fn: Callable[[], int]
check_for_cancel_every: int = 50
_cancelled_tasks: set[TaskId] = field(default_factory=set, init=False)
_maybe_queue: list[TextGeneration] = field(default_factory=list, init=False)
_maybe_cancel: list[TextGeneration] = field(default_factory=list, init=False)
_all_tasks: dict[TaskId, TextGeneration] = field(default_factory=dict, init=False)
_queue: deque[TextGeneration] = field(default_factory=deque, init=False)
_active: (
tuple[
TextGeneration,
# mlx generator that does work
Generator[GenerationResponse],
# queue that the 1st generator should push to and 3rd generator should pull from
GeneratorQueue[GenerationResponse],
# generator to get parsed outputs
Generator[GenerationResponse | ToolCallResponse | None],
]
| None
) = field(default=None, init=False)
def warmup(self) -> None:
self.check_for_cancel_every = self._warmup_fn()
def submit(
self,
task: TextGeneration,
) -> None:
self._cancelled_tasks.discard(CANCEL_ALL_TASKS)
self._all_tasks[task.task_id] = task
self._maybe_queue.append(task)
def agree_on_tasks(self) -> None:
"""Agree between all ranks about the task ordering (some may have received in different order or not at all)."""
agreed, different = mx_all_gather_tasks(self._maybe_queue, self.group)
self._queue.extend(task for task in self._maybe_queue if task in agreed)
self._maybe_queue = [task for task in self._maybe_queue if task in different]
def agree_on_cancellations(self) -> None:
"""Agree between all ranks about which tasks to cancel."""
has_cancel_all = False
for task_id in self.cancel_receiver.collect():
if task_id == CANCEL_ALL_TASKS:
has_cancel_all = True
continue
if task_id in self._all_tasks:
self._maybe_cancel.append(self._all_tasks[task_id])
if mx_any(has_cancel_all, self.group):
self._cancelled_tasks.add(CANCEL_ALL_TASKS)
agreed, different = mx_all_gather_tasks(self._maybe_cancel, self.group)
self._cancelled_tasks.update(task.task_id for task in agreed)
self._maybe_cancel = list(different)
def step(
self,
) -> Iterable[
tuple[TaskId, GenerationResponse | ToolCallResponse | Cancelled | Finished]
]:
if self._active is None:
self.agree_on_tasks()
if self._queue:
self._start_next()
else:
return map(lambda task: (task, Cancelled()), self._cancelled_tasks)
assert self._active is not None
task, mlx_gen, queue, output_generator = self._active
response = None
try:
queue.push(next(mlx_gen))
response = next(output_generator)
except (StopIteration, PrefillCancelled):
response = Finished()
self._active = None
if self._queue:
self._start_next()
except Exception as e:
self._send_error(task, e)
self._active = None
raise
return itertools.chain(
[] if response is None else [(task.task_id, response)],
map(lambda task: (task, Cancelled()), self._cancelled_tasks),
)
def _start_next(self) -> None:
task = self._queue.popleft()
try:
mlx_gen = self._build_generator(task)
except Exception as e:
self._send_error(task, e)
raise
queue = GeneratorQueue[GenerationResponse]()
if task.task_params.bench:
output_generator = queue.gen()
else:
output_generator = apply_all_parsers(
queue.gen(),
apply_chat_template(self.tokenizer, task.task_params),
self.tool_parser,
self.tokenizer,
self.model_id,
task.task_params.tools,
)
self._active = (task, mlx_gen, queue, output_generator)
def _send_error(self, task: TextGeneration, e: Exception) -> None:
if self.device_rank == 0:
# TODO: sync channels?
self.event_sender.send(
(
task.command_id,
ErrorChunk(
model=self.model_id,
finish_reason="error",
error_message=str(e),
),
)
)
def _build_generator(self, task: TextGeneration) -> Generator[GenerationResponse]:
_check_for_debug_prompts(task.task_params)
prompt = apply_chat_template(self.tokenizer, task.task_params)
def on_prefill_progress(processed: int, total: int) -> None:
if self.device_rank == 0:
self.event_sender.send(
(
task.command_id,
PrefillProgressChunk(
model=self.model_id,
processed_tokens=processed,
total_tokens=total,
),
)
)
def distributed_prompt_progress_callback() -> None:
self.agree_on_cancellations()
if self.should_cancel(task.task_id):
raise PrefillCancelled()
self.agree_on_tasks()
tokens_since_cancel_check = self.check_for_cancel_every
def on_generation_token() -> None:
nonlocal tokens_since_cancel_check
tokens_since_cancel_check += 1
if tokens_since_cancel_check >= self.check_for_cancel_every:
tokens_since_cancel_check = 0
self.agree_on_cancellations()
if self.should_cancel(task.task_id):
raise PrefillCancelled()
self.agree_on_tasks()
return self._generate_fn(
task=task.task_params,
prompt=prompt,
kv_prefix_cache=self.kv_prefix_cache,
on_prefill_progress=on_prefill_progress,
distributed_prompt_progress_callback=distributed_prompt_progress_callback,
on_generation_token=on_generation_token,
group=self.group,
)
def close(self) -> None:
del self.tokenizer, self.group
@dataclass(eq=False)
class BatchGenerator(Engine[TextGeneration, GenerationResponse | ToolCallResponse]):
tokenizer: TokenizerWrapper
group: mx.distributed.Group | None
kv_prefix_cache: KVPrefixCache | None
tool_parser: ToolParser | None
model_id: ModelId
device_rank: int
cancel_receiver: MpReceiver[TaskId]
event_sender: MpSender[tuple[CommandId, ErrorChunk | PrefillProgressChunk]]
_gen: "ExoBatchGenerator | VllmBatchEngine"
max_concurrent_requests: int = EXO_MAX_CONCURRENT_REQUESTS
check_for_cancel_every: int = 50
_cancelled_tasks: set[TaskId] = field(default_factory=set, init=False)
_maybe_queue: list[TextGeneration] = field(default_factory=list, init=False)
_maybe_cancel: list[TextGeneration] = field(default_factory=list, init=False)
_all_tasks: dict[TaskId, TextGeneration] = field(default_factory=dict, init=False)
_queue: deque[TextGeneration] = field(default_factory=deque, init=False)
_active_tasks: dict[
TaskId,
tuple[
TextGeneration,
GeneratorQueue[GenerationResponse],
Generator[GenerationResponse | ToolCallResponse | None],
],
] = field(default_factory=dict, init=False)
def warmup(self) -> None:
self.check_for_cancel_every = self._gen.warmup()
def submit(
self,
task: TextGeneration,
) -> None:
self._cancelled_tasks.discard(CANCEL_ALL_TASKS)
self._all_tasks[task.task_id] = task
self._maybe_queue.append(task)
def agree_on_tasks(self) -> None:
"""Agree between all ranks about the task ordering (some may have received in different order or not at all)."""
agreed, different = mx_all_gather_tasks(self._maybe_queue, self.group)
self._queue.extend(task for task in self._maybe_queue if task in agreed)
self._maybe_queue = [task for task in self._maybe_queue if task in different]
def agree_on_cancellations(self) -> None:
"""Agree between all ranks about which tasks to cancel."""
has_cancel_all = False
for task_id in self.cancel_receiver.collect():
if task_id == CANCEL_ALL_TASKS:
has_cancel_all = True
continue
if task_id in self._all_tasks:
self._maybe_cancel.append(self._all_tasks[task_id])
if mx_any(has_cancel_all, self.group):
self._cancelled_tasks.add(CANCEL_ALL_TASKS)
agreed, different = mx_all_gather_tasks(self._maybe_cancel, self.group)
self._cancelled_tasks.update(task.task_id for task in agreed)
self._maybe_cancel = list(different)
def step(
self,
) -> Iterable[
tuple[TaskId, GenerationResponse | ToolCallResponse | Cancelled | Finished]
]:
if not self._queue:
self.agree_on_tasks()
# Submit any queued tasks to the engine
while self._queue and len(self._active_tasks) < self.max_concurrent_requests:
task = self._queue.popleft()
try:
task_id = self._start_task(task)
except PrefillCancelled:
continue
except Exception as e:
self._send_error(task, e)
raise
queue = GeneratorQueue[GenerationResponse]()
if task.task_params.bench:
output_generator = queue.gen()
else:
output_generator = apply_all_parsers(
queue.gen(),
apply_chat_template(self.tokenizer, task.task_params),
self.tool_parser,
self.tokenizer,
self.model_id,
task.task_params.tools,
)
self._active_tasks[task_id] = (task, queue, output_generator)
if not self._gen.has_work:
return self._apply_cancellations()
results = self._gen.step()
output: list[
tuple[TaskId, GenerationResponse | ToolCallResponse | Cancelled | Finished]
] = []
for uid, response in results:
if uid not in self._active_tasks:
# should we error here?
logger.warning(f"{uid=} not found in active tasks")
continue
task, queue, output_generator = self._active_tasks[uid]
queue.push(response)
parsed = next(output_generator)
if parsed is not None:
output.append((task.task_id, parsed))
if response.finish_reason is not None:
output.append((task.task_id, Finished()))
del self._active_tasks[uid]
return itertools.chain(output, self._apply_cancellations())
def _apply_cancellations(
self,
) -> list[tuple[TaskId, Cancelled]]:
if not self._cancelled_tasks:
return []
cancel_all = CANCEL_ALL_TASKS in self._cancelled_tasks
ids_to_cancel: list[TaskId] = []
results: list[tuple[TaskId, Cancelled]] = []
for tid, (task, _, _) in list(self._active_tasks.items()):
if task.task_id in self._cancelled_tasks or cancel_all:
ids_to_cancel.append(tid)
results.append((task.task_id, Cancelled()))
del self._active_tasks[tid]
if ids_to_cancel:
self._gen.cancel(ids_to_cancel)
already_cancelled = {tid for tid, _ in results}
for tid in self._cancelled_tasks:
if tid != CANCEL_ALL_TASKS and tid not in already_cancelled:
results.append((tid, Cancelled()))
self._cancelled_tasks.clear()
return results
def _send_error(self, task: TextGeneration, e: Exception) -> None:
if self.device_rank == 0:
self.event_sender.send(
(
task.command_id,
ErrorChunk(
model=self.model_id,
finish_reason="error",
error_message=str(e),
),
)
)
def _start_task(self, task: TextGeneration) -> TaskId:
_check_for_debug_prompts(task.task_params)
prompt = apply_chat_template(self.tokenizer, task.task_params)
def on_prefill_progress(processed: int, total: int) -> None:
if self.device_rank == 0:
self.event_sender.send(
(
task.command_id,
PrefillProgressChunk(
model=self.model_id,
processed_tokens=processed,
total_tokens=total,
),
)
)
def distributed_prompt_progress_callback() -> None:
self.agree_on_cancellations()
if self.should_cancel(task.task_id):
raise PrefillCancelled()
self.agree_on_tasks()
tokens_since_cancel_check = self.check_for_cancel_every
def on_generation_token() -> None:
nonlocal tokens_since_cancel_check
tokens_since_cancel_check += 1
if tokens_since_cancel_check >= self.check_for_cancel_every:
tokens_since_cancel_check = 0
self.agree_on_cancellations()
if self.should_cancel(task.task_id):
self._cancelled_tasks.add(task.task_id)
self.agree_on_tasks()
return self._gen.submit(
task_id=task.task_id,
task_params=task.task_params,
prompt=prompt,
on_prefill_progress=on_prefill_progress,
distributed_prompt_progress_callback=distributed_prompt_progress_callback,
on_generation_token=on_generation_token,
)
def close(self) -> None:
self._gen.close()
del self.tokenizer, self.group
+23 -15
View File
@@ -1,27 +1,36 @@
import contextlib
import os
from dataclasses import dataclass
from typing import Self, Callable
from exo_core.engine import EngineBuilder, Engine
from exo_core.types.common import ModelId
from typing import Callable, Self
import mlx.core as mx
from exo_core.engine import EngineBuilder
from exo_core.tokenizers.tool_parsers import make_mlx_parser
from exo_core.types.chunks import ErrorChunk, PrefillProgressChunk
from exo_core.types.common import CommandId, ModelId
from exo_core.types.instances import BoundInstance
from exo_core.types.tasks import TextGeneration
from exo_core.types.runner_response import GenerationResponse
from mlx_engine.utils_mlx import initialize_mlx, load_mlx_items
from mlx_engine.types import Model
from exo_core.types.runner_response import GenerationResponse, ToolCallResponse
from exo_core.types.tasks import TaskId, TextGeneration
from exo_core.utils.channels import MpReceiver, MpSender
from loguru import logger
from mlx_lm.tokenizer_utils import TokenizerWrapper
from mlx_engine.batch_generator import BatchGenerator, SequentialGenerator
from mlx_engine.cache import KVPrefixCache
from mlx_engine.generator.batch_generate import ExoBatchGenerator
from mlx_engine.generator.generate import (
mlx_generate,
warmup_inference,
)
from exo_core.utils.tool_parsers import make_mlx_parser
from mlx_engine.types import Model
from mlx_engine.utils_mlx import initialize_mlx, load_mlx_items
@dataclass
class MlxBuilder(EngineBuilder[BoundInstance, TextGeneration, GenerationResponse]):
import mlx.core as mx
from mlx_lm.tokenizer_utils import TokenizerWrapper
class MlxBuilder(EngineBuilder[BoundInstance, TextGeneration, GenerationResponse | ToolCallResponse]):
model_id: ModelId
bound_instance: BoundInstance
event_sender: MpSender[Event]
event_sender: MpSender[tuple[CommandId, ErrorChunk | PrefillProgressChunk]]
cancel_receiver: MpReceiver[TaskId]
inference_model: Model | None = None
tokenizer: TokenizerWrapper | None = None
@@ -31,7 +40,7 @@ class MlxBuilder(EngineBuilder[BoundInstance, TextGeneration, GenerationResponse
def create(
cls,
bound_instance: BoundInstance,
event_sender: MpSender[Event],
event_sender: MpSender[tuple[CommandId, ErrorChunk | PrefillProgressChunk]],
cancel_receiver: MpReceiver[TaskId],
) -> Self:
return cls(
@@ -105,7 +114,6 @@ class MlxBuilder(EngineBuilder[BoundInstance, TextGeneration, GenerationResponse
_generate_fn=generate_fn,
_warmup_fn=warmup_fn,
)
from exo.worker.runner.llm_inference.batch_generator import ExoBatchGenerator
logger.info("using BatchGenerator")
gen = ExoBatchGenerator(
+3 -4
View File
@@ -25,7 +25,7 @@ if TYPE_CHECKING:
# Fraction of device memory above which LRU eviction kicks in.
# Smaller machines need more aggressive eviction.
def _default_memory_threshold() -> float:
total_gb = Memory.from_bytes(psutil.virtual_memory().total).in_gb
total_gb = Memory.from_bytes(psutil.virtual_memory().total).in_gb # pyright: ignore[reportAny]
if total_gb >= 128:
return 0.85
if total_gb >= 64:
@@ -220,7 +220,6 @@ class KVPrefixCache:
def lookup(
self, prompt_token_ids: list[int]
) -> tuple["TorchKVCache | None", int, int | None]:
from exo.worker.engines.vllm.kv_cache import TorchKVCache
prompt_mx = mx.array(prompt_token_ids)
max_length = len(prompt_token_ids)
@@ -352,14 +351,14 @@ def get_prefix_length(prompt: mx.array, cached_prompt: mx.array) -> int:
def get_available_memory() -> Memory:
mem: int = psutil.virtual_memory().available
mem: int = psutil.virtual_memory().available # pyright: ignore[reportAny]
return Memory.from_bytes(mem)
def get_memory_used_percentage() -> float:
mem = psutil.virtual_memory()
# percent is 0-100
return float(mem.percent / 100)
return float(mem.percent / 100) # pyright: ignore[reportAny]
def make_kv_cache(
@@ -2,9 +2,8 @@ import json
import re
from typing import Any
from mlx_lm.chat_templates import deepseek_v32
from exo_core.types.runner_response import ToolCallItem
from mlx_lm.chat_templates import deepseek_v32
BOS_TOKEN: str = deepseek_v32.bos_token
EOS_TOKEN: str = deepseek_v32.eos_token
@@ -4,7 +4,15 @@ from typing import Callable, cast
import mlx.core as mx
from exo_core.types.common import ModelId
from exo_core.types.runner_response import GenerationResponse
from exo_core.types.runner_response import (
CompletionTokensDetails,
FinishReason,
GenerationResponse,
GenerationStats,
PromptTokensDetails,
TopLogprobItem,
Usage,
)
from exo_core.types.tasks import TaskId
from exo_core.types.text_generation import TextGenerationTaskParams
from exo_core.utils.memory import Memory
@@ -16,14 +24,6 @@ from mlx_lm.models.cache import RotatingKVCache
from mlx_lm.sample_utils import make_logits_processors, make_sampler
from mlx_lm.tokenizer_utils import StreamingDetokenizer, TokenizerWrapper
from exo.api.types import (
CompletionTokensDetails,
FinishReason,
GenerationStats,
PromptTokensDetails,
TopLogprobItem,
Usage,
)
from mlx_engine.cache import (
CacheSnapshot,
KVPrefixCache,
@@ -7,7 +7,13 @@ from typing import Callable, Generator, cast, get_args
import mlx.core as mx
from exo_core.types.common import ModelId
from exo_core.types.runner_response import (
CompletionTokensDetails,
FinishReason,
GenerationResponse,
GenerationStats,
PromptTokensDetails,
TopLogprobItem,
Usage,
)
from exo_core.types.text_generation import InputMessage, TextGenerationTaskParams
from exo_core.utils.memory import Memory
@@ -20,14 +26,6 @@ from mlx_lm.models.cache import ArraysCache, RotatingKVCache
from mlx_lm.sample_utils import make_logits_processors, make_sampler
from mlx_lm.tokenizer_utils import TokenizerWrapper
from exo.api.types import (
CompletionTokensDetails,
FinishReason,
GenerationStats,
PromptTokensDetails,
TopLogprobItem,
Usage,
)
from mlx_engine.auto_parallel import (
PipelineFirstLayer,
PipelineLastLayer,