diff --git a/src/exo/worker/engines/mlx/cache.py b/src/exo/worker/engines/mlx/cache.py index 17f65202..fc2fc5c2 100644 --- a/src/exo/worker/engines/mlx/cache.py +++ b/src/exo/worker/engines/mlx/cache.py @@ -1,6 +1,6 @@ import os from copy import deepcopy -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, cast import mlx.core as mx import psutil @@ -135,8 +135,7 @@ class KVPrefixCache: def _get_mlx_cache(self, index: int) -> MLXCacheType: cached = self.caches[index] - assert not isinstance(cached, TorchKVCache) - return cached + return cast(MLXCacheType, cached) def _get_snapshot( self, entry_index: int, target_token_count: int @@ -220,7 +219,7 @@ class KVPrefixCache: def lookup( self, prompt_token_ids: list[int] - ) -> tuple[TorchKVCache | None, int, int | None]: + ) -> tuple["TorchKVCache | None", int, int | None]: from exo.worker.engines.kv_cache import TorchKVCache prompt_mx = mx.array(prompt_token_ids) @@ -251,7 +250,9 @@ class KVPrefixCache: torch_cache = TorchKVCache.from_mlx_cache(cached) return torch_cache.trim_to(best_length), best_length, best_index - def add_from_torch(self, prompt_token_ids: list[int], cache: TorchKVCache) -> None: + def add_from_torch( + self, prompt_token_ids: list[int], cache: "TorchKVCache" + ) -> None: self._evict_if_needed() self.prompts.append(mx.array(prompt_token_ids)) self.caches.append(cache.detach_cpu()) diff --git a/src/exo/worker/engines/vllm/growable_cache.py b/src/exo/worker/engines/vllm/growable_cache.py index e2763d64..d412ccc1 100644 --- a/src/exo/worker/engines/vllm/growable_cache.py +++ b/src/exo/worker/engines/vllm/growable_cache.py @@ -1,13 +1,8 @@ -from typing import TYPE_CHECKING - import torch +from vllm.v1.worker.gpu_model_runner import GPUModelRunner from exo.shared.logging import logger - -if TYPE_CHECKING: - from vllm.v1.worker.gpu_model_runner import GPUModelRunner - - from exo.worker.engines.mlx.cache import KVPrefixCache +from exo.worker.engines.mlx.cache import KVPrefixCache INITIAL_FRACTION = 0.05 GROWTH_HEADROOM_BYTES = 512 * 1024 * 1024 diff --git a/src/exo/worker/runner/llm_inference/batch_generator.py b/src/exo/worker/runner/llm_inference/batch_generator.py index b329aeba..8eba561d 100644 --- a/src/exo/worker/runner/llm_inference/batch_generator.py +++ b/src/exo/worker/runner/llm_inference/batch_generator.py @@ -317,7 +317,7 @@ class BatchGenerator(InferenceGenerator): device_rank: int cancel_receiver: MpReceiver[TaskId] event_sender: MpSender[Event] - _gen: ExoBatchGenerator | VllmBatchEngine + _gen: "ExoBatchGenerator | VllmBatchEngine" max_concurrent_requests: int = EXO_MAX_CONCURRENT_REQUESTS check_for_cancel_every: int = 50 diff --git a/src/exo/worker/runner/llm_inference/runner.py b/src/exo/worker/runner/llm_inference/runner.py index 352d6b75..5aac0a00 100644 --- a/src/exo/worker/runner/llm_inference/runner.py +++ b/src/exo/worker/runner/llm_inference/runner.py @@ -6,6 +6,7 @@ from abc import ABC, abstractmethod from collections.abc import Callable from dataclasses import dataclass from enum import Enum +from typing import TYPE_CHECKING import mlx.core as mx from anyio import WouldBlock @@ -61,7 +62,6 @@ from exo.worker.engines.mlx.utils_mlx import ( initialize_mlx, load_mlx_items, ) -from exo.worker.engines.vllm.vllm_generator import VllmBatchEngine from exo.worker.runner.bootstrap import logger from exo.worker.runner.llm_inference.batch_generator import ( BatchGenerator, @@ -72,6 +72,9 @@ from exo.worker.runner.llm_inference.batch_generator import ( from .batch_generator import Cancelled, Finished from .tool_parsers import make_mlx_parser +if TYPE_CHECKING: + pass + class ExitCode(str, Enum): AllTasksComplete = "AllTasksComplete" @@ -481,7 +484,7 @@ class MlxBuilder(Builder): ) def close(self): - with contextlib.suppress(NameError): + with contextlib.suppress(NameError, AttributeError): del self.inference_model, self.tokenizer @@ -515,9 +518,8 @@ class VllmBuilder(Builder): ) def build(self) -> InferenceGenerator: - from mlx_lm.tokenizer_utils import TokenizerWrapper - from exo.worker.engines.vllm.vllm_generator import ( + VllmBatchEngine, warmup_vllm_engine, ) @@ -546,5 +548,5 @@ class VllmBuilder(Builder): ) def close(self) -> None: - with contextlib.suppress(NameError): + with contextlib.suppress(NameError, AttributeError): del self._engine, self._prefix_cache, self._tool_parser