Type error lol

This commit is contained in:
Ryuichi Leo Takashige
2026-03-17 18:00:39 +00:00
parent 6a3eb2f37d
commit cacd26e63c
4 changed files with 16 additions and 18 deletions
+6 -5
View File
@@ -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