[tokenizer] improve non streaming request processing + some small fixes. (#20310)

This commit is contained in:
Alex Nails
2026-04-10 15:46:12 -07:00
committed by GitHub
parent 451320596f
commit 0af9166474
4 changed files with 233 additions and 53 deletions
@@ -120,8 +120,9 @@ class AsyncDynamicbatchTokenizer:
) -> None:
"""Process a dynamic batch of encode requests for single string prompts."""
# Check if all kwargs are identical for efficient batch processing
can_batch = len(set(str(sorted(kw.items())) for kw in kwargs_list)) == 1
kwargs = kwargs_list[0] if can_batch else None
first_kw = kwargs_list[0]
can_batch = all(kw == first_kw for kw in kwargs_list[1:])
kwargs = first_kw if can_batch else None
try:
# If every request uses identical kwargs we can run a single
@@ -185,8 +185,9 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
# fast path
first_skip, first_space = skip_list[0], space_list[0]
if all(s == first_skip for s in skip_list) and all(
sp == first_space for sp in space_list
if all(
s == first_skip and sp == first_space
for s, sp in zip(skip_list, space_list)
):
return self.tokenizer.batch_decode(
ids_list,
@@ -294,8 +295,8 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
new_text = read_texts[i][len(surr_texts[i]) :]
if recv_obj.finished_reasons[i] is None:
# Streaming chunk: update the decode status
if len(new_text) > 0 and not new_text.endswith("�"):
s.decoded_text = s.decoded_text + new_text
if new_text and not new_text.endswith("�"):
s.decoded_text += new_text
s.surr_offset = s.read_offset
s.read_offset = len(s.decode_ids)
new_text = ""
+96 -18
View File
@@ -149,9 +149,32 @@ class ReqState:
last_output_offset: int = 0
last_text_offset: int = 0
# Buffer non-streaming text until the final response.
buffer_text: bool = False
text: str = ""
text_chunks: List[str] = dataclasses.field(default_factory=list)
def append_text(self, chunk: str):
if self.buffer_text:
self.text_chunks.append(chunk)
else:
self.text += chunk
def get_text(self) -> str:
if self.buffer_text:
return "".join(self.text_chunks)
return self.text
def get_crash_dump_output(self) -> Dict[Any, Any]:
out = {}
if self.text or self.text_chunks:
out["text"] = self.get_text()
if self.output_ids:
out["output_ids"] = self.output_ids.copy()
return out
# For incremental state update.
# TODO(lianmin): do not initialize some lists if not needed.
text: str = ""
output_ids: List[int] = dataclasses.field(default_factory=list)
input_token_logprobs_val: List[float] = dataclasses.field(default_factory=list)
input_token_logprobs_idx: List[int] = dataclasses.field(default_factory=list)
@@ -175,6 +198,24 @@ class ReqState:
output_token_ids_logprobs: List[Any] = dataclasses.field(default_factory=list)
def make_req_state(
out_list: List[Dict[Any, Any]],
finished: bool,
event: asyncio.Event,
obj: Union[GenerateReqInput, EmbeddingReqInput],
time_stats: APIServerReqTimeStats,
) -> ReqState:
is_streaming_request = getattr(obj, "stream", False)
return ReqState(
out_list,
finished,
event,
obj,
time_stats,
buffer_text=not is_streaming_request,
)
def _slice_streaming_output_meta_info(
meta_info: Dict[Any, Any],
last_output_offset: int,
@@ -1669,45 +1710,67 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
if getattr(recv_obj, "dp_ranks", None):
meta_info["dp_rank"] = recv_obj.dp_ranks[i]
state.finished = recv_obj.finished_reasons[i] is not None
if isinstance(recv_obj, BatchStrOutput):
state.text += recv_obj.output_strs[i]
# Not all request types have `stream` (e.g., EmbeddingReqInput). Default to non-streaming.
is_stream = getattr(state.obj, "stream", False)
if self.server_args.incremental_streaming_output and is_stream:
incremental = (
self.server_args.incremental_streaming_output and is_stream
)
output_offset = state.last_output_offset
state.append_text(recv_obj.output_strs[i])
state.output_ids.extend(recv_obj.output_ids[i])
if is_stream:
if incremental:
output_token_ids = state.output_ids[output_offset:]
_slice_streaming_output_meta_info(meta_info, output_offset)
state.last_output_offset = len(state.output_ids)
output_text = state.text[state.last_text_offset :]
state.last_text_offset = len(state.text)
text = state.get_text()
output_text = text[state.last_text_offset :]
state.last_text_offset = len(text)
else:
state.output_ids.extend(recv_obj.output_ids[i])
output_token_ids = state.output_ids.copy()
output_text = state.text
output_text = state.get_text()
out_dict = {
"text": output_text,
"output_ids": output_token_ids,
"meta_info": meta_info,
}
elif state.finished:
out_dict = {
"text": state.get_text(),
"output_ids": state.output_ids.copy(),
"meta_info": meta_info,
}
else:
out_dict = None
elif isinstance(recv_obj, BatchTokenIDOutput):
is_stream = getattr(state.obj, "stream", False)
if self.server_args.incremental_streaming_output and is_stream:
incremental = (
self.server_args.incremental_streaming_output and is_stream
)
output_offset = state.last_output_offset
state.output_ids.extend(recv_obj.output_ids[i])
if is_stream:
if incremental:
output_token_ids = state.output_ids[output_offset:]
_slice_streaming_output_meta_info(meta_info, output_offset)
state.last_output_offset = len(state.output_ids)
else:
state.output_ids.extend(recv_obj.output_ids[i])
output_token_ids = state.output_ids.copy()
out_dict = {
"output_ids": output_token_ids,
"meta_info": meta_info,
}
elif state.finished:
out_dict = {
"output_ids": state.output_ids.copy(),
"meta_info": meta_info,
}
else:
out_dict = None
else:
assert isinstance(recv_obj, BatchEmbeddingOutput)
out_dict = {
@@ -1715,8 +1778,6 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
"meta_info": meta_info,
}
state.finished = recv_obj.finished_reasons[i] is not None
# Set first_token_time on the first output batch.
# This is the single write point for first_token_time.
if state.time_stats.first_token_time == 0.0:
@@ -1754,6 +1815,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
if self.server_args.enable_lora and state.obj.lora_path:
asyncio.create_task(self.lora_registry.release(state.obj.lora_id))
if out_dict is not None:
state.out_list.append(out_dict)
state.event.set()
@@ -2167,7 +2229,11 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
unfinished_requests.append(
(
state.obj,
state.out_list[-1] if state.out_list else {},
(
state.out_list[-1]
if state.out_list
else state.get_crash_dump_output()
),
convert_time_to_realtime(state.time_stats.created_time),
convert_time_to_realtime(state.time_stats.finished_time),
)
@@ -2277,7 +2343,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
if is_stream:
output_ids = [output_ids[-1]] if len(output_ids) > 0 else []
out = {
"text": state.text,
"text": state.get_text(),
"output_ids": output_ids,
"meta_info": meta_info,
}
@@ -2397,7 +2463,13 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
if not hasattr(obj, "is_single") or obj.is_single:
time_stats = APIServerReqTimeStats(disagg_mode=self.disaggregation_mode)
state = ReqState([], False, asyncio.Event(), obj, time_stats)
state = make_req_state(
[],
False,
asyncio.Event(),
obj,
time_stats,
)
self.rid_to_state[obj.rid] = state
if self.server_args.enable_trace:
@@ -2413,7 +2485,13 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
else:
for i in range(len(obj.rid)):
time_stats = APIServerReqTimeStats(disagg_mode=self.disaggregation_mode)
state = ReqState([], False, asyncio.Event(), obj[i], time_stats)
state = make_req_state(
[],
False,
asyncio.Event(),
obj[i],
time_stats,
)
self.rid_to_state[obj.rid[i]] = state
if self.server_args.enable_trace:
+106 -6
View File
@@ -2,19 +2,30 @@
Unit tests for TokenizerManager helper methods.
This tests the refactored tokenization functionality including input format detection,
tokenizer input preparation, and result extraction logic.
tokenizer input preparation, result extraction logic, and ReqState text buffering.
Usage:
python3 -m unittest test_tokenizer_manager.TestInputFormatDetection
python3 -m unittest test_tokenizer_manager.TestTokenizerInputPreparation
python3 -m unittest test_tokenizer_manager.TestTokenizerResultExtraction
python3 -m unittest test_tokenizer_manager.TestTokenizerManagerIntegration
python3 -m unittest test_tokenizer_manager.TestReqStateTextBuffering
python3 -m unittest test_tokenizer_manager.TestReqStateCrashDump
python3 -m unittest test_tokenizer_manager.TestMakeReqState
"""
import asyncio
import unittest
from unittest.mock import Mock, patch
from sglang.srt.managers.tokenizer_manager import InputFormat, TokenizerManager
from sglang.srt.managers.io_struct import EmbeddingReqInput, GenerateReqInput
from sglang.srt.managers.tokenizer_manager import (
InputFormat,
ReqState,
TokenizerManager,
make_req_state,
)
from sglang.srt.observability.req_time_stats import APIServerReqTimeStats
from sglang.srt.server_args import PortArgs, ServerArgs
from sglang.test.test_utils import DEFAULT_SMALL_MODEL_NAME_FOR_TEST
@@ -29,7 +40,7 @@ class TestInputFormatDetection(unittest.TestCase):
self.port_args = PortArgs.init_new(self.server_args)
with patch("zmq.asyncio.Context"), patch(
"sglang.srt.utils.get_zmq_socket"
"sglang.srt.utils.network.get_zmq_socket"
), patch(
"sglang.srt.utils.hf_transformers_utils.get_tokenizer"
) as mock_tokenizer:
@@ -125,7 +136,7 @@ class TestTokenizerInputPreparation(unittest.TestCase):
self.port_args = PortArgs.init_new(self.server_args)
with patch("zmq.asyncio.Context"), patch(
"sglang.srt.utils.get_zmq_socket"
"sglang.srt.utils.network.get_zmq_socket"
), patch(
"sglang.srt.utils.hf_transformers_utils.get_tokenizer"
) as mock_tokenizer:
@@ -183,7 +194,7 @@ class TestTokenizerResultExtraction(unittest.TestCase):
self.port_args = PortArgs.init_new(self.server_args)
with patch("zmq.asyncio.Context"), patch(
"sglang.srt.utils.get_zmq_socket"
"sglang.srt.utils.network.get_zmq_socket"
), patch(
"sglang.srt.utils.hf_transformers_utils.get_tokenizer"
) as mock_tokenizer:
@@ -305,7 +316,7 @@ class TestTokenizerManagerIntegration(unittest.TestCase):
self.port_args = PortArgs.init_new(self.server_args)
with patch("zmq.asyncio.Context"), patch(
"sglang.srt.utils.get_zmq_socket"
"sglang.srt.utils.network.get_zmq_socket"
), patch(
"sglang.srt.utils.hf_transformers_utils.get_tokenizer"
) as mock_tokenizer:
@@ -404,5 +415,94 @@ class TestTokenizerManagerIntegration(unittest.TestCase):
self.assertIsNone(result_token_type_ids)
def _make_state(buffer_text: bool = False) -> ReqState:
"""Create a minimal ReqState for testing."""
obj = Mock(spec=GenerateReqInput)
return ReqState(
out_list=[],
finished=False,
event=asyncio.Event(),
obj=obj,
time_stats=APIServerReqTimeStats(),
buffer_text=buffer_text,
)
class TestReqStateTextBuffering(unittest.TestCase):
"""Test ReqState.append_text / get_text in both buffering modes."""
def test_streaming_mode_concatenates_directly(self):
state = _make_state(buffer_text=False)
state.append_text("hello ")
state.append_text("world")
self.assertEqual(state.get_text(), "hello world")
self.assertEqual(state.text_chunks, [])
def test_buffer_mode_collects_chunks(self):
state = _make_state(buffer_text=True)
state.append_text("hello ")
state.append_text("world")
self.assertEqual(state.text, "")
self.assertEqual(state.get_text(), "hello world")
class TestReqStateCrashDump(unittest.TestCase):
"""Test ReqState.get_crash_dump_output."""
def test_empty_state(self):
state = _make_state(buffer_text=False)
self.assertEqual(state.get_crash_dump_output(), {})
def test_with_text_only(self):
state = _make_state(buffer_text=False)
state.append_text("partial output")
self.assertEqual(state.get_crash_dump_output(), {"text": "partial output"})
def test_with_output_ids_only(self):
state = _make_state(buffer_text=False)
state.output_ids = [1, 2, 3]
self.assertEqual(state.get_crash_dump_output(), {"output_ids": [1, 2, 3]})
def test_with_text_and_output_ids(self):
state = _make_state(buffer_text=False)
state.append_text("hello")
state.output_ids = [10, 20]
self.assertEqual(
state.get_crash_dump_output(),
{"text": "hello", "output_ids": [10, 20]},
)
class TestMakeReqState(unittest.TestCase):
"""Test make_req_state factory function."""
def _call(self, *, obj_stream=None):
if obj_stream is not None:
obj = Mock(spec=GenerateReqInput)
obj.stream = obj_stream
else:
obj = Mock(spec=EmbeddingReqInput)
del obj.stream
return make_req_state(
out_list=[],
finished=False,
event=asyncio.Event(),
obj=obj,
time_stats=APIServerReqTimeStats(),
)
def test_streaming_request_does_not_buffer(self):
state = self._call(obj_stream=True)
self.assertFalse(state.buffer_text)
def test_non_streaming_request_buffers(self):
state = self._call(obj_stream=False)
self.assertTrue(state.buffer_text)
def test_embedding_request_always_buffers(self):
state = self._call()
self.assertTrue(state.buffer_text)
if __name__ == "__main__":
unittest.main(verbosity=2)