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:
co-authored by
Shiyan Deng
Yinghai Lu
parent
f15748d965
commit
ff04a00d73
@@ -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),
|
||||
|
||||
@@ -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,16 +1554,18 @@ 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(
|
||||
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)
|
||||
time_stats = tokenized_obj.time_stats
|
||||
@@ -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,9 +1588,11 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
prepared_mm_items = []
|
||||
dispatched = False
|
||||
try:
|
||||
prepared_mm_items = self.cuda_vmm_feature_transport.prepare_for_dispatch(
|
||||
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")
|
||||
time_stats = [tokenized_obj.time_stats for tokenized_obj in tokenized_objs]
|
||||
@@ -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,17 +477,27 @@ 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):
|
||||
if packed_source is not None:
|
||||
self.memory_pool[
|
||||
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
|
||||
@@ -570,21 +613,32 @@ 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()
|
||||
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
|
||||
]
|
||||
)
|
||||
if ack_count == self.consumer_count:
|
||||
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)
|
||||
@@ -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()
|
||||
|
||||
@@ -9,7 +9,7 @@ import pickle
|
||||
import queue
|
||||
import threading
|
||||
import unittest
|
||||
from unittest.mock import MagicMock, patch
|
||||
from unittest.mock import patch
|
||||
|
||||
import torch
|
||||
|
||||
@@ -184,24 +184,25 @@ class TestCudaVmmTransport(CustomTestCase):
|
||||
finally:
|
||||
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)
|
||||
sources = [
|
||||
torch.arange(24, dtype=torch.float32, device="cuda:0")
|
||||
.reshape(4, 6)
|
||||
.transpose(0, 1),
|
||||
torch.arange(24, dtype=torch.float32).reshape(4, 6).transpose(0, 1),
|
||||
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]
|
||||
proxies = reconstructed = None
|
||||
try:
|
||||
stream = MagicMock(wraps=torch.cuda.current_stream(0))
|
||||
with patch("torch.cuda.current_stream", return_value=stream):
|
||||
with patch.object(
|
||||
pool._publish_stream,
|
||||
"synchronize",
|
||||
wraps=pool._publish_stream.synchronize,
|
||||
) as synchronize:
|
||||
proxies = pool.wrap_tensors(sources)
|
||||
|
||||
self.assertIsNotNone(proxies)
|
||||
self.assertEqual(stream.synchronize.call_count, 1)
|
||||
self.assertEqual(synchronize.call_count, 1)
|
||||
self.assertEqual(len(pool.occupied_chunks), 1)
|
||||
self.assertTrue(
|
||||
all(
|
||||
@@ -337,11 +338,13 @@ class TestCudaVmmTransport(CustomTestCase):
|
||||
|
||||
def test_failed_cleanup_sync_quarantines_pool(self):
|
||||
pool = CudaVmmMemoryPool(4 << 20, 60, 0, 1, allow_posix_fallback=True)
|
||||
stream = MagicMock()
|
||||
stream.synchronize.side_effect = RuntimeError("forced sync failure")
|
||||
try:
|
||||
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"),
|
||||
):
|
||||
pool.wrap_tensor(torch.ones(16, device="cuda:0"))
|
||||
@@ -359,20 +362,21 @@ class TestCudaVmmTransport(CustomTestCase):
|
||||
shutdown_entered = threading.Event()
|
||||
shutdown_finished = threading.Event()
|
||||
errors = []
|
||||
real_stream = torch.cuda.current_stream(0)
|
||||
stream = MagicMock(wraps=real_stream)
|
||||
real_synchronize = pool._publish_stream.synchronize
|
||||
|
||||
def synchronize():
|
||||
publisher_entered.set()
|
||||
if not allow_publisher_to_finish.wait(timeout=10):
|
||||
raise TimeoutError("publisher was not released")
|
||||
real_stream.synchronize()
|
||||
|
||||
stream.synchronize.side_effect = synchronize
|
||||
real_synchronize()
|
||||
|
||||
def publish():
|
||||
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"))
|
||||
except Exception as error: # pragma: no cover
|
||||
errors.append(error)
|
||||
|
||||
@@ -1,7 +1,10 @@
|
||||
import asyncio
|
||||
import concurrent.futures
|
||||
import threading
|
||||
import unittest
|
||||
from contextlib import nullcontext
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock, call, patch
|
||||
from unittest.mock import AsyncMock, MagicMock, call, patch
|
||||
|
||||
import torch
|
||||
|
||||
@@ -214,6 +217,7 @@ class TestCudaVmmFeatureTransport(unittest.TestCase):
|
||||
transport = CudaVmmFeatureTransport(SimpleNamespace(), None)
|
||||
|
||||
self.assertEqual(transport.prepare_for_dispatch([None]), [])
|
||||
self.assertEqual(asyncio.run(transport.prepare_for_dispatch_async([None])), [])
|
||||
transport.cancel_for_dispatch([])
|
||||
transport.shutdown()
|
||||
self.assertIsNone(transport.pool)
|
||||
@@ -252,6 +256,28 @@ class TestCudaVmmFeatureTransport(unittest.TestCase):
|
||||
transport.pool.wrap_tensor.assert_not_called()
|
||||
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):
|
||||
from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem
|
||||
from sglang.srt.utils.cuda_ipc_transport_utils import (
|
||||
@@ -351,7 +377,7 @@ class TestCudaVmmFeatureTransport(unittest.TestCase):
|
||||
|
||||
manager = object.__new__(TokenizerManager)
|
||||
transport = MagicMock()
|
||||
transport.prepare_for_dispatch.return_value = []
|
||||
transport.prepare_for_dispatch_async = AsyncMock(return_value=[])
|
||||
manager.cuda_vmm_feature_transport = transport
|
||||
manager._dispatch_to_scheduler = MagicMock()
|
||||
tokenized_obj = SimpleNamespace(
|
||||
@@ -362,10 +388,10 @@ class TestCudaVmmFeatureTransport(unittest.TestCase):
|
||||
)
|
||||
|
||||
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)
|
||||
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()
|
||||
|
||||
def test_failed_dispatch_cancels_published_items(self):
|
||||
@@ -388,16 +414,16 @@ class TestCudaVmmFeatureTransport(unittest.TestCase):
|
||||
time_stats=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
|
||||
|
||||
with (
|
||||
patch.object(tokenizer_manager, "wrap_shm_features", lambda obj: obj),
|
||||
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,)
|
||||
)
|
||||
transport.cancel_for_dispatch.assert_called_once_with(items)
|
||||
@@ -424,18 +450,158 @@ class TestCudaVmmFeatureTransport(unittest.TestCase):
|
||||
time_stats=time_stats,
|
||||
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
|
||||
|
||||
with (
|
||||
patch.object(tokenizer_manager, "wrap_shm_features", lambda obj: obj),
|
||||
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)
|
||||
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):
|
||||
from sglang.srt.utils.cuda_vmm_transport_utils import (
|
||||
CudaVmmFeatureTransport,
|
||||
@@ -495,6 +661,7 @@ class TestCudaVmmFeatureTransport(unittest.TestCase):
|
||||
transport = object.__new__(CudaVmmFeatureTransport)
|
||||
pool = MagicMock()
|
||||
transport.pool = pool
|
||||
transport._publisher_executor = None
|
||||
manager.cuda_vmm_feature_transport = transport
|
||||
manager._subprocess_watchdog = None
|
||||
engine = object.__new__(Engine)
|
||||
@@ -524,6 +691,7 @@ class TestCudaVmmFeatureTransport(unittest.TestCase):
|
||||
transport = object.__new__(CudaVmmFeatureTransport)
|
||||
pool = MagicMock()
|
||||
transport.pool = pool
|
||||
transport._publisher_executor = None
|
||||
manager.cuda_vmm_feature_transport = transport
|
||||
# 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
|
||||
@@ -576,6 +744,7 @@ class TestCudaVmmFeatureTransport(unittest.TestCase):
|
||||
pool = MagicMock()
|
||||
pool.shutdown.side_effect = [RuntimeError("shutdown failed"), None]
|
||||
transport.pool = pool
|
||||
transport._publisher_executor = None
|
||||
|
||||
with self.assertRaisesRegex(RuntimeError, "shutdown failed"):
|
||||
transport.shutdown()
|
||||
|
||||
Reference in New Issue
Block a user