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,
|
||||||
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),
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
Reference in New Issue
Block a user