Strip vllm generator

This commit is contained in:
Ryuichi Leo Takashige
2026-03-16 22:11:35 +00:00
parent e96f084051
commit ec5d62f935
4 changed files with 313 additions and 669 deletions
+251 -582
View File
@@ -1,11 +1,9 @@
import gc
import itertools
import json
import re
import sys
import time
from collections import deque
from collections.abc import Callable, Generator, Iterable
from collections.abc import Callable, Generator
from dataclasses import dataclass, field
from pathlib import Path
@@ -20,16 +18,10 @@ from exo.shared.types.api import (
PromptTokensDetails,
Usage,
)
from exo.shared.types.chunks import PrefillProgressChunk
from exo.shared.types.common import ModelId
from exo.shared.types.events import ChunkGenerated, Event
from exo.shared.types.memory import Memory
from exo.shared.types.tasks import CANCEL_ALL_TASKS, TaskId, TextGeneration
from exo.shared.types.worker.runner_response import (
GenerationResponse,
ToolCallResponse,
)
from exo.utils.channels import MpReceiver, MpSender
from exo.shared.types.text_generation import TextGenerationTaskParams
from exo.shared.types.worker.runner_response import GenerationResponse
from exo.worker.engines.kv_cache import TorchKVCache
from exo.worker.engines.mlx.cache import KVPrefixCache
from exo.worker.engines.vllm.growable_cache import (
@@ -42,13 +34,6 @@ from exo.worker.engines.vllm.prompt_format import (
make_vllm_sampling_params,
)
from exo.worker.runner.bootstrap import logger
from exo.worker.runner.llm_inference.batch_generator import (
Cancelled,
Finished,
GeneratorQueue,
InferenceGenerator,
)
from exo.worker.runner.llm_inference.model_output_parsers import apply_vllm_parsers
from exo.worker.runner.llm_inference.tool_parsers import ToolParser, infer_tool_parser
@@ -66,628 +51,317 @@ def _build_layer_groups(kv_cache_config: object) -> list[int]:
@dataclass
class _ActiveRequest:
task: TextGeneration
class _EngineRequest:
uid: int
request_id: str
prompt_token_count: int
prompt_token_ids: list[int]
queue: GeneratorQueue[GenerationResponse]
parsed_gen: Generator[GenerationResponse | ToolCallResponse | None]
prefill_done: bool = False
prefill_steps: int = 0
prev_text: str = ""
prev_token_count: int = 0
start_time: float = field(default_factory=time.perf_counter)
first_token_time: float | None = None
on_generation_token: Callable[[], None] | None = None
on_prefill_progress: Callable[[int, int], None] | None = None
@dataclass(eq=False)
class VllmSequentialGenerator(InferenceGenerator):
engine: LLMEngine
model_id: ModelId
tool_parser: ToolParser | None
cancel_receiver: MpReceiver[TaskId]
event_sender: MpSender[Event]
prefix_cache: KVPrefixCache
def _save_prefix_cache(
engine: LLMEngine,
prefix_cache: KVPrefixCache,
request_id: str,
prompt_token_ids: list[int],
prompt_token_count: int,
) -> None:
try:
coordinator = None
model_runner = _growable_model_runner_ref[0]
kv_cache_config = None
try:
engine_core = engine.engine_core.engine_core # type: ignore
coordinator = engine_core.scheduler.kv_cache_manager.coordinator # type: ignore
kv_cache_config = engine_core.scheduler.kv_cache_manager.kv_cache_config # type: ignore
except Exception:
pass
if coordinator is None or model_runner is None or kv_cache_config is None:
return
_cancelled_tasks: set[TaskId] = field(default_factory=set, init=False)
_all_tasks: dict[TaskId, TextGeneration] = field(default_factory=dict, init=False)
_maybe_queue: list[TextGeneration] = field(default_factory=list, init=False)
_queue: deque[TextGeneration] = field(default_factory=deque, init=False)
_maybe_cancel: list[TextGeneration] = field(default_factory=list, init=False)
_active: _ActiveRequest | None = field(default=None, init=False)
def warmup(self) -> None:
tokenizer = self.engine.get_tokenizer()
messages = [{"role": "user", "content": "Prompt to warm up the inference engine. Repeat this."}]
prompt_text: str = tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=True) # type: ignore
token_ids: list[int] = tokenizer.encode(prompt_text, add_special_tokens=False) # type: ignore
params = SamplingParams(max_tokens=50, detokenize=False)
self.engine.add_request("warmup", {"prompt_token_ids": token_ids}, params)
tokens_generated = 0
while self.engine.has_unfinished_requests():
self.engine.step()
tokens_generated += 1
logger.info(f"vLLM warmup complete, generated {tokens_generated} tokens")
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:
self._queue.extend(self._maybe_queue)
self._maybe_queue.clear()
def agree_on_cancellations(self) -> None:
for task_id in self.cancel_receiver.collect():
if task_id == CANCEL_ALL_TASKS:
self._cancelled_tasks.add(CANCEL_ALL_TASKS)
else:
if task_id in self._all_tasks:
self._maybe_cancel.append(self._all_tasks[task_id])
self._cancelled_tasks.update(task.task_id for task in self._maybe_cancel)
self._maybe_cancel.clear()
def step(
self,
) -> Iterable[
tuple[TaskId, GenerationResponse | ToolCallResponse | Cancelled | Finished]
]:
self.agree_on_cancellations()
if self._active is None and not self._queue:
self.agree_on_tasks()
if self._active is None and not self._queue:
return []
tokenizer = self.engine.get_tokenizer()
think_start: str | None = getattr(tokenizer, "think_start", None)
think_end: str | None = getattr(tokenizer, "think_end", None)
if self._active is None:
task = self._queue.popleft()
if self.should_cancel(task.task_id):
self._cancelled_tasks.discard(task.task_id)
return [(task.task_id, Cancelled())]
token_ids, prompt_text, prompt_token_count = format_vllm_prompt(
self.engine, task.task_params
)
logger.info(prompt_text)
request_id = str(task.task_id)
sampling_params = make_vllm_sampling_params(self.engine, task.task_params, self.model_id)
self.engine.add_request(
request_id,
{"prompt_token_ids": token_ids},
sampling_params,
)
queue: GeneratorQueue[GenerationResponse] = GeneratorQueue()
parsed_gen = apply_vllm_parsers(
queue.gen(),
self.model_id,
prompt_text,
self.tool_parser,
task.task_params.tools,
think_start=think_start,
think_end=think_end,
)
self._active = _ActiveRequest(
task=task,
request_id=request_id,
prompt_token_count=prompt_token_count,
prompt_token_ids=token_ids,
queue=queue,
parsed_gen=parsed_gen,
)
active = self._active
if self.should_cancel(active.task.task_id):
self._active = None
self._cancelled_tasks.discard(active.task.task_id)
return [(active.task.task_id, Cancelled())]
if not self.engine.has_unfinished_requests():
self._active = None
return [(active.task.task_id, Finished())]
if not active.prefill_done:
max_batch_tokens: int = getattr(self.engine.model_config, "max_num_batched_tokens", 2048) or 2048 # type: ignore[reportUnknownMemberType]
prefill_steps = 0
while self.engine.has_unfinished_requests():
self.agree_on_cancellations()
if self.should_cancel(active.task.task_id):
self.engine.abort_request([active.request_id])
self._active = None
self._cancelled_tasks.discard(active.task.task_id)
return [(active.task.task_id, Cancelled())]
outputs = self.engine.step()
prefill_steps += 1
for output in outputs:
if output.request_id != active.request_id:
continue
if len(output.outputs[0].token_ids) > 0:
active.first_token_time = time.perf_counter()
active.prefill_done = True
self._save_prefix_cache(active)
break
if not active.prefill_done:
self.event_sender.send(ChunkGenerated(
command_id=active.task.command_id,
chunk=PrefillProgressChunk(
model=self.model_id,
processed_tokens=min(prefill_steps * max_batch_tokens, active.prompt_token_count),
total_tokens=active.prompt_token_count,
),
))
if active.prefill_done:
internal_id: str | None = None
for mgr in coordinator.single_type_managers: # type: ignore
for key in mgr.req_to_blocks: # type: ignore
if str(key).startswith(request_id): # type: ignore
internal_id = str(key) # type: ignore
break
if not active.prefill_done:
self._active = None
return [(active.task.task_id, Finished())]
# Fall through to process the outputs from the final prefill step
else:
outputs = self.engine.step()
finished = False
results: list[
tuple[
TaskId,
GenerationResponse | ToolCallResponse | Cancelled | Finished,
]
] = []
if internal_id:
break
if internal_id is None:
return
null_block = coordinator.block_pool.null_block # type: ignore
block_ids_per_group: list[list[int]] = []
token_offset_per_group: list[int] = []
for mgr in coordinator.single_type_managers: # type: ignore
blocks = mgr.req_to_blocks.get(internal_id) # type: ignore
if not blocks:
block_ids_per_group.append([])
token_offset_per_group.append(0)
continue
block_size: int = mgr.block_size # type: ignore
num_leading_nulls = 0
for b in blocks: # type: ignore
if b is null_block or b.is_null: # type: ignore
num_leading_nulls += 1
else:
break
real_blocks = [b for b in blocks if b is not null_block and not b.is_null] # type: ignore
block_ids_per_group.append([b.block_id for b in real_blocks]) # type: ignore
token_offset_per_group.append(num_leading_nulls * block_size)
layer_to_group = _build_layer_groups(kv_cache_config)
torch_cache = TorchKVCache.from_vllm_cache(
model_runner.kv_caches, # type: ignore
block_ids_per_group,
layer_to_group,
prompt_token_count,
token_offset_per_group,
)
prefix_cache.add_from_torch(prompt_token_ids, torch_cache)
except Exception:
logger.opt(exception=True).warning("Failed to save prefix cache")
def _build_generation_response(
tokenizer: object,
token_id: int,
finish_reason: str | None,
prompt_token_count: int,
completion_tokens: int,
start_time: float,
first_token_time: float | None,
) -> GenerationResponse:
token_text: str = tokenizer.decode([token_id]) # type: ignore[reportUnknownMemberType]
finish_usage: Usage | None = None
finish_stats: GenerationStats | None = None
mapped_finish_reason: str | None = None
if finish_reason:
now = time.perf_counter()
prefill_elapsed = (first_token_time or now) - start_time
decode_elapsed = now - (first_token_time or now)
finish_usage = Usage(
prompt_tokens=prompt_token_count,
completion_tokens=completion_tokens,
total_tokens=prompt_token_count + completion_tokens,
prompt_tokens_details=PromptTokensDetails(),
completion_tokens_details=CompletionTokensDetails(),
)
finish_stats = GenerationStats(
prompt_tps=prompt_token_count / prefill_elapsed if prefill_elapsed > 0 else 0.0,
generation_tps=completion_tokens / decode_elapsed if decode_elapsed > 0 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]
),
)
mapped_finish_reason = finish_reason if finish_reason in ("stop", "length", "content_filter") else "stop"
return GenerationResponse(
text=token_text,
token=token_id,
finish_reason=mapped_finish_reason,
usage=finish_usage,
stats=finish_stats,
)
def vllm_generate(
engine: LLMEngine,
model_id: ModelId,
task: TextGenerationTaskParams,
prompt: str,
prefix_cache: KVPrefixCache,
on_prefill_progress: Callable[[int, int], None] | None = None,
distributed_prompt_progress_callback: Callable[[], None] | None = None,
on_generation_token: Callable[[], None] | None = None,
) -> Generator[GenerationResponse, None, None]:
token_ids, prompt_text, prompt_token_count = format_vllm_prompt(engine, task)
logger.info(prompt_text)
request_id = f"vllm-seq-{time.monotonic_ns()}"
sampling_params = make_vllm_sampling_params(engine, task, model_id)
engine.add_request(request_id, {"prompt_token_ids": token_ids}, sampling_params)
tokenizer = engine.get_tokenizer()
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
prefill_done = False
prefill_steps = 0
while engine.has_unfinished_requests():
if distributed_prompt_progress_callback and not prefill_done:
distributed_prompt_progress_callback()
outputs = engine.step()
for output in outputs:
if output.request_id != active.request_id:
if output.request_id != request_id:
continue
completion = output.outputs[0]
new_token_count = len(completion.token_ids)
new_tokens = completion.token_ids[active.prev_token_count :]
new_text = completion.text[len(active.prev_text) :]
new_tokens = completion.token_ids[prev_token_count:]
finish_reason = completion.finish_reason
prev_token_count = new_token_count
active.prev_text = completion.text
active.prev_token_count = new_token_count
if active.first_token_time is None and new_text:
active.first_token_time = time.perf_counter()
if not prefill_done and not new_tokens:
prefill_steps += 1
if on_prefill_progress:
on_prefill_progress(
min(prefill_steps * max_batch_tokens, prompt_token_count),
prompt_token_count,
)
continue
finish_usage: Usage | None = None
finish_stats: GenerationStats | None = None
mapped_finish_reason: str | None = None
if finish_reason:
now = time.perf_counter()
prefill_elapsed = (active.first_token_time or now) - active.start_time
decode_elapsed = now - (active.first_token_time or now)
finish_usage = Usage(
prompt_tokens=active.prompt_token_count,
completion_tokens=new_token_count,
total_tokens=active.prompt_token_count + new_token_count,
prompt_tokens_details=PromptTokensDetails(),
completion_tokens_details=CompletionTokensDetails(),
)
finish_stats = GenerationStats(
prompt_tps=active.prompt_token_count / prefill_elapsed
if prefill_elapsed > 0
else 0.0,
generation_tps=new_token_count / decode_elapsed
if decode_elapsed > 0
else 0.0,
prompt_tokens=active.prompt_token_count,
generation_tokens=new_token_count,
peak_memory_usage=Memory.from_bytes(
torch.cuda.max_memory_allocated() # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType, reportAttributeAccessIssue]
),
)
mapped_finish_reason = (
finish_reason
if finish_reason in ("stop", "length", "content_filter")
else "stop"
)
finished = True
if not prefill_done and new_tokens:
first_token_time = time.perf_counter()
prefill_done = True
_save_prefix_cache(engine, prefix_cache, request_id, token_ids, prompt_token_count)
tokenizer = self.engine.get_tokenizer()
for i, token_id in enumerate(new_tokens):
is_last = i == len(new_tokens) - 1
token_text: str = tokenizer.decode([token_id]) # type: ignore[reportUnknownMemberType]
active.queue.push(
GenerationResponse(
text=token_text,
token=token_id,
finish_reason=mapped_finish_reason if is_last and finished else None,
usage=finish_usage if is_last and finished else None,
stats=finish_stats if is_last and finished else None,
)
if on_generation_token:
on_generation_token()
yield _build_generation_response(
tokenizer, token_id,
finish_reason if is_last and finish_reason else None,
prompt_token_count, new_token_count,
start_time, first_token_time,
)
try:
parsed = next(active.parsed_gen)
except StopIteration:
self.engine.abort_request([active.request_id])
results.append((active.task.task_id, Finished()))
self._active = None
return results
if parsed is not None:
results.append((active.task.task_id, parsed))
if finished:
logger.info(f"vLLM generation done for request {active.request_id}")
results.append((active.task.task_id, Finished()))
self._active = None
return results
def _get_coordinator(self) -> object | None:
if not hasattr(self, "_coordinator_cached"):
try:
engine_core = self.engine.engine_core.engine_core # type: ignore
self._coordinator_cached: object | None = engine_core.scheduler.kv_cache_manager.coordinator # type: ignore
except Exception:
self._coordinator_cached = None
return self._coordinator_cached
def _get_kv_cache_config(self) -> object | None:
if not hasattr(self, "_kv_cache_config_cached"):
try:
engine_core = self.engine.engine_core.engine_core # type: ignore
self._kv_cache_config_cached: object | None = engine_core.scheduler.kv_cache_manager.kv_cache_config # type: ignore
except Exception:
self._kv_cache_config_cached = None
return self._kv_cache_config_cached
def _save_prefix_cache(self, active: _ActiveRequest) -> None:
try:
coordinator = self._get_coordinator()
model_runner = _growable_model_runner_ref[0]
kv_cache_config = self._get_kv_cache_config()
if coordinator is None or model_runner is None or kv_cache_config is None:
return
internal_id: str | None = None
for mgr in coordinator.single_type_managers: # type: ignore
for key in mgr.req_to_blocks: # type: ignore
if str(key).startswith(active.request_id): # type: ignore
internal_id = str(key) # type: ignore
break
if internal_id:
break
if internal_id is None:
return
null_block = coordinator.block_pool.null_block # type: ignore
block_ids_per_group: list[list[int]] = []
token_offset_per_group: list[int] = []
for mgr in coordinator.single_type_managers: # type: ignore
blocks = mgr.req_to_blocks.get(internal_id) # type: ignore
if not blocks:
block_ids_per_group.append([])
token_offset_per_group.append(0)
continue
block_size: int = mgr.block_size # type: ignore
num_leading_nulls = 0
for b in blocks: # type: ignore
if b is null_block or b.is_null: # type: ignore
num_leading_nulls += 1
else:
break
real_blocks = [b for b in blocks if b is not null_block and not b.is_null] # type: ignore
block_ids_per_group.append([b.block_id for b in real_blocks]) # type: ignore
token_offset_per_group.append(num_leading_nulls * block_size)
layer_to_group = _build_layer_groups(kv_cache_config)
torch_cache = TorchKVCache.from_vllm_cache(
model_runner.kv_caches, # type: ignore
block_ids_per_group,
layer_to_group,
active.prompt_token_count,
token_offset_per_group,
)
self.prefix_cache.add_from_torch(active.prompt_token_ids, torch_cache)
except Exception:
logger.opt(exception=True).warning("Failed to save prefix cache")
def close(self) -> None:
del self.engine
gc.collect()
torch.cuda.empty_cache()
if torch.distributed.is_initialized():
torch.distributed.destroy_process_group()
_EXO_MAX_CONCURRENT_VLLM_REQUESTS = 8
def warmup_vllm_engine(engine: LLMEngine) -> None:
tokenizer = engine.get_tokenizer()
messages = [{"role": "user", "content": "Prompt to warm up the inference engine. Repeat this."}]
prompt_text: str = tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=True) # type: ignore
token_ids: list[int] = tokenizer.encode(prompt_text, add_special_tokens=False) # type: ignore
params = SamplingParams(max_tokens=50, detokenize=False)
engine.add_request("warmup", {"prompt_token_ids": token_ids}, params)
while engine.has_unfinished_requests():
engine.step()
logger.info("vLLM warmup complete")
@dataclass(eq=False)
class VllmBatchGenerator(InferenceGenerator):
class VllmBatchEngine:
engine: LLMEngine
model_id: ModelId
tool_parser: ToolParser | None
cancel_receiver: MpReceiver[TaskId]
event_sender: MpSender[Event]
prefix_cache: KVPrefixCache
_cancelled_tasks: set[TaskId] = field(default_factory=set, init=False)
_all_tasks: dict[TaskId, TextGeneration] = field(default_factory=dict, init=False)
_maybe_queue: list[TextGeneration] = field(default_factory=list, init=False)
_queue: deque[TextGeneration] = field(default_factory=deque, init=False)
_maybe_cancel: list[TextGeneration] = field(default_factory=list, init=False)
_active: dict[str, _ActiveRequest] = field(default_factory=dict, init=False)
_active: dict[int, _EngineRequest] = field(default_factory=dict, init=False)
_next_uid: int = field(default=0, init=False)
def warmup(self) -> None:
tokenizer = self.engine.get_tokenizer()
messages = [{"role": "user", "content": "Prompt to warm up the inference engine. Repeat this."}]
prompt_text: str = tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=True) # type: ignore
token_ids: list[int] = tokenizer.encode(prompt_text, add_special_tokens=False) # type: ignore
params = SamplingParams(max_tokens=50, detokenize=False)
self.engine.add_request("warmup", {"prompt_token_ids": token_ids}, params)
while self.engine.has_unfinished_requests():
self.engine.step()
logger.info("vLLM batch warmup complete")
@property
def has_work(self) -> bool:
return bool(self._active) or self.engine.has_unfinished_requests()
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:
self._queue.extend(self._maybe_queue)
self._maybe_queue.clear()
def agree_on_cancellations(self) -> None:
for task_id in self.cancel_receiver.collect():
if task_id == CANCEL_ALL_TASKS:
self._cancelled_tasks.add(CANCEL_ALL_TASKS)
else:
if task_id in self._all_tasks:
self._maybe_cancel.append(self._all_tasks[task_id])
self._cancelled_tasks.update(task.task_id for task in self._maybe_cancel)
self._maybe_cancel.clear()
def _start_request(self, task: TextGeneration) -> _ActiveRequest:
token_ids, prompt_text, prompt_token_count = format_vllm_prompt(
self.engine, task.task_params
)
def submit(
self,
task_params: TextGenerationTaskParams,
prompt: str,
on_prefill_progress: Callable[[int, int], None] | None = None,
distributed_prompt_progress_callback: Callable[[], None] | None = None,
on_generation_token: Callable[[], None] | None = None,
) -> int:
token_ids, prompt_text, prompt_token_count = format_vllm_prompt(self.engine, task_params)
logger.info(prompt_text)
request_id = str(task.task_id)
sampling_params = make_vllm_sampling_params(self.engine, task.task_params, self.model_id)
self.engine.add_request(
request_id,
{"prompt_token_ids": token_ids},
sampling_params,
)
tokenizer = self.engine.get_tokenizer()
think_start: str | None = getattr(tokenizer, "think_start", None)
think_end: str | None = getattr(tokenizer, "think_end", None)
queue: GeneratorQueue[GenerationResponse] = GeneratorQueue()
parsed_gen = apply_vllm_parsers(
queue.gen(),
self.model_id,
prompt_text,
self.tool_parser,
task.task_params.tools,
think_start=think_start,
think_end=think_end,
)
return _ActiveRequest(
task=task,
uid = self._next_uid
self._next_uid += 1
request_id = f"vllm-batch-{uid}"
sampling_params = make_vllm_sampling_params(self.engine, task_params, self.model_id)
self.engine.add_request(request_id, {"prompt_token_ids": token_ids}, sampling_params)
self._active[uid] = _EngineRequest(
uid=uid,
request_id=request_id,
prompt_token_count=prompt_token_count,
prompt_token_ids=token_ids,
queue=queue,
parsed_gen=parsed_gen,
on_generation_token=on_generation_token,
on_prefill_progress=on_prefill_progress,
)
return uid
def _apply_cancellations(
self,
) -> list[tuple[TaskId, Cancelled]]:
if not self._cancelled_tasks:
def step(self) -> list[tuple[int, GenerationResponse]]:
if not self.has_work:
return []
cancel_all = CANCEL_ALL_TASKS in self._cancelled_tasks
rids_to_abort: list[str] = []
results: list[tuple[TaskId, Cancelled]] = []
for rid, active in list(self._active.items()):
if active.task.task_id in self._cancelled_tasks or cancel_all:
rids_to_abort.append(rid)
results.append((active.task.task_id, Cancelled()))
del self._active[rid]
if rids_to_abort:
self.engine.abort_request(rids_to_abort)
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 step(
self,
) -> Iterable[
tuple[TaskId, GenerationResponse | ToolCallResponse | Cancelled | Finished]
]:
self.agree_on_cancellations()
if not self._queue:
self.agree_on_tasks()
results: list[
tuple[TaskId, GenerationResponse | ToolCallResponse | Cancelled | Finished]
] = []
while self._queue and len(self._active) < _EXO_MAX_CONCURRENT_VLLM_REQUESTS:
task = self._queue.popleft()
if self.should_cancel(task.task_id):
self._cancelled_tasks.discard(task.task_id)
results.append((task.task_id, Cancelled()))
continue
active = self._start_request(task)
self._active[active.request_id] = active
if not self._active:
return itertools.chain(results, self._apply_cancellations())
outputs = self.engine.step()
tokenizer = self.engine.get_tokenizer()
max_batch_tokens: int = getattr(self.engine.model_config, "max_num_batched_tokens", 2048) or 2048 # type: ignore[reportUnknownMemberType]
results: list[tuple[int, GenerationResponse]] = []
rid_to_uid = {req.request_id: uid for uid, req in self._active.items()}
for output in outputs:
rid = output.request_id
if rid not in self._active:
uid = rid_to_uid.get(output.request_id)
if uid is None:
continue
active = self._active[rid]
req = self._active[uid]
completion = output.outputs[0]
new_token_count = len(completion.token_ids)
new_tokens = completion.token_ids[active.prev_token_count:]
new_tokens = completion.token_ids[req.prev_token_count:]
finish_reason = completion.finish_reason
req.prev_token_count = new_token_count
active.prev_text = completion.text
active.prev_token_count = new_token_count
if active.first_token_time is None and new_tokens:
active.first_token_time = time.perf_counter()
if not active.prefill_done:
active.prefill_done = True
self._save_prefix_cache(active)
if not req.prefill_done and not new_tokens:
req.prefill_steps += 1
if req.on_prefill_progress:
req.on_prefill_progress(
min(req.prefill_steps * max_batch_tokens, req.prompt_token_count),
req.prompt_token_count,
)
continue
finish_usage: Usage | None = None
finish_stats: GenerationStats | None = None
mapped_finish_reason: str | None = None
finished = False
if finish_reason:
now = time.perf_counter()
prefill_elapsed = (active.first_token_time or now) - active.start_time
decode_elapsed = now - (active.first_token_time or now)
finish_usage = Usage(
prompt_tokens=active.prompt_token_count,
completion_tokens=new_token_count,
total_tokens=active.prompt_token_count + new_token_count,
prompt_tokens_details=PromptTokensDetails(),
completion_tokens_details=CompletionTokensDetails(),
if not req.prefill_done and new_tokens:
req.first_token_time = time.perf_counter()
req.prefill_done = True
_save_prefix_cache(
self.engine, self.prefix_cache,
req.request_id, req.prompt_token_ids, req.prompt_token_count,
)
finish_stats = GenerationStats(
prompt_tps=active.prompt_token_count / prefill_elapsed if prefill_elapsed > 0 else 0.0,
generation_tps=new_token_count / decode_elapsed if decode_elapsed > 0 else 0.0,
prompt_tokens=active.prompt_token_count,
generation_tokens=new_token_count,
peak_memory_usage=Memory.from_bytes(
torch.cuda.max_memory_allocated() # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType, reportAttributeAccessIssue]
),
)
mapped_finish_reason = (
finish_reason if finish_reason in ("stop", "length", "content_filter") else "stop"
)
finished = True
for i, token_id in enumerate(new_tokens):
is_last = i == len(new_tokens) - 1
token_text: str = tokenizer.decode([token_id]) # type: ignore[reportUnknownMemberType]
active.queue.push(
GenerationResponse(
text=token_text,
token=token_id,
finish_reason=mapped_finish_reason if is_last and finished else None,
usage=finish_usage if is_last and finished else None,
stats=finish_stats if is_last and finished else None,
if req.on_generation_token:
req.on_generation_token()
results.append((uid, _build_generation_response(
tokenizer, token_id,
finish_reason if is_last and finish_reason else None,
req.prompt_token_count, new_token_count,
req.start_time, req.first_token_time,
)))
if finish_reason:
del self._active[uid]
for req in self._active.values():
if not req.prefill_done:
req.prefill_steps += 1
if req.on_prefill_progress:
req.on_prefill_progress(
min(req.prefill_steps * max_batch_tokens, req.prompt_token_count),
req.prompt_token_count,
)
)
try:
parsed = next(active.parsed_gen)
except StopIteration:
self.engine.abort_request([rid])
results.append((active.task.task_id, Finished()))
del self._active[rid]
break
if parsed is not None:
results.append((active.task.task_id, parsed))
else:
if finished:
logger.info(f"vLLM generation done for request {rid}")
results.append((active.task.task_id, Finished()))
del self._active[rid]
max_batch_tokens: int = getattr(self.engine.model_config, "max_num_batched_tokens", 2048) or 2048 # type: ignore[reportUnknownMemberType]
for active in self._active.values():
if not active.prefill_done:
active.prefill_steps += 1
self.event_sender.send(ChunkGenerated(
command_id=active.task.command_id,
chunk=PrefillProgressChunk(
model=self.model_id,
processed_tokens=min(active.prefill_steps * max_batch_tokens, active.prompt_token_count),
total_tokens=active.prompt_token_count,
),
))
return results
return itertools.chain(results, self._apply_cancellations())
def _get_coordinator(self) -> object | None:
if not hasattr(self, "_coordinator_cached"):
try:
engine_core = self.engine.engine_core.engine_core # type: ignore
self._coordinator_cached: object | None = engine_core.scheduler.kv_cache_manager.coordinator # type: ignore
except Exception:
self._coordinator_cached = None
return self._coordinator_cached
def _get_kv_cache_config(self) -> object | None:
if not hasattr(self, "_kv_cache_config_cached"):
try:
engine_core = self.engine.engine_core.engine_core # type: ignore
self._kv_cache_config_cached: object | None = engine_core.scheduler.kv_cache_manager.kv_cache_config # type: ignore
except Exception:
self._kv_cache_config_cached = None
return self._kv_cache_config_cached
def _save_prefix_cache(self, active: _ActiveRequest) -> None:
try:
coordinator = self._get_coordinator()
model_runner = _growable_model_runner_ref[0]
kv_cache_config = self._get_kv_cache_config()
if coordinator is None or model_runner is None or kv_cache_config is None:
return
internal_id: str | None = None
for mgr in coordinator.single_type_managers: # type: ignore
for key in mgr.req_to_blocks: # type: ignore
if str(key).startswith(active.request_id): # type: ignore
internal_id = str(key) # type: ignore
break
if internal_id:
break
if internal_id is None:
return
null_block = coordinator.block_pool.null_block # type: ignore
block_ids_per_group: list[list[int]] = []
token_offset_per_group: list[int] = []
for mgr in coordinator.single_type_managers: # type: ignore
blocks = mgr.req_to_blocks.get(internal_id) # type: ignore
if not blocks:
block_ids_per_group.append([])
token_offset_per_group.append(0)
continue
block_size: int = mgr.block_size # type: ignore
num_leading_nulls = 0
for b in blocks: # type: ignore
if b is null_block or b.is_null: # type: ignore
num_leading_nulls += 1
else:
break
real_blocks = [b for b in blocks if b is not null_block and not b.is_null] # type: ignore
block_ids_per_group.append([b.block_id for b in real_blocks]) # type: ignore
token_offset_per_group.append(num_leading_nulls * block_size)
layer_to_group = _build_layer_groups(kv_cache_config)
torch_cache = TorchKVCache.from_vllm_cache(
model_runner.kv_caches, # type: ignore
block_ids_per_group,
layer_to_group,
active.prompt_token_count,
token_offset_per_group,
)
self.prefix_cache.add_from_torch(active.prompt_token_ids, torch_cache)
except Exception:
logger.opt(exception=True).warning("Failed to save prefix cache")
def cancel(self, uids: list[int]) -> None:
rids = [self._active[uid].request_id for uid in uids if uid in self._active]
if rids:
self.engine.abort_request(rids)
for uid in uids:
self._active.pop(uid, None)
def close(self) -> None:
for rid in list(self._active):
self.engine.abort_request([rid])
rids = [req.request_id for req in self._active.values()]
if rids:
self.engine.abort_request(rids)
self._active.clear()
del self.engine
gc.collect()
@@ -796,11 +470,6 @@ def load_vllm_engine(
load_format="fastsafetensors",
enable_prefix_caching=False,
attention_backend="TRITON_ATTN",
# torch.compile's piecewise compilation produces kernels that assume
# KV cache blocks are written by previous compiled forward passes.
# Our prefix cache writes KV externally via tensor indexing and skips
# the forward pass, which leaves the compiled kernel's internal state
# inconsistent with the actual block data, producing garbage output.
enforce_eager=True,
disable_log_stats=True,
)
@@ -140,13 +140,14 @@ class SequentialGenerator(InferenceGenerator):
| None
) = field(default=None, init=False)
def warmup(self):
self.check_for_cancel_every = warmup_inference(
model=self.model,
tokenizer=self.tokenizer,
group=self.group,
model_id=self.model_id,
)
def warmup(self) -> None:
if self.model is not None:
self.check_for_cancel_every = warmup_inference(
model=self.model,
tokenizer=self.tokenizer,
group=self.group,
model_id=self.model_id,
)
def submit(
self,
@@ -304,7 +305,7 @@ class SequentialGenerator(InferenceGenerator):
@dataclass(eq=False)
class BatchGenerator(InferenceGenerator):
model: Model
model: Model | None
tokenizer: TokenizerWrapper
group: mx.distributed.Group | None
kv_prefix_cache: KVPrefixCache | None
@@ -313,6 +314,8 @@ class BatchGenerator(InferenceGenerator):
device_rank: int
cancel_receiver: MpReceiver[TaskId]
event_sender: MpSender[Event]
_gen: ExoBatchGenerator # ExoBatchGenerator or 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)
@@ -320,7 +323,6 @@ class BatchGenerator(InferenceGenerator):
_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)
_mlx_gen: ExoBatchGenerator = field(init=False)
_active_tasks: dict[
int,
tuple[
@@ -330,21 +332,14 @@ class BatchGenerator(InferenceGenerator):
],
] = field(default_factory=dict, init=False)
def __post_init__(self) -> None:
self._mlx_gen = ExoBatchGenerator(
model=self.model,
tokenizer=self.tokenizer,
group=self.group,
kv_prefix_cache=self.kv_prefix_cache,
)
def warmup(self):
self.check_for_cancel_every = warmup_inference(
model=self.model,
tokenizer=self.tokenizer,
group=self.group,
model_id=self.model_id,
)
def warmup(self) -> None:
if self.model is not None:
self.check_for_cancel_every = warmup_inference(
model=self.model,
tokenizer=self.tokenizer,
group=self.group,
model_id=self.model_id,
)
def submit(
self,
@@ -386,7 +381,7 @@ class BatchGenerator(InferenceGenerator):
self.agree_on_tasks()
# Submit any queued tasks to the engine
while self._queue and len(self._active_tasks) < EXO_MAX_CONCURRENT_REQUESTS:
while self._queue and len(self._active_tasks) < self.max_concurrent_requests:
task = self._queue.popleft()
try:
uid = self._start_task(task)
@@ -411,10 +406,10 @@ class BatchGenerator(InferenceGenerator):
)
self._active_tasks[uid] = (task, queue, output_generator)
if not self._mlx_gen.has_work:
if not self._gen.has_work:
return self._apply_cancellations()
results = self._mlx_gen.step()
results = self._gen.step()
output: list[
tuple[TaskId, GenerationResponse | ToolCallResponse | Cancelled | Finished]
@@ -456,7 +451,7 @@ class BatchGenerator(InferenceGenerator):
del self._active_tasks[uid]
if uids_to_cancel:
self._mlx_gen.cancel(uids_to_cancel)
self._gen.cancel(uids_to_cancel)
already_cancelled = {tid for tid, _ in results}
for tid in self._cancelled_tasks:
@@ -516,7 +511,7 @@ class BatchGenerator(InferenceGenerator):
self.agree_on_tasks()
return self._mlx_gen.submit(
return self._gen.submit(
task_params=task.task_params,
prompt=prompt,
on_prefill_progress=on_prefill_progress,
@@ -525,5 +520,7 @@ class BatchGenerator(InferenceGenerator):
)
def close(self) -> None:
self._mlx_gen.close()
del self.model, self.tokenizer, self.group
self._gen.close()
if self.model is not None:
del self.model
del self.tokenizer, self.group
@@ -49,12 +49,10 @@ def apply_all_parsers(
starts_in_thinking=detect_thinking_prompt_suffix(prompt, tokenizer),
)
if issubclass(model_type, GptOssModel):
lower = model_id.normalize().lower()
if issubclass(model_type, GptOssModel) or "gpt-oss" in lower or "gpt_oss" in lower:
mlx_generator = parse_gpt_oss(mlx_generator)
elif (
issubclass(model_type, DeepseekV32Model)
and "deepseek" in model_id.normalize().lower()
):
elif issubclass(model_type, DeepseekV32Model) or "deepseek" in lower:
mlx_generator = parse_deepseek_v32(mlx_generator)
elif tool_parser:
mlx_generator = parse_tool_calls(mlx_generator, tool_parser, tools)
@@ -62,42 +60,6 @@ def apply_all_parsers(
return mlx_generator
def apply_vllm_parsers(
receiver: Generator[GenerationResponse | None],
model_id: ModelId,
prompt: str,
tool_parser: ToolParser | None,
tools: list[dict[str, Any]] | None,
think_start: str | None = None,
think_end: str | None = None,
) -> Generator[GenerationResponse | ToolCallResponse | None]:
gen = receiver
lower = model_id.normalize().lower()
if "gpt-oss" in lower or "gpt_oss" in lower:
return parse_gpt_oss(gen)
if "deepseek" in lower:
gen = parse_thinking_models(
gen,
think_start or "<think>",
think_end or "</think>",
starts_in_thinking=prompt.rstrip().endswith(think_start or "<think>"),
)
return parse_deepseek_v32(gen)
if think_start is not None:
gen = parse_thinking_models(
gen,
think_start,
think_end,
starts_in_thinking=prompt.rstrip().endswith(think_start),
)
if tool_parser:
gen = parse_tool_calls(gen, tool_parser, tools)
return gen
_GPT_OSS_CHANNEL_TOKEN = 200005
_GPT_OSS_MESSAGE_TOKEN = 200008
+31 -15
View File
@@ -462,7 +462,15 @@ class MlxBuilder(Builder):
cancel_receiver=self.cancel_receiver,
event_sender=self.event_sender,
)
from exo.worker.runner.llm_inference.batch_generator import ExoBatchGenerator
logger.info("using BatchGenerator")
gen = ExoBatchGenerator(
model=self.inference_model,
tokenizer=self.tokenizer,
group=self.group,
kv_prefix_cache=kv_prefix_cache,
)
return BatchGenerator(
model=self.inference_model,
tokenizer=self.tokenizer,
@@ -473,6 +481,7 @@ class MlxBuilder(Builder):
device_rank=device_rank,
cancel_receiver=self.cancel_receiver,
event_sender=self.event_sender,
_gen=gen,
)
def shutdown_cleanup(self) -> None:
@@ -509,28 +518,35 @@ class VllmBuilder(Builder):
)
def build(self) -> InferenceGenerator:
if os.environ.get("EXO_NO_BATCH"):
from exo.worker.engines.vllm.vllm_generator import VllmSequentialGenerator
from mlx_lm.tokenizer_utils import TokenizerWrapper
logger.info("using VllmSequentialGenerator (batching disabled)")
return VllmSequentialGenerator(
engine=self._engine,
model_id=self.model_id,
tool_parser=self._tool_parser,
cancel_receiver=self.cancel_receiver,
event_sender=self.event_sender,
prefix_cache=self._prefix_cache,
)
from exo.worker.engines.vllm.vllm_generator import VllmBatchGenerator
from exo.worker.engines.vllm.vllm_generator import (
VllmBatchEngine,
warmup_vllm_engine,
)
logger.info("using VllmBatchGenerator")
return VllmBatchGenerator(
warmup_vllm_engine(self._engine)
gen = VllmBatchEngine(
engine=self._engine,
model_id=self.model_id,
prefix_cache=self._prefix_cache,
)
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(
model=None,
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,
prefix_cache=self._prefix_cache,
_gen=gen,
max_concurrent_requests=max_concurrent,
)
def shutdown_cleanup(self) -> None: