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),
@@ -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()