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
+5 -1
View File
@@ -5,7 +5,11 @@ description = "Add your description here"
readme = "README.md"
authors = [{ name = "Evan", email = "[email protected]" }]
requires-python = ">=3.13"
dependencies = ["pydantic>=2.13.0b2"]
dependencies = [
"mlx-lm", # TODO: depend on transformers or other
"openai-harmony", # inherit from workspace
"pydantic", # inherit from workspace
]
[build-system]
requires = ["uv_build>=0.9.24,<0.10.0"]
+10 -9
View File
@@ -1,14 +1,15 @@
from abc import ABC, abstractmethod
from collections.abc import Callable, Iterable
from typing import Self
from exo_core.types.tasks import TaskId
class Cancelled: pass
class Cancelled:
pass
class Finished: pass
class Finished:
pass
CANCEL_ALL_TASKS = TaskId("CANCEL_TALL_TASKS")
@@ -45,12 +46,12 @@ class Engine[TaskType, ResponseType](ABC):
class EngineBuilder[SetupType, TaskType, ResponseType](ABC):
@classmethod
@abstractmethod
def create(
cls,
bound_instance: SetupType,
) -> Self: ...
# @classmethod
# @abstractmethod
# def create(
# cls,
# bound_instance: SetupType,
# ) -> Self: ...
@abstractmethod
def connect(self) -> None: ...
@@ -0,0 +1,382 @@
from collections.abc import Generator
from functools import cache
from typing import Any
from loguru import logger
from mlx_engine.utils_mlx import (
detect_thinking_prompt_suffix,
)
from mlx_lm.tokenizer_utils import TokenizerWrapper
from openai_harmony import (
HarmonyEncodingName,
HarmonyError,
Role,
StreamableParser,
load_harmony_encoding,
)
from exo_core.tokenizers.tool_parsers import ToolParser
from exo_core.types.common import ModelId
from exo_core.types.runner_response import (
GenerationResponse,
ToolCallItem,
ToolCallResponse,
)
@cache
def get_gpt_oss_encoding():
encoding = load_harmony_encoding(HarmonyEncodingName.HARMONY_GPT_OSS)
return encoding
def apply_all_parsers(
receiver: Generator[GenerationResponse | None],
prompt: str,
tool_parser: ToolParser | None,
tokenizer: TokenizerWrapper,
model_id: ModelId,
tools: list[dict[str, Any]] | None,
) -> Generator[GenerationResponse | ToolCallResponse | None]:
gen = receiver
if tokenizer.has_thinking:
gen = parse_thinking_models(
gen,
tokenizer.think_start,
tokenizer.think_end,
starts_in_thinking=detect_thinking_prompt_suffix(prompt, tokenizer),
)
lower = model_id.normalize().lower()
if "gpt-oss" in lower or "gpt_oss" in lower:
gen = parse_gpt_oss(gen)
elif "deepseek" in lower:
gen = parse_deepseek_v32(gen)
elif tool_parser:
gen = parse_tool_calls(gen, tool_parser, tools)
return gen
_GPT_OSS_CHANNEL_TOKEN = 200005
_GPT_OSS_MESSAGE_TOKEN = 200008
def parse_gpt_oss(
responses: Generator[GenerationResponse | None],
) -> Generator[GenerationResponse | ToolCallResponse | None]:
encoding = get_gpt_oss_encoding()
stream = StreamableParser(encoding, role=Role.ASSISTANT)
thinking = False
current_tool_name: str | None = None
tool_arg_parts: list[str] = []
for response in responses:
if response is None:
yield None
continue
token_id = response.token
try:
stream.process(token_id)
except HarmonyError as e:
logger.error(
f"HarmonyError on token_id={response.token} text={response.text!r}: {e}"
)
return
delta = stream.last_content_delta
ch = stream.current_channel
recipient = stream.current_recipient
effective_recipient = (
recipient
if (recipient is not None and recipient.startswith("functions."))
else None
)
if effective_recipient != current_tool_name:
if current_tool_name is not None:
tool_name = current_tool_name.removeprefix("functions.")
logger.info(f"parse_gpt_oss yielding tool call: name={tool_name!r}")
yield ToolCallResponse(
tool_calls=[
ToolCallItem(
name=tool_name,
arguments="".join(tool_arg_parts).strip(),
)
],
usage=response.usage,
)
tool_arg_parts = []
current_tool_name = effective_recipient
if current_tool_name is not None:
if delta:
tool_arg_parts.append(delta)
if response.finish_reason is not None:
yield response.model_copy(update={"text": "".join(tool_arg_parts)})
tool_arg_parts = []
continue
is_suppressed = ch == "analysis" or (
recipient is not None and recipient.startswith("!")
)
if is_suppressed and not thinking:
thinking = True
if not is_suppressed and thinking:
thinking = False
if delta:
yield response.model_copy(update={"text": delta, "is_thinking": thinking})
if response.finish_reason is not None:
yield response.model_copy(update={"text": ""})
def parse_deepseek_v32(
responses: Generator[GenerationResponse | None],
) -> Generator[GenerationResponse | ToolCallResponse | None]:
"""Parse DeepSeek V3.2 DSML tool calls from the generation stream.
Uses accumulated-text matching (not per-token marker checks) because
DSML markers like <DSMLfunction_calls> may span multiple tokens.
Also handles <think>...</think> blocks for thinking mode.
"""
from mlx_engine.dsml_encoding import (
THINKING_END,
THINKING_START,
TOOL_CALLS_END,
TOOL_CALLS_START,
parse_dsml_output,
)
accumulated = ""
in_tool_call = False
thinking = False
# Tokens buffered while we detect the start of a DSML block
pending_buffer: list[GenerationResponse] = []
# Text accumulated during a tool call block
tool_call_text = ""
for response in responses:
if response is None:
yield None
continue
# ── Handle thinking tags ──
if not thinking and THINKING_START in response.text:
thinking = True
# Yield any text before the <think> tag
before = response.text[: response.text.index(THINKING_START)]
if before:
yield response.model_copy(update={"text": before})
continue
if thinking and THINKING_END in response.text:
thinking = False
# Yield any text after the </think> tag
after = response.text[
response.text.index(THINKING_END) + len(THINKING_END) :
]
if after:
yield response.model_copy(update={"text": after, "is_thinking": False})
continue
if thinking:
yield response.model_copy(update={"is_thinking": True})
continue
# ── Handle tool call accumulation ──
if in_tool_call:
tool_call_text += response.text
if TOOL_CALLS_END in tool_call_text:
# Parse the accumulated DSML block
parsed = parse_dsml_output(tool_call_text)
if parsed is not None:
logger.info(f"parsed DSML tool calls: {parsed}")
yield ToolCallResponse(
tool_calls=parsed,
usage=response.usage,
stats=response.stats,
)
else:
logger.warning(
f"DSML tool call parsing failed for: {tool_call_text}"
)
yield response.model_copy(update={"text": tool_call_text})
in_tool_call = False
tool_call_text = ""
continue
# EOS reached before end marker — yield buffered text as-is
if response.finish_reason is not None:
logger.info("DSML tool call parsing interrupted by EOS")
yield response.model_copy(update={"text": tool_call_text})
in_tool_call = False
tool_call_text = ""
continue
# ── Detect start of tool call block ──
accumulated += response.text
if TOOL_CALLS_START in accumulated:
# The start marker might be split across pending_buffer + current token
start_idx = accumulated.index(TOOL_CALLS_START)
# Yield any pending tokens that are purely before the marker
pre_text = accumulated[:start_idx]
if pre_text:
# Flush pending buffer tokens that contributed text before the marker
for buf_resp in pending_buffer:
if pre_text:
chunk = buf_resp.text
if len(chunk) <= len(pre_text):
yield buf_resp
pre_text = pre_text[len(chunk) :]
else:
yield buf_resp.model_copy(update={"text": pre_text})
pre_text = ""
pending_buffer = []
tool_call_text = accumulated[start_idx:]
accumulated = ""
# Check if the end marker is already present (entire tool call in one token)
if TOOL_CALLS_END in tool_call_text:
parsed = parse_dsml_output(tool_call_text)
if parsed is not None:
logger.info(f"parsed DSML tool calls: {parsed}")
yield ToolCallResponse(
tool_calls=parsed,
usage=response.usage,
stats=response.stats,
)
else:
logger.warning(
f"DSML tool call parsing failed for: {tool_call_text}"
)
yield response.model_copy(update={"text": tool_call_text})
tool_call_text = ""
else:
in_tool_call = True
continue
# Check if accumulated text might be the start of a DSML marker
# Buffer tokens if we see a partial match at the end
if _could_be_dsml_prefix(accumulated):
pending_buffer.append(response)
continue
# No partial match — flush all pending tokens and the current one
for buf_resp in pending_buffer:
yield buf_resp
pending_buffer = []
accumulated = ""
yield response
# Flush any remaining pending buffer at generator end
for buf_resp in pending_buffer:
yield buf_resp
def _could_be_dsml_prefix(text: str) -> bool:
"""Check if the end of text could be the start of a DSML function_calls marker.
We look for suffixes of text that are prefixes of the TOOL_CALLS_START pattern.
This allows us to buffer tokens until we can determine if a tool call is starting.
"""
from mlx_engine.dsml_encoding import TOOL_CALLS_START
# Only check the last portion of text that could overlap with the marker
max_check = len(TOOL_CALLS_START)
tail = text[-max_check:] if len(text) > max_check else text
# Check if any suffix of tail is a prefix of TOOL_CALLS_START
for i in range(len(tail)):
suffix = tail[i:]
if TOOL_CALLS_START.startswith(suffix):
return True
return False
def parse_thinking_models(
responses: Generator[GenerationResponse | None],
think_start: str | None,
think_end: str | None,
starts_in_thinking: bool = True,
) -> Generator[GenerationResponse | None]:
"""Route thinking tokens via is_thinking flag.
Swallows think tag tokens, sets is_thinking on all others.
Always yields tokens with finish_reason to avoid hanging the chunk stream.
"""
is_thinking = starts_in_thinking
for response in responses:
if response is None:
yield None
continue
if response.finish_reason is not None:
yield response.model_copy(update={"is_thinking": False})
continue
if response.text == think_start:
is_thinking = True
continue
if response.text == think_end:
is_thinking = False
continue
yield response.model_copy(update={"is_thinking": is_thinking})
def parse_tool_calls(
responses: Generator[GenerationResponse | None],
tool_parser: ToolParser,
tools: list[dict[str, Any]] | None,
) -> Generator[GenerationResponse | ToolCallResponse | None]:
in_tool_call = False
tool_call_text_parts: list[str] = []
for response in responses:
if response is None:
yield None
continue
if not in_tool_call and response.text.startswith(tool_parser.start_parsing):
in_tool_call = True
if not in_tool_call:
yield response
continue
tool_call_text_parts.append(response.text)
if response.text.endswith(tool_parser.end_parsing):
# parse the actual tool calls from the tool call text
combined = "".join(tool_call_text_parts)
parsed = tool_parser.parse(combined.strip(), tools=tools)
logger.info(f"parsed {tool_call_text_parts=} into {parsed=}")
in_tool_call = False
tool_call_text_parts = []
if parsed is None:
logger.warning(f"tool call parsing failed for text {combined}")
yield response.model_copy(update={"text": combined})
continue
yield ToolCallResponse(
tool_calls=parsed, usage=response.usage, stats=response.stats
)
continue
if response.finish_reason is not None:
logger.info(
"tool call parsing interrupted, yield partial tool call as text"
)
response = response.model_copy(
update={
"text": "".join(tool_call_text_parts),
"token": 0,
}
)
yield response
+1 -1
View File
@@ -1,8 +1,8 @@
from collections.abc import Generator
from typing import Any, Literal
from exo_core.model_cards import ModelId
from exo_core.models import TaggedModel
from exo_core.types.common import ModelId
from exo_core.types.runner_response import (
FinishReason,
GenerationStats,
@@ -5,7 +5,8 @@ from pydantic import model_validator
from exo_core.model_cards import ModelTask
from exo_core.models import CamelCaseModel, TaggedModel
from exo_core.types.common import Host, Id, NodeId
from exo_core.types.runners import RunnerId, ShardAssignments, ShardMetadata
from exo_core.types.runners import RunnerId, ShardAssignments
from exo_core.types.shards import ShardMetadata
class InstanceId(Id):
@@ -2,9 +2,8 @@ from collections.abc import Mapping
from pydantic import model_validator
from exo_core.model_cards import ModelId
from exo_core.models import CamelCaseModel, TaggedModel
from exo_core.types.common import Id, NodeId
from exo_core.types.common import Id, ModelId, NodeId
from exo_core.types.shards import ShardMetadata
@@ -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,
+2 -6
View File
@@ -188,12 +188,8 @@
autoPatchelfIgnoreMissingDeps = (old.autoPatchelfIgnoreMissingDeps or [ ]) ++ [ "libcuda.so.1" ];
});
xgrammar = prev.xgrammar.overrideAttrs (old: {
nativeBuildInputs = (old.nativeBuildInputs or [ ]) ++ [ final.setuptools final.scikit-build-core final.packaging final.pathspec pkgs.cmake final.nanobind ];
prePatch = ''
cat cpp/nanobind/CMakeLists.txt
'';
patches = (old.patches or [ ]) ++ [ ../nix/nanobind_cmake.patch ];
nativeBuildInputs = (old.nativeBuildInputs or [ ]) ++ [ pkgs.cmake ];
patches = (old.patches or [ ]) ++ [ ../nix/xgrammar_cmake.patch ];
});
vllm = prev.vllm.overrideAttrs (old: {
patches = (old.patches or [ ]) ++ [ ../nix/vllm_uv2nix_cmake.patch ];
+18 -11
View File
@@ -1,30 +1,37 @@
import contextlib
import os
from dataclasses import dataclass
from typing import Self, Callable
from typing import Callable, Self
from exo_core.constants import EXO_MODELS_DIR
from exo_core.engine import EngineBuilder, Engine
from exo_core.types.common import ModelId
from exo_core.engine import Engine, EngineBuilder
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 vllm_engine.vllm_generator import VllmBatchEngine
from vllm_engine.vllm_generator import load_vllm_engine
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_engine.batch_generator import BatchGenerator
from vllm_engine.vllm_generator import VllmBatchEngine, load_vllm_engine
@dataclass
class VllmBuilder(EngineBuilder[BoundInstance, TextGeneration, GenerationResponse]):
class VllmBuilder(EngineBuilder[BoundInstance, TextGeneration, GenerationResponse | ToolCallResponse]):
model_id: ModelId
model_path: str
trust_remote_code: bool
cancel_receiver: MpReceiver[TaskId]
event_sender: MpSender[Event]
event_sender: MpSender[tuple[CommandId, ErrorChunk | PrefillProgressChunk]]
bound_instance: BoundInstance
@classmethod
def create(
cls,
bound_instance: BoundInstance,
event_sender: MpSender[Event],
cancel_receiver: MpReceiver[TaskId],
cancel_receiver: MpReceiver[TaskId],
event_sender: MpSender[tuple[CommandId, ErrorChunk | PrefillProgressChunk]],
) -> Self:
mid = bound_instance.instance.shard_assignments.model_id
return cls(
@@ -1,9 +1,8 @@
import torch
from loguru import logger
from mlx_engine.cache import KVPrefixCache
from vllm.v1.worker.gpu_model_runner import GPUModelRunner
from loguru import logger
INITIAL_FRACTION = 0.05
GROWTH_HEADROOM_BYTES = 512 * 1024 * 1024
MIN_GROWTH_BLOCKS = 16
@@ -8,28 +8,27 @@ from collections.abc import Callable, Generator
from dataclasses import dataclass, field
import torch
from mlx_engine.cache import KVPrefixCache
from mlx_engine.utils_mlx import get_eos_token_ids_for_model
from exo_core.tokenizers.tool_parsers import ToolParser, infer_tool_parser
from exo_core.types.common import ModelId
from exo_core.types.runner_response import GenerationResponse
from exo_core.types.tasks import TaskId
from exo_core.types.text_generation import TextGenerationTaskParams
from exo_core.utils.memory import Memory
from exo_core.engine import Engine
from loguru import logger
from vllm.engine.arg_utils import EngineArgs
from vllm.v1.attention.backends.registry import AttentionBackendEnum
from vllm.sampling_params import SamplingParams
from vllm.v1.engine.llm_engine import LLMEngine
from vllm.v1.kv_cache_interface import KVCacheConfig
from exo_core.types.runner_response import (
CompletionTokensDetails,
GenerationResponse,
GenerationStats,
PromptTokensDetails,
Usage,
)
from exo.worker.runner.llm_inference.tool_parsers import ToolParser, infer_tool_parser
from exo_core.types.tasks import TaskId
from exo_core.types.text_generation import TextGenerationTaskParams
from exo_core.utils.memory import Memory
from loguru import logger
from mlx_engine.cache import KVPrefixCache
from mlx_engine.utils_mlx import get_eos_token_ids_for_model
from vllm.engine.arg_utils import EngineArgs
from vllm.sampling_params import SamplingParams
from vllm.v1.attention.backends.registry import AttentionBackendEnum
from vllm.v1.engine.llm_engine import LLMEngine
from vllm.v1.kv_cache_interface import KVCacheConfig
from vllm_engine.growable_cache import (
get_model_runner,
patch_vllm,