Reduce tokenizer overhead and offload CUDA VMM publication (#37330)

Co-authored-by: Shiyan Deng <dsy842974287@meta.com>
Co-authored-by: Yinghai Lu <yinghai@meta.com>
This commit is contained in:
Lianmin Zheng
2026-09-02 17:21:08 -07:00
committed by GitHub
co-authored by Shiyan Deng Yinghai Lu
parent f15748d965
commit ff04a00d73
5 changed files with 350 additions and 67 deletions
@@ -28,6 +28,7 @@ from llguidance.torch import (
fill_next_token_bitmask_par,
fill_next_token_bitmask_par_with_draft_tokens,
)
from transformers import PreTrainedTokenizerFast
from sglang.srt.constrained.base_grammar_backend import (
BaseGrammarBackend,
@@ -95,6 +96,22 @@ def _normalize_eos_token_ids(
return list(eos_token_ids)
def _create_llguidance_tokenizer(
tokenizer,
n_vocab: Optional[int],
eos_token: Optional[Union[int, List[int]]],
) -> LLTokenizer:
if isinstance(tokenizer, PreTrainedTokenizerFast):
backend_tokenizer = tokenizer.backend_tokenizer
if backend_tokenizer.padding is None and backend_tokenizer.truncation is None:
return LLTokenizer(
backend_tokenizer.to_str(),
n_vocab=n_vocab,
eos_token=(tokenizer.eos_token_id if eos_token is None else eos_token),
)
return from_tokenizer(tokenizer, n_vocab, eos_token=eos_token)
class GuidanceGrammar(BaseGrammarObject):
def __init__(
@@ -223,7 +240,7 @@ class GuidanceBackend(BaseGrammarBackend):
self.tokenizer = tokenizer
self.any_whitespace = any_whitespace
self.whitespace_pattern = whitespace_pattern
self.llguidance_tokenizer = from_tokenizer(
self.llguidance_tokenizer = _create_llguidance_tokenizer(
self.tokenizer,
n_vocab,
eos_token=_normalize_eos_token_ids(eos_token_ids),
+15 -11
View File
@@ -818,7 +818,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
state = self.rid_to_state[obj.rid]
if obj.return_prompt_token_ids:
state.prompt_token_ids = list(tokenized_obj.input_ids)
self._send_one_request(tokenized_obj)
await self._send_one_request(tokenized_obj)
async for response in self._wait_one_response(obj, request):
yield response
else:
@@ -1554,15 +1554,17 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
)
)
def _send_one_request(
async def _send_one_request(
self,
tokenized_obj: Union[TokenizedGenerateReqInput, TokenizedEmbeddingReqInput],
):
prepared_mm_items = []
dispatched = False
try:
prepared_mm_items = self.cuda_vmm_feature_transport.prepare_for_dispatch(
(tokenized_obj.mm_inputs,)
prepared_mm_items = (
await self.cuda_vmm_feature_transport.prepare_for_dispatch_async(
(tokenized_obj.mm_inputs,)
)
)
tokenized_obj.time_stats.set_api_server_dispatch_time()
tokenized_obj = wrap_shm_features(tokenized_obj)
@@ -1576,7 +1578,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
if not dispatched:
self.cuda_vmm_feature_transport.cancel_for_dispatch(prepared_mm_items)
def _send_batch_request(
async def _send_batch_request(
self,
tokenized_objs: List[
Union[TokenizedGenerateReqInput, TokenizedEmbeddingReqInput]
@@ -1586,8 +1588,10 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
prepared_mm_items = []
dispatched = False
try:
prepared_mm_items = self.cuda_vmm_feature_transport.prepare_for_dispatch(
tokenized_obj.mm_inputs for tokenized_obj in tokenized_objs
prepared_mm_items = (
await self.cuda_vmm_feature_transport.prepare_for_dispatch_async(
tokenized_obj.mm_inputs for tokenized_obj in tokenized_objs
)
)
set_time_batch(tokenized_objs, "set_api_server_dispatch_time")
@@ -1823,7 +1827,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
if getattr(obj, "parallel_sample_num", 1) == 1:
if self._should_use_batch_tokenization(batch_size, obj):
tokenized_objs = await self._batch_tokenize_and_process(batch_size, obj)
self._send_batch_request(tokenized_objs)
await self._send_batch_request(tokenized_objs)
# Set up generators for each request in the batch
for i in range(batch_size):
@@ -1848,7 +1852,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
state = self.rid_to_state[tmp_obj.rid]
if tmp_obj.return_prompt_token_ids:
state.prompt_token_ids = list(tokenized_obj.input_ids)
self._send_one_request(tokenized_obj)
await self._send_one_request(tokenized_obj)
generators.append(self._wait_one_response(tmp_obj, request))
rids.append(tmp_obj.rid)
else:
@@ -1881,7 +1885,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
tokenized_obj.sampling_params.max_new_tokens = 0
tokenized_obj.stream = False
self._init_req_state(tmp_obj)
self._send_one_request(tokenized_obj)
await self._send_one_request(tokenized_obj)
await self._wait_one_response(tmp_obj, request).__anext__()
# Expand requests, assign new rids for them, and send them
@@ -1901,7 +1905,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
tokenized_obj.time_stats = state.time_stats
if tmp_obj.return_prompt_token_ids:
state.prompt_token_ids = list(tokenized_objs[i].input_ids)
self._send_one_request(tokenized_obj)
await self._send_one_request(tokenized_obj)
generators.append(self._wait_one_response(tmp_obj, request))
rids.append(tmp_obj.rid)
@@ -1,5 +1,7 @@
from __future__ import annotations
import asyncio
import concurrent.futures
import logging
import os
import secrets
@@ -155,6 +157,33 @@ def _build_packed_tensor_layout(
return layouts, next_offset
def _prepare_pinned_copy_source(tensor: torch.Tensor) -> torch.Tensor:
if not tensor.is_contiguous():
tensor = tensor.contiguous()
if tensor.device.type == "cpu" and not tensor.is_pinned():
tensor = tensor.pin_memory()
return tensor
def _pack_pinned_copy_sources(
tensors: Sequence[torch.Tensor],
layouts: Sequence[_CudaVmmPackedTensorLayout],
packed_data_nbytes: int,
) -> torch.Tensor | None:
if not all(tensor.device.type == "cpu" for tensor in tensors):
return None
staging = torch.empty(
packed_data_nbytes, dtype=torch.uint8, device="cpu", pin_memory=True
)
for tensor, layout in zip(tensors, layouts, strict=True):
source = tensor if tensor.is_contiguous() else tensor.contiguous()
staging[
layout.relative_offset : layout.relative_offset + layout.data_nbytes
].copy_(source.reshape(-1).view(torch.uint8))
return staging
def _contains_tensor_container(value) -> bool:
return isinstance(value, (list, tuple)) and any(
isinstance(item, torch.Tensor) or _contains_tensor_container(item)
@@ -205,6 +234,7 @@ class CudaVmmMemoryPool:
self.memory_pool = None
self._fd_broker: _PosixFdBroker | None = None
self.posix_socket_path: str | None = None
self._publish_stream = None
self._recycle_stream = None
self._recycle_thread = None
@@ -240,6 +270,7 @@ class CudaVmmMemoryPool:
self.available_chunks = [_CudaVmmMemoryChunk(0, self.allocation_size)]
self.occupied_chunks = []
self._publish_stream = torch.cuda.Stream(device=self.device_index)
self._recycle_stream = torch.cuda.Stream(device=self.device_index)
self._recycle_thread = threading.Thread(
target=self._recycle_loop,
@@ -360,11 +391,8 @@ class CudaVmmMemoryPool:
def wrap_tensor(self, tensor: torch.Tensor):
self._raise_if_failed()
if not tensor.is_contiguous():
tensor = tensor.contiguous()
data_nbytes = tensor.numel() * tensor.element_size()
required_size = align_up(self.control_size + data_nbytes, _CONTROL_ALIGNMENT)
source_bytes = tensor.reshape(-1).view(torch.uint8)
chunk = self._reserve_for_publish(required_size)
if chunk is None:
@@ -374,8 +402,13 @@ class CudaVmmMemoryPool:
producer_stream = None
copy_synchronized = False
try:
with torch.cuda.device(self.device_index):
producer_stream = torch.cuda.current_stream(self.device_index)
with (
torch.cuda.device(self.device_index),
torch.cuda.stream(self._publish_stream),
):
producer_stream = self._publish_stream
copy_source = _prepare_pinned_copy_source(tensor)
source_bytes = copy_source.reshape(-1).view(torch.uint8)
control_offset = chunk.start
data_offset = control_offset + self.control_size
control_end = control_offset + self.control_size
@@ -444,21 +477,31 @@ class CudaVmmMemoryPool:
producer_stream = None
copy_synchronized = False
try:
contiguous_tensors = [
tensor if tensor.is_contiguous() else tensor.contiguous()
for tensor in tensors
]
with torch.cuda.device(self.device_index):
producer_stream = torch.cuda.current_stream(self.device_index)
with (
torch.cuda.device(self.device_index),
torch.cuda.stream(self._publish_stream),
):
producer_stream = self._publish_stream
packed_source = _pack_pinned_copy_sources(
tensors, layouts, packed_data_nbytes
)
control_offset = chunk.start
data_offset = control_offset + self.control_size
control_end = control_offset + self.control_size
self.memory_pool[control_offset:control_end].zero_()
for tensor, layout in zip(contiguous_tensors, layouts):
data_start = data_offset + layout.relative_offset
if packed_source is not None:
self.memory_pool[
data_start : data_start + layout.data_nbytes
].copy_(tensor.reshape(-1).view(torch.uint8), non_blocking=True)
data_offset : data_offset + packed_data_nbytes
].copy_(packed_source, non_blocking=True)
else:
copy_sources = [
_prepare_pinned_copy_source(tensor) for tensor in tensors
]
for tensor, layout in zip(copy_sources, layouts, strict=True):
data_start = data_offset + layout.relative_offset
self.memory_pool[
data_start : data_start + layout.data_nbytes
].copy_(tensor.reshape(-1).view(torch.uint8), non_blocking=True)
# A single synchronization publishes every child together.
producer_stream.synchronize()
copy_synchronized = True
@@ -570,24 +613,35 @@ class CudaVmmMemoryPool:
self._stop_recycler.set()
def _recycle_chunks(self) -> None:
if not self.occupied_chunks:
return
remaining = []
recycled = []
with (
torch.cuda.device(self.device_index),
torch.cuda.stream(self._recycle_stream),
):
for chunk in self.occupied_chunks:
ack_start = chunk.start
ack_end = ack_start + self.consumer_count * _CONTROL_WORD_BYTES
ack_count = int(
torch.count_nonzero(
self.memory_pool[ack_start:ack_end].view(torch.int32)
).item()
)
if ack_count == self.consumer_count:
recycled.append(_CudaVmmMemoryChunk(chunk.start, chunk.end))
else:
remaining.append(chunk)
acknowledgement_words = torch.stack(
[
self.memory_pool[
chunk.start : chunk.start
+ self.consumer_count * _CONTROL_WORD_BYTES
].view(torch.int32)
for chunk in self.occupied_chunks
]
)
acknowledgement_counts = (
torch.count_nonzero(acknowledgement_words, dim=1).cpu().tolist()
)
for chunk, acknowledgement_count in zip(
self.occupied_chunks, acknowledgement_counts, strict=True
):
if acknowledgement_count == self.consumer_count:
recycled.append(_CudaVmmMemoryChunk(chunk.start, chunk.end))
else:
remaining.append(chunk)
self.available_chunks.extend(recycled)
self.occupied_chunks = remaining
@@ -941,6 +995,7 @@ class CudaVmmFeatureTransport:
def __init__(self, server_args, mm_processor) -> None:
self.pool: CudaVmmMemoryPool | None = None
self._publisher_executor: concurrent.futures.ThreadPoolExecutor | None = None
if get_mm().mm_feature_transport != "cuda_vmm":
return
if mm_processor is None:
@@ -959,6 +1014,36 @@ class CudaVmmFeatureTransport:
consumer_count=get_vmm_feature_consumer_count(),
allow_posix_fallback=get_parallel().nnodes == 1,
)
self._publisher_executor = concurrent.futures.ThreadPoolExecutor(
max_workers=1,
thread_name_prefix="cuda-vmm-publisher",
)
async def prepare_for_dispatch_async(
self,
mm_inputs_batch: Iterable[MultimodalProcessorOutput | None],
) -> list[MultimodalDataItem]:
"""Publish features without blocking the tokenizer event loop."""
mm_inputs_batch = tuple(mm_inputs_batch)
if self.pool is None or not any(
mm_inputs is not None and mm_inputs.mm_items
for mm_inputs in mm_inputs_batch
):
return []
if self._publisher_executor is None:
raise RuntimeError("CUDA VMM feature transport is shutting down")
future = asyncio.get_running_loop().run_in_executor(
self._publisher_executor,
self.prepare_for_dispatch,
mm_inputs_batch,
)
try:
return await asyncio.shield(future)
except asyncio.CancelledError:
prepared_mm_items = await future
self.cancel_for_dispatch(prepared_mm_items)
raise
def prepare_for_dispatch(
self,
@@ -991,7 +1076,7 @@ class CudaVmmFeatureTransport:
pack_candidates = [
(item, item.feature)
for item in mm_items
if item.modality == Modality.IMAGE
if item.modality in (Modality.IMAGE, Modality.VIDEO)
and isinstance(item.feature, torch.Tensor)
and item.feature.numel() > 0
and not item.model_specific_data.get(
@@ -1069,4 +1154,8 @@ class CudaVmmFeatureTransport:
def shutdown(self) -> None:
if self.pool is None:
return
publisher_executor = self._publisher_executor
if publisher_executor is not None:
publisher_executor.shutdown(wait=True, cancel_futures=True)
self._publisher_executor = None
self.pool.shutdown()