urgk
This commit is contained in:
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user