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,
fill_next_token_bitmask_par_with_draft_tokens, fill_next_token_bitmask_par_with_draft_tokens,
) )
from transformers import PreTrainedTokenizerFast
from sglang.srt.constrained.base_grammar_backend import ( from sglang.srt.constrained.base_grammar_backend import (
BaseGrammarBackend, BaseGrammarBackend,
@@ -95,6 +96,22 @@ def _normalize_eos_token_ids(
return list(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): class GuidanceGrammar(BaseGrammarObject):
def __init__( def __init__(
@@ -223,7 +240,7 @@ class GuidanceBackend(BaseGrammarBackend):
self.tokenizer = tokenizer self.tokenizer = tokenizer
self.any_whitespace = any_whitespace self.any_whitespace = any_whitespace
self.whitespace_pattern = whitespace_pattern self.whitespace_pattern = whitespace_pattern
self.llguidance_tokenizer = from_tokenizer( self.llguidance_tokenizer = _create_llguidance_tokenizer(
self.tokenizer, self.tokenizer,
n_vocab, n_vocab,
eos_token=_normalize_eos_token_ids(eos_token_ids), 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] state = self.rid_to_state[obj.rid]
if obj.return_prompt_token_ids: if obj.return_prompt_token_ids:
state.prompt_token_ids = list(tokenized_obj.input_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): async for response in self._wait_one_response(obj, request):
yield response yield response
else: else:
@@ -1554,15 +1554,17 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
) )
) )
def _send_one_request( async def _send_one_request(
self, self,
tokenized_obj: Union[TokenizedGenerateReqInput, TokenizedEmbeddingReqInput], tokenized_obj: Union[TokenizedGenerateReqInput, TokenizedEmbeddingReqInput],
): ):
prepared_mm_items = [] prepared_mm_items = []
dispatched = False dispatched = False
try: try:
prepared_mm_items = self.cuda_vmm_feature_transport.prepare_for_dispatch( prepared_mm_items = (
(tokenized_obj.mm_inputs,) 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.time_stats.set_api_server_dispatch_time()
tokenized_obj = wrap_shm_features(tokenized_obj) tokenized_obj = wrap_shm_features(tokenized_obj)
@@ -1576,7 +1578,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
if not dispatched: if not dispatched:
self.cuda_vmm_feature_transport.cancel_for_dispatch(prepared_mm_items) self.cuda_vmm_feature_transport.cancel_for_dispatch(prepared_mm_items)
def _send_batch_request( async def _send_batch_request(
self, self,
tokenized_objs: List[ tokenized_objs: List[
Union[TokenizedGenerateReqInput, TokenizedEmbeddingReqInput] Union[TokenizedGenerateReqInput, TokenizedEmbeddingReqInput]
@@ -1586,8 +1588,10 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
prepared_mm_items = [] prepared_mm_items = []
dispatched = False dispatched = False
try: try:
prepared_mm_items = self.cuda_vmm_feature_transport.prepare_for_dispatch( prepared_mm_items = (
tokenized_obj.mm_inputs for tokenized_obj in tokenized_objs 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") 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 getattr(obj, "parallel_sample_num", 1) == 1:
if self._should_use_batch_tokenization(batch_size, obj): if self._should_use_batch_tokenization(batch_size, obj):
tokenized_objs = await self._batch_tokenize_and_process(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 # Set up generators for each request in the batch
for i in range(batch_size): for i in range(batch_size):
@@ -1848,7 +1852,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
state = self.rid_to_state[tmp_obj.rid] state = self.rid_to_state[tmp_obj.rid]
if tmp_obj.return_prompt_token_ids: if tmp_obj.return_prompt_token_ids:
state.prompt_token_ids = list(tokenized_obj.input_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)) generators.append(self._wait_one_response(tmp_obj, request))
rids.append(tmp_obj.rid) rids.append(tmp_obj.rid)
else: else:
@@ -1881,7 +1885,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
tokenized_obj.sampling_params.max_new_tokens = 0 tokenized_obj.sampling_params.max_new_tokens = 0
tokenized_obj.stream = False tokenized_obj.stream = False
self._init_req_state(tmp_obj) 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__() await self._wait_one_response(tmp_obj, request).__anext__()
# Expand requests, assign new rids for them, and send them # 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 tokenized_obj.time_stats = state.time_stats
if tmp_obj.return_prompt_token_ids: if tmp_obj.return_prompt_token_ids:
state.prompt_token_ids = list(tokenized_objs[i].input_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)) generators.append(self._wait_one_response(tmp_obj, request))
rids.append(tmp_obj.rid) rids.append(tmp_obj.rid)
@@ -1,5 +1,7 @@
from __future__ import annotations from __future__ import annotations
import asyncio
import concurrent.futures
import logging import logging
import os import os
import secrets import secrets
@@ -155,6 +157,33 @@ def _build_packed_tensor_layout(
return layouts, next_offset 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: def _contains_tensor_container(value) -> bool:
return isinstance(value, (list, tuple)) and any( return isinstance(value, (list, tuple)) and any(
isinstance(item, torch.Tensor) or _contains_tensor_container(item) isinstance(item, torch.Tensor) or _contains_tensor_container(item)
@@ -205,6 +234,7 @@ class CudaVmmMemoryPool:
self.memory_pool = None self.memory_pool = None
self._fd_broker: _PosixFdBroker | None = None self._fd_broker: _PosixFdBroker | None = None
self.posix_socket_path: str | None = None self.posix_socket_path: str | None = None
self._publish_stream = None
self._recycle_stream = None self._recycle_stream = None
self._recycle_thread = None self._recycle_thread = None
@@ -240,6 +270,7 @@ class CudaVmmMemoryPool:
self.available_chunks = [_CudaVmmMemoryChunk(0, self.allocation_size)] self.available_chunks = [_CudaVmmMemoryChunk(0, self.allocation_size)]
self.occupied_chunks = [] 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_stream = torch.cuda.Stream(device=self.device_index)
self._recycle_thread = threading.Thread( self._recycle_thread = threading.Thread(
target=self._recycle_loop, target=self._recycle_loop,
@@ -360,11 +391,8 @@ class CudaVmmMemoryPool:
def wrap_tensor(self, tensor: torch.Tensor): def wrap_tensor(self, tensor: torch.Tensor):
self._raise_if_failed() self._raise_if_failed()
if not tensor.is_contiguous():
tensor = tensor.contiguous()
data_nbytes = tensor.numel() * tensor.element_size() data_nbytes = tensor.numel() * tensor.element_size()
required_size = align_up(self.control_size + data_nbytes, _CONTROL_ALIGNMENT) 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) chunk = self._reserve_for_publish(required_size)
if chunk is None: if chunk is None:
@@ -374,8 +402,13 @@ class CudaVmmMemoryPool:
producer_stream = None producer_stream = None
copy_synchronized = False copy_synchronized = False
try: try:
with torch.cuda.device(self.device_index): with (
producer_stream = torch.cuda.current_stream(self.device_index) 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 control_offset = chunk.start
data_offset = control_offset + self.control_size data_offset = control_offset + self.control_size
control_end = control_offset + self.control_size control_end = control_offset + self.control_size
@@ -444,21 +477,31 @@ class CudaVmmMemoryPool:
producer_stream = None producer_stream = None
copy_synchronized = False copy_synchronized = False
try: try:
contiguous_tensors = [ with (
tensor if tensor.is_contiguous() else tensor.contiguous() torch.cuda.device(self.device_index),
for tensor in tensors torch.cuda.stream(self._publish_stream),
] ):
with torch.cuda.device(self.device_index): producer_stream = self._publish_stream
producer_stream = torch.cuda.current_stream(self.device_index) packed_source = _pack_pinned_copy_sources(
tensors, layouts, packed_data_nbytes
)
control_offset = chunk.start control_offset = chunk.start
data_offset = control_offset + self.control_size data_offset = control_offset + self.control_size
control_end = control_offset + self.control_size control_end = control_offset + self.control_size
self.memory_pool[control_offset:control_end].zero_() self.memory_pool[control_offset:control_end].zero_()
for tensor, layout in zip(contiguous_tensors, layouts): if packed_source is not None:
data_start = data_offset + layout.relative_offset
self.memory_pool[ self.memory_pool[
data_start : data_start + layout.data_nbytes data_offset : data_offset + packed_data_nbytes
].copy_(tensor.reshape(-1).view(torch.uint8), non_blocking=True) ].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. # A single synchronization publishes every child together.
producer_stream.synchronize() producer_stream.synchronize()
copy_synchronized = True copy_synchronized = True
@@ -570,24 +613,35 @@ class CudaVmmMemoryPool:
self._stop_recycler.set() self._stop_recycler.set()
def _recycle_chunks(self) -> None: def _recycle_chunks(self) -> None:
if not self.occupied_chunks:
return
remaining = [] remaining = []
recycled = [] recycled = []
with ( with (
torch.cuda.device(self.device_index), torch.cuda.device(self.device_index),
torch.cuda.stream(self._recycle_stream), torch.cuda.stream(self._recycle_stream),
): ):
for chunk in self.occupied_chunks: acknowledgement_words = torch.stack(
ack_start = chunk.start [
ack_end = ack_start + self.consumer_count * _CONTROL_WORD_BYTES self.memory_pool[
ack_count = int( chunk.start : chunk.start
torch.count_nonzero( + self.consumer_count * _CONTROL_WORD_BYTES
self.memory_pool[ack_start:ack_end].view(torch.int32) ].view(torch.int32)
).item() for chunk in self.occupied_chunks
) ]
if ack_count == self.consumer_count: )
recycled.append(_CudaVmmMemoryChunk(chunk.start, chunk.end)) acknowledgement_counts = (
else: torch.count_nonzero(acknowledgement_words, dim=1).cpu().tolist()
remaining.append(chunk) )
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.available_chunks.extend(recycled)
self.occupied_chunks = remaining self.occupied_chunks = remaining
@@ -941,6 +995,7 @@ class CudaVmmFeatureTransport:
def __init__(self, server_args, mm_processor) -> None: def __init__(self, server_args, mm_processor) -> None:
self.pool: CudaVmmMemoryPool | None = None self.pool: CudaVmmMemoryPool | None = None
self._publisher_executor: concurrent.futures.ThreadPoolExecutor | None = None
if get_mm().mm_feature_transport != "cuda_vmm": if get_mm().mm_feature_transport != "cuda_vmm":
return return
if mm_processor is None: if mm_processor is None:
@@ -959,6 +1014,36 @@ class CudaVmmFeatureTransport:
consumer_count=get_vmm_feature_consumer_count(), consumer_count=get_vmm_feature_consumer_count(),
allow_posix_fallback=get_parallel().nnodes == 1, 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( def prepare_for_dispatch(
self, self,
@@ -991,7 +1076,7 @@ class CudaVmmFeatureTransport:
pack_candidates = [ pack_candidates = [
(item, item.feature) (item, item.feature)
for item in mm_items 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 isinstance(item.feature, torch.Tensor)
and item.feature.numel() > 0 and item.feature.numel() > 0
and not item.model_specific_data.get( and not item.model_specific_data.get(
@@ -1069,4 +1154,8 @@ class CudaVmmFeatureTransport:
def shutdown(self) -> None: def shutdown(self) -> None:
if self.pool is None: if self.pool is None:
return 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() self.pool.shutdown()
@@ -9,7 +9,7 @@ import pickle
import queue import queue
import threading import threading
import unittest import unittest
from unittest.mock import MagicMock, patch from unittest.mock import patch
import torch import torch
@@ -184,24 +184,25 @@ class TestCudaVmmTransport(CustomTestCase):
finally: finally:
pool.shutdown() pool.shutdown()
def test_packed_tensors_round_trip_through_one_shared_buffer(self): def test_packed_cpu_tensors_round_trip_through_one_shared_buffer(self):
pool = CudaVmmMemoryPool(4 << 20, 60, 0, 1, allow_posix_fallback=True) pool = CudaVmmMemoryPool(4 << 20, 60, 0, 1, allow_posix_fallback=True)
sources = [ sources = [
torch.arange(24, dtype=torch.float32, device="cuda:0") torch.arange(24, dtype=torch.float32).reshape(4, 6).transpose(0, 1),
.reshape(4, 6)
.transpose(0, 1),
torch.arange(7, dtype=torch.bfloat16), torch.arange(7, dtype=torch.bfloat16),
torch.arange(5, dtype=torch.int64, device="cuda:0"), torch.arange(5, dtype=torch.int64),
] ]
expected = [source.contiguous().cpu() for source in sources] expected = [source.contiguous().cpu() for source in sources]
proxies = reconstructed = None proxies = reconstructed = None
try: try:
stream = MagicMock(wraps=torch.cuda.current_stream(0)) with patch.object(
with patch("torch.cuda.current_stream", return_value=stream): pool._publish_stream,
"synchronize",
wraps=pool._publish_stream.synchronize,
) as synchronize:
proxies = pool.wrap_tensors(sources) proxies = pool.wrap_tensors(sources)
self.assertIsNotNone(proxies) self.assertIsNotNone(proxies)
self.assertEqual(stream.synchronize.call_count, 1) self.assertEqual(synchronize.call_count, 1)
self.assertEqual(len(pool.occupied_chunks), 1) self.assertEqual(len(pool.occupied_chunks), 1)
self.assertTrue( self.assertTrue(
all( all(
@@ -337,11 +338,13 @@ class TestCudaVmmTransport(CustomTestCase):
def test_failed_cleanup_sync_quarantines_pool(self): def test_failed_cleanup_sync_quarantines_pool(self):
pool = CudaVmmMemoryPool(4 << 20, 60, 0, 1, allow_posix_fallback=True) pool = CudaVmmMemoryPool(4 << 20, 60, 0, 1, allow_posix_fallback=True)
stream = MagicMock()
stream.synchronize.side_effect = RuntimeError("forced sync failure")
try: try:
with ( with (
patch("torch.cuda.current_stream", return_value=stream), patch.object(
pool._publish_stream,
"synchronize",
side_effect=RuntimeError("forced sync failure"),
),
self.assertRaisesRegex(RuntimeError, "forced sync failure"), self.assertRaisesRegex(RuntimeError, "forced sync failure"),
): ):
pool.wrap_tensor(torch.ones(16, device="cuda:0")) pool.wrap_tensor(torch.ones(16, device="cuda:0"))
@@ -359,20 +362,21 @@ class TestCudaVmmTransport(CustomTestCase):
shutdown_entered = threading.Event() shutdown_entered = threading.Event()
shutdown_finished = threading.Event() shutdown_finished = threading.Event()
errors = [] errors = []
real_stream = torch.cuda.current_stream(0) real_synchronize = pool._publish_stream.synchronize
stream = MagicMock(wraps=real_stream)
def synchronize(): def synchronize():
publisher_entered.set() publisher_entered.set()
if not allow_publisher_to_finish.wait(timeout=10): if not allow_publisher_to_finish.wait(timeout=10):
raise TimeoutError("publisher was not released") raise TimeoutError("publisher was not released")
real_stream.synchronize() real_synchronize()
stream.synchronize.side_effect = synchronize
def publish(): def publish():
try: try:
with patch("torch.cuda.current_stream", return_value=stream): with patch.object(
pool._publish_stream,
"synchronize",
side_effect=synchronize,
):
pool.wrap_tensor(torch.ones(16, device="cuda:0")) pool.wrap_tensor(torch.ones(16, device="cuda:0"))
except Exception as error: # pragma: no cover except Exception as error: # pragma: no cover
errors.append(error) errors.append(error)
@@ -1,7 +1,10 @@
import asyncio
import concurrent.futures
import threading
import unittest import unittest
from contextlib import nullcontext from contextlib import nullcontext
from types import SimpleNamespace from types import SimpleNamespace
from unittest.mock import MagicMock, call, patch from unittest.mock import AsyncMock, MagicMock, call, patch
import torch import torch
@@ -214,6 +217,7 @@ class TestCudaVmmFeatureTransport(unittest.TestCase):
transport = CudaVmmFeatureTransport(SimpleNamespace(), None) transport = CudaVmmFeatureTransport(SimpleNamespace(), None)
self.assertEqual(transport.prepare_for_dispatch([None]), []) self.assertEqual(transport.prepare_for_dispatch([None]), [])
self.assertEqual(asyncio.run(transport.prepare_for_dispatch_async([None])), [])
transport.cancel_for_dispatch([]) transport.cancel_for_dispatch([])
transport.shutdown() transport.shutdown()
self.assertIsNone(transport.pool) self.assertIsNone(transport.pool)
@@ -252,6 +256,28 @@ class TestCudaVmmFeatureTransport(unittest.TestCase):
transport.pool.wrap_tensor.assert_not_called() transport.pool.wrap_tensor.assert_not_called()
self.assertEqual([item.feature for item in items], proxies) self.assertEqual([item.feature for item in items], proxies)
def test_video_clip_features_are_packed_per_request(self):
from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem
from sglang.srt.utils.cuda_vmm_transport_utils import (
CudaVmmFeatureTransport,
)
transport = object.__new__(CudaVmmFeatureTransport)
transport.pool = MagicMock()
features = [torch.arange(4), torch.arange(4, 8)]
proxies = [object(), object()]
transport.pool.wrap_tensors.return_value = proxies
items = [
MultimodalDataItem(modality=Modality.VIDEO, feature=feature)
for feature in features
]
transport.wrap_items(items)
transport.pool.wrap_tensors.assert_called_once_with(features)
transport.pool.wrap_tensor.assert_not_called()
self.assertEqual([item.feature for item in items], proxies)
def test_deferred_features_are_not_packed(self): def test_deferred_features_are_not_packed(self):
from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem
from sglang.srt.utils.cuda_ipc_transport_utils import ( from sglang.srt.utils.cuda_ipc_transport_utils import (
@@ -351,7 +377,7 @@ class TestCudaVmmFeatureTransport(unittest.TestCase):
manager = object.__new__(TokenizerManager) manager = object.__new__(TokenizerManager)
transport = MagicMock() transport = MagicMock()
transport.prepare_for_dispatch.return_value = [] transport.prepare_for_dispatch_async = AsyncMock(return_value=[])
manager.cuda_vmm_feature_transport = transport manager.cuda_vmm_feature_transport = transport
manager._dispatch_to_scheduler = MagicMock() manager._dispatch_to_scheduler = MagicMock()
tokenized_obj = SimpleNamespace( tokenized_obj = SimpleNamespace(
@@ -362,10 +388,10 @@ class TestCudaVmmFeatureTransport(unittest.TestCase):
) )
with patch.object(tokenizer_manager, "wrap_shm_features", lambda obj: obj): with patch.object(tokenizer_manager, "wrap_shm_features", lambda obj: obj):
manager._send_one_request(tokenized_obj) asyncio.run(manager._send_one_request(tokenized_obj))
manager._dispatch_to_scheduler.assert_called_once_with(tokenized_obj) manager._dispatch_to_scheduler.assert_called_once_with(tokenized_obj)
transport.prepare_for_dispatch.assert_called_once_with((None,)) transport.prepare_for_dispatch_async.assert_awaited_once_with((None,))
transport.cancel_for_dispatch.assert_not_called() transport.cancel_for_dispatch.assert_not_called()
def test_failed_dispatch_cancels_published_items(self): def test_failed_dispatch_cancels_published_items(self):
@@ -388,16 +414,16 @@ class TestCudaVmmFeatureTransport(unittest.TestCase):
time_stats=MagicMock(), time_stats=MagicMock(),
wrap_pickle_fields=MagicMock(), wrap_pickle_fields=MagicMock(),
) )
transport.prepare_for_dispatch.return_value = items transport.prepare_for_dispatch_async = AsyncMock(return_value=items)
manager.cuda_vmm_feature_transport = transport manager.cuda_vmm_feature_transport = transport
with ( with (
patch.object(tokenizer_manager, "wrap_shm_features", lambda obj: obj), patch.object(tokenizer_manager, "wrap_shm_features", lambda obj: obj),
self.assertRaisesRegex(RuntimeError, "send failed"), self.assertRaisesRegex(RuntimeError, "send failed"),
): ):
manager._send_one_request(tokenized_obj) asyncio.run(manager._send_one_request(tokenized_obj))
transport.prepare_for_dispatch.assert_called_once_with( transport.prepare_for_dispatch_async.assert_awaited_once_with(
(tokenized_obj.mm_inputs,) (tokenized_obj.mm_inputs,)
) )
transport.cancel_for_dispatch.assert_called_once_with(items) transport.cancel_for_dispatch.assert_called_once_with(items)
@@ -424,18 +450,158 @@ class TestCudaVmmFeatureTransport(unittest.TestCase):
time_stats=time_stats, time_stats=time_stats,
wrap_pickle_fields=MagicMock(), wrap_pickle_fields=MagicMock(),
) )
transport.prepare_for_dispatch.return_value = items transport.prepare_for_dispatch_async = AsyncMock(return_value=items)
manager.cuda_vmm_feature_transport = transport manager.cuda_vmm_feature_transport = transport
with ( with (
patch.object(tokenizer_manager, "wrap_shm_features", lambda obj: obj), patch.object(tokenizer_manager, "wrap_shm_features", lambda obj: obj),
self.assertRaisesRegex(RuntimeError, "bookkeeping failed"), self.assertRaisesRegex(RuntimeError, "bookkeeping failed"),
): ):
manager._send_one_request(tokenized_obj) asyncio.run(manager._send_one_request(tokenized_obj))
manager._dispatch_to_scheduler.assert_called_once_with(tokenized_obj) manager._dispatch_to_scheduler.assert_called_once_with(tokenized_obj)
transport.cancel_for_dispatch.assert_not_called() transport.cancel_for_dispatch.assert_not_called()
def test_async_publication_keeps_event_loop_responsive(self):
from sglang.srt.utils.cuda_vmm_transport_utils import (
CudaVmmFeatureTransport,
)
transport = object.__new__(CudaVmmFeatureTransport)
transport.pool = object()
transport._publisher_executor = concurrent.futures.ThreadPoolExecutor(
max_workers=1
)
started = threading.Event()
event_loop_responsive = threading.Event()
release = threading.Event()
prepared = [object()]
def block_publication(_):
started.set()
release.wait()
return prepared
def unblock_if_event_loop_stalls():
started.wait()
if not event_loop_responsive.wait(timeout=0.5):
release.set()
transport.prepare_for_dispatch = MagicMock(side_effect=block_publication)
mm_inputs = SimpleNamespace(mm_items=[object()])
async def run():
watchdog = threading.Thread(target=unblock_if_event_loop_stalls)
watchdog.start()
try:
task = asyncio.create_task(
transport.prepare_for_dispatch_async([mm_inputs])
)
while not started.is_set():
await asyncio.sleep(0)
event_loop_responsive.set()
self.assertFalse(task.done())
release.set()
self.assertEqual(await task, prepared)
finally:
started.set()
event_loop_responsive.set()
release.set()
watchdog.join()
try:
asyncio.run(run())
finally:
transport._publisher_executor.shutdown(wait=True)
def test_text_only_publication_bypasses_publisher_executor(self):
from sglang.srt.utils.cuda_vmm_transport_utils import (
CudaVmmFeatureTransport,
)
transport = object.__new__(CudaVmmFeatureTransport)
transport.pool = object()
transport._publisher_executor = MagicMock()
result = asyncio.run(
transport.prepare_for_dispatch_async([None, SimpleNamespace(mm_items=[])])
)
self.assertEqual(result, [])
transport._publisher_executor.submit.assert_not_called()
def test_pageable_copy_source_is_pinned(self):
from sglang.srt.utils import cuda_vmm_transport_utils as vmm
source = torch.arange(4)
pinned = torch.arange(4)
with patch.object(
torch.Tensor, "pin_memory", autospec=True, return_value=pinned
) as pin_memory:
result = vmm._prepare_pinned_copy_source(source)
self.assertIs(result, pinned)
pin_memory.assert_called_once_with()
def test_packed_cpu_sources_use_one_pinned_staging_buffer(self):
from sglang.srt.utils import cuda_vmm_transport_utils as vmm
sources = [torch.arange(4, dtype=torch.int32), torch.arange(2)]
layouts, packed_data_nbytes = vmm._build_packed_tensor_layout(sources)
staging = torch.empty(packed_data_nbytes, dtype=torch.uint8)
with patch.object(vmm.torch, "empty", return_value=staging) as empty:
result = vmm._pack_pinned_copy_sources(sources, layouts, packed_data_nbytes)
self.assertIs(result, staging)
empty.assert_called_once_with(
packed_data_nbytes,
dtype=torch.uint8,
device="cpu",
pin_memory=True,
)
for source, layout in zip(sources, layouts, strict=True):
actual = staging[
layout.relative_offset : layout.relative_offset + layout.data_nbytes
]
self.assertTrue(torch.equal(actual, source.reshape(-1).view(torch.uint8)))
def test_recycler_polls_all_chunks_in_one_batch(self):
from sglang.srt.utils import cuda_vmm_transport_utils as vmm
pool = object.__new__(vmm.CudaVmmMemoryPool)
pool.device_index = 0
pool.consumer_count = 2
pool._recycle_stream = object()
pool.memory_pool = MagicMock()
pool.available_chunks = []
pool.occupied_chunks = [
vmm._CudaVmmMemoryChunk(0, 64),
vmm._CudaVmmMemoryChunk(64, 128),
]
acknowledgement_words = object()
acknowledgement_counts = MagicMock()
acknowledgement_counts.cpu.return_value.tolist.return_value = [2, 1]
with (
patch.object(vmm.torch.cuda, "device", return_value=nullcontext()),
patch.object(vmm.torch.cuda, "stream", return_value=nullcontext()),
patch.object(
vmm.torch, "stack", return_value=acknowledgement_words
) as stack,
patch.object(
vmm.torch,
"count_nonzero",
return_value=acknowledgement_counts,
) as count_nonzero,
):
pool._recycle_chunks()
stack.assert_called_once()
count_nonzero.assert_called_once_with(acknowledgement_words, dim=1)
self.assertEqual(pool.available_chunks, [vmm._CudaVmmMemoryChunk(0, 64)])
self.assertEqual(pool.occupied_chunks, [vmm._CudaVmmMemoryChunk(64, 128)])
def test_prepare_batch_cancels_prior_groups_on_failure(self): def test_prepare_batch_cancels_prior_groups_on_failure(self):
from sglang.srt.utils.cuda_vmm_transport_utils import ( from sglang.srt.utils.cuda_vmm_transport_utils import (
CudaVmmFeatureTransport, CudaVmmFeatureTransport,
@@ -495,6 +661,7 @@ class TestCudaVmmFeatureTransport(unittest.TestCase):
transport = object.__new__(CudaVmmFeatureTransport) transport = object.__new__(CudaVmmFeatureTransport)
pool = MagicMock() pool = MagicMock()
transport.pool = pool transport.pool = pool
transport._publisher_executor = None
manager.cuda_vmm_feature_transport = transport manager.cuda_vmm_feature_transport = transport
manager._subprocess_watchdog = None manager._subprocess_watchdog = None
engine = object.__new__(Engine) engine = object.__new__(Engine)
@@ -524,6 +691,7 @@ class TestCudaVmmFeatureTransport(unittest.TestCase):
transport = object.__new__(CudaVmmFeatureTransport) transport = object.__new__(CudaVmmFeatureTransport)
pool = MagicMock() pool = MagicMock()
transport.pool = pool transport.pool = pool
transport._publisher_executor = None
manager.cuda_vmm_feature_transport = transport manager.cuda_vmm_feature_transport = transport
# A real record: the launcher publishes it partway through, and what it # A real record: the launcher publishes it partway through, and what it
# reads after that comes out of the bags, which only project from a # reads after that comes out of the bags, which only project from a
@@ -576,6 +744,7 @@ class TestCudaVmmFeatureTransport(unittest.TestCase):
pool = MagicMock() pool = MagicMock()
pool.shutdown.side_effect = [RuntimeError("shutdown failed"), None] pool.shutdown.side_effect = [RuntimeError("shutdown failed"), None]
transport.pool = pool transport.pool = pool
transport._publisher_executor = None
with self.assertRaisesRegex(RuntimeError, "shutdown failed"): with self.assertRaisesRegex(RuntimeError, "shutdown failed"):
transport.shutdown() transport.shutdown()