Type error lol
This commit is contained in:
@@ -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())
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user