man i gotta commit

This commit is contained in:
Evan
2026-03-23 16:00:57 +00:00
parent 8cd6191c70
commit 6355f3d8fb
328 changed files with 941 additions and 3347 deletions
-1
View File
@@ -1 +0,0 @@
3.13
+4 -1
View File
@@ -10,6 +10,8 @@ dependencies = [
"mlx-cuda-13==0.30.6; sys_platform == 'linux'",
"fastsafetensors>=0.1.10; sys_platform == 'linux'",
"exo_core",
"mlx_engine", # TODO(evan): remove this dependency!
"loguru>=0.7.3",
]
[tool.uv]
@@ -17,7 +19,8 @@ environments = ["sys_platform == 'linux'"]
[tool.uv.sources]
exo_core = { workspace = true }
exo_core = { workspace = true, editable = true }
mlx_engine = { workspace = true, editable = true }
vllm = { git = "https://github.com/hmellor/vllm.git", rev = "b99bedc737166ae5ca98cb9e3534b96e0c8c69aa" }
torch = [
{ index = "pytorch-cu130", marker = "platform_machine == 'aarch64'" },
@@ -0,0 +1,84 @@
from dataclasses import dataclass
from typing import Self, Callable
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.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
@dataclass
class VllmBuilder(EngineBuilder[BoundInstance, TextGeneration, GenerationResponse]):
model_id: ModelId
model_path: str
trust_remote_code: bool
cancel_receiver: MpReceiver[TaskId]
event_sender: MpSender[Event]
bound_instance: BoundInstance
@classmethod
def create(
cls,
bound_instance: BoundInstance,
event_sender: MpSender[Event],
cancel_receiver: MpReceiver[TaskId],
) -> Self:
mid = bound_instance.instance.shard_assignments.model_id
return cls(
mid,
str(EXO_MODELS_DIR / mid.normalize()),
bound_instance.bound_shard.model_card.trust_remote_code,
cancel_receiver,
event_sender,
bound_instance,
)
def connect(self) -> None:
raise NotImplementedError(
"Multiple node VLLM instances are not supported at the moment!"
)
def load(
self,
on_timeout: Callable[[], None],
on_layer_loaded: Callable[[int, int], None],
) -> None:
self._engine, self._tool_parser, self._prefix_cache = load_vllm_engine(
model_path=self.model_path,
model_id=self.model_id,
trust_remote_code=self.trust_remote_code,
n_layers=self.bound_instance.bound_shard.model_card.n_layers,
on_layer_loaded=on_layer_loaded,
)
def build(self) -> Engine[TextGeneration, GenerationResponse]:
gen = VllmBatchEngine(
engine=self._engine,
model_id=self.model_id,
prefix_cache=self._prefix_cache,
)
from mlx_lm.tokenizer_utils import TokenizerWrapper
tokenizer = TokenizerWrapper(self._engine.get_tokenizer())
max_concurrent = 1 if os.environ.get("EXO_NO_BATCH") else 8
logger.info(f"using BatchGenerator (vLLM, max_concurrent={max_concurrent})")
return BatchGenerator(
tokenizer=tokenizer,
group=None,
tool_parser=self._tool_parser,
kv_prefix_cache=None,
model_id=self.model_id,
device_rank=0,
cancel_receiver=self.cancel_receiver,
event_sender=self.event_sender,
_gen=gen,
max_concurrent_requests=max_concurrent,
)
def close(self) -> None:
with contextlib.suppress(NameError, AttributeError):
del self._engine, self._prefix_cache, self._tool_parser
@@ -1,8 +1,8 @@
import torch
from mlx_engine.cache import KVPrefixCache
from vllm.v1.worker.gpu_model_runner import GPUModelRunner
from exo.shared.logging import logger
from exo.worker.engines.mlx.cache import KVPrefixCache
from loguru import logger
INITIAL_FRACTION = 0.05
GROWTH_HEADROOM_BYTES = 512 * 1024 * 1024
@@ -1,14 +1,13 @@
from exo.worker.engines.mlx.utils_mlx import (
apply_chat_template,
get_eos_token_ids_for_model,
)
from exo_core.types.common import ModelId
from exo_core.types.text_generation import TextGenerationTaskParams
from mlx_lm.tokenizer_utils import TokenizerWrapper
from vllm.sampling_params import SamplingParams
from vllm.v1.engine.llm_engine import LLMEngine
from exo.worker.engines.mlx.utils_mlx import (
apply_chat_template,
get_eos_token_ids_for_model,
)
def format_vllm_prompt(
engine: LLMEngine, params: TextGenerationTaskParams
@@ -8,25 +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.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.api.types import (
from exo_core.types.runner_response import (
CompletionTokensDetails,
GenerationStats,
PromptTokensDetails,
Usage,
)
from exo.worker.engines.mlx.cache import KVPrefixCache
from exo.worker.engines.mlx.utils_mlx import get_eos_token_ids_for_model
from exo.worker.runner.bootstrap import logger
from exo.worker.runner.llm_inference.tool_parsers import ToolParser, infer_tool_parser
from vllm_engine.growable_cache import (
get_model_runner,
@@ -177,9 +179,7 @@ def _build_generation_response(
else 0.0,
prompt_tokens=prompt_token_count,
generation_tokens=completion_tokens,
peak_memory_usage=Memory.from_bytes(
torch.cuda.max_memory_allocated() # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType, reportAttributeAccessIssue]
),
peak_memory_usage=Memory.from_bytes(torch.cuda.max_memory_allocated()),
)
mapped_finish_reason = (
finish_reason
@@ -215,7 +215,7 @@ def vllm_generate(
stop_ids = _stop_token_ids(tokenizer, model_id)
max_batch_tokens: int = (
getattr(engine.model_config, "max_num_batched_tokens", 2048) or 2048
) # type: ignore[reportUnknownMemberType]
)
start_time = time.perf_counter()
first_token_time: float | None = None
prev_token_count = 0
@@ -433,7 +433,7 @@ class VllmBatchEngine:
return results
def cancel(self, task_ids: list[TaskId]) -> None:
to_abort = [tid for tid in task_ids if tid in self._active]
to_abort = [str(tid) for tid in task_ids if tid in self._active]
if to_abort:
self.engine.abort_request(to_abort)
for tid in task_ids:
@@ -480,18 +480,18 @@ def set_n_layers(n: int) -> None:
def _wrap_weights_iterator(
original: Callable[..., Generator[tuple[str, "torch.Tensor"], None, None]],
) -> Callable[..., Generator[tuple[str, "torch.Tensor"], None, None]]: # pyright: ignore[reportUnknownParameterType]
original: Callable[..., Generator[tuple[str, torch.Tensor], None, None]],
) -> Callable[..., Generator[tuple[str, torch.Tensor], None, None]]:
def patched(
hf_weights_files: list[str], *args: object, **kwargs: object
) -> Generator[tuple[str, "torch.Tensor"], None, None]: # pyright: ignore[reportUnknownParameterType]
) -> Generator[tuple[str, torch.Tensor], None, None]:
callback = get_weight_loading_callback()
if callback is not None and hf_weights_files:
total_layers = get_n_layers()
seen_layers: set[int] = set()
last_reported = 0
for name, tensor in original(hf_weights_files, *args, **kwargs): # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType]
yield name, tensor # pyright: ignore[reportUnknownArgumentType]
for name, tensor in original(hf_weights_files, *args, **kwargs):
yield name, tensor
match = _LAYER_INDEX_PATTERN.search(name)
if match:
seen_layers.add(int(match.group(1)))
@@ -501,19 +501,19 @@ def _wrap_weights_iterator(
last_reported = current
callback(total_layers, total_layers)
else:
yield from original(hf_weights_files, *args, **kwargs) # pyright: ignore[reportUnknownMemberType]
yield from original(hf_weights_files, *args, **kwargs)
return patched
def _monkey_patch_iterator(weight_utils: object, attr_name: str) -> None: # pyright: ignore[reportUnknownParameterType]
def _monkey_patch_iterator(weight_utils: object, attr_name: str) -> None:
original = getattr(weight_utils, attr_name, None)
if original is None:
return
patched = _wrap_weights_iterator(original) # pyright: ignore[reportUnknownArgumentType]
patched = _wrap_weights_iterator(original) # pyright: ignore[reportAny]
setattr(weight_utils, attr_name, patched)
for mod in list(sys.modules.values()):
if mod is None or mod is weight_utils:
if mod is weight_utils:
continue
for name in list(vars(mod)):
if vars(mod)[name] is original:
@@ -527,21 +527,21 @@ def _patch_weight_loading_progress() -> None:
_weight_loading_patched = True
from vllm.model_executor.model_loader import (
weight_utils, # pyright: ignore[reportMissingImports]
weight_utils,
)
_monkey_patch_iterator(weight_utils, "safetensors_weights_iterator")
_monkey_patch_iterator(weight_utils, "fastsafetensors_weights_iterator")
import huggingface_hub # pyright: ignore[reportMissingImports]
import huggingface_hub
def _noop_metadata(*_a: object, **_kw: object) -> None:
pass # pyright: ignore[reportUnknownParameterType]
pass
original_metadata = huggingface_hub.get_safetensors_metadata # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType]
huggingface_hub.get_safetensors_metadata = _noop_metadata # pyright: ignore[reportAttributeAccessIssue]
original_metadata = huggingface_hub.get_safetensors_metadata
huggingface_hub.get_safetensors_metadata = _noop_metadata
for mod in list(sys.modules.values()):
if mod is None or mod is huggingface_hub:
if mod is huggingface_hub:
continue
for attr in list(vars(mod)):
if vars(mod)[attr] is original_metadata:
@@ -571,7 +571,7 @@ def load_vllm_engine(
trust_remote_code=trust_remote_code,
load_format="fastsafetensors",
enable_prefix_caching=False,
attention_backend="TRITON_ATTN",
attention_backend=AttentionBackendEnum.TRITON_ATTN,
enforce_eager=True,
disable_log_stats=True,
)
-2348
View File
File diff suppressed because it is too large Load Diff