urgk
This commit is contained in:
@@ -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"]
|
||||
|
||||
@@ -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 <|DSML|function_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,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
|
||||
@@ -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(
|
||||
|
||||
@@ -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
@@ -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 ];
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user