Strip vllm generator
This commit is contained in:
@@ -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
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user