diff --git a/python/sglang/srt/constrained/llguidance_backend.py b/python/sglang/srt/constrained/llguidance_backend.py index dbc223c62..1f5bb72ad 100644 --- a/python/sglang/srt/constrained/llguidance_backend.py +++ b/python/sglang/srt/constrained/llguidance_backend.py @@ -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), diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py index 2acc9efd0..2d61e1417 100644 --- a/python/sglang/srt/managers/tokenizer_manager.py +++ b/python/sglang/srt/managers/tokenizer_manager.py @@ -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) diff --git a/python/sglang/srt/utils/cuda_vmm_transport_utils.py b/python/sglang/srt/utils/cuda_vmm_transport_utils.py index b81805de4..67c5ee1a7 100644 --- a/python/sglang/srt/utils/cuda_vmm_transport_utils.py +++ b/python/sglang/srt/utils/cuda_vmm_transport_utils.py @@ -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() diff --git a/test/registered/unit/multimodal/test_cuda_vmm_transport.py b/test/registered/unit/multimodal/test_cuda_vmm_transport.py index e3531a658..235f63e7e 100644 --- a/test/registered/unit/multimodal/test_cuda_vmm_transport.py +++ b/test/registered/unit/multimodal/test_cuda_vmm_transport.py @@ -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) diff --git a/test/registered/unit/multimodal/test_gpu_feature_transport.py b/test/registered/unit/multimodal/test_gpu_feature_transport.py index 6504ba075..1c488832f 100644 --- a/test/registered/unit/multimodal/test_gpu_feature_transport.py +++ b/test/registered/unit/multimodal/test_gpu_feature_transport.py @@ -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()