[tokenizer] improve non streaming request processing + some small fixes. (#20310)
This commit is contained in:
@@ -120,8 +120,9 @@ class AsyncDynamicbatchTokenizer:
|
|||||||
) -> None:
|
) -> None:
|
||||||
"""Process a dynamic batch of encode requests for single string prompts."""
|
"""Process a dynamic batch of encode requests for single string prompts."""
|
||||||
# Check if all kwargs are identical for efficient batch processing
|
# Check if all kwargs are identical for efficient batch processing
|
||||||
can_batch = len(set(str(sorted(kw.items())) for kw in kwargs_list)) == 1
|
first_kw = kwargs_list[0]
|
||||||
kwargs = kwargs_list[0] if can_batch else None
|
can_batch = all(kw == first_kw for kw in kwargs_list[1:])
|
||||||
|
kwargs = first_kw if can_batch else None
|
||||||
|
|
||||||
try:
|
try:
|
||||||
# If every request uses identical kwargs we can run a single
|
# If every request uses identical kwargs we can run a single
|
||||||
|
|||||||
@@ -185,8 +185,9 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
|
|||||||
|
|
||||||
# fast path
|
# fast path
|
||||||
first_skip, first_space = skip_list[0], space_list[0]
|
first_skip, first_space = skip_list[0], space_list[0]
|
||||||
if all(s == first_skip for s in skip_list) and all(
|
if all(
|
||||||
sp == first_space for sp in space_list
|
s == first_skip and sp == first_space
|
||||||
|
for s, sp in zip(skip_list, space_list)
|
||||||
):
|
):
|
||||||
return self.tokenizer.batch_decode(
|
return self.tokenizer.batch_decode(
|
||||||
ids_list,
|
ids_list,
|
||||||
@@ -294,8 +295,8 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
|
|||||||
new_text = read_texts[i][len(surr_texts[i]) :]
|
new_text = read_texts[i][len(surr_texts[i]) :]
|
||||||
if recv_obj.finished_reasons[i] is None:
|
if recv_obj.finished_reasons[i] is None:
|
||||||
# Streaming chunk: update the decode status
|
# Streaming chunk: update the decode status
|
||||||
if len(new_text) > 0 and not new_text.endswith("�"):
|
if new_text and not new_text.endswith("�"):
|
||||||
s.decoded_text = s.decoded_text + new_text
|
s.decoded_text += new_text
|
||||||
s.surr_offset = s.read_offset
|
s.surr_offset = s.read_offset
|
||||||
s.read_offset = len(s.decode_ids)
|
s.read_offset = len(s.decode_ids)
|
||||||
new_text = ""
|
new_text = ""
|
||||||
|
|||||||
@@ -149,9 +149,32 @@ class ReqState:
|
|||||||
last_output_offset: int = 0
|
last_output_offset: int = 0
|
||||||
last_text_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.
|
# For incremental state update.
|
||||||
# TODO(lianmin): do not initialize some lists if not needed.
|
# TODO(lianmin): do not initialize some lists if not needed.
|
||||||
text: str = ""
|
|
||||||
output_ids: List[int] = dataclasses.field(default_factory=list)
|
output_ids: List[int] = dataclasses.field(default_factory=list)
|
||||||
input_token_logprobs_val: List[float] = 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)
|
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)
|
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(
|
def _slice_streaming_output_meta_info(
|
||||||
meta_info: Dict[Any, Any],
|
meta_info: Dict[Any, Any],
|
||||||
last_output_offset: int,
|
last_output_offset: int,
|
||||||
@@ -1669,45 +1710,67 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
|
|||||||
if getattr(recv_obj, "dp_ranks", None):
|
if getattr(recv_obj, "dp_ranks", None):
|
||||||
meta_info["dp_rank"] = recv_obj.dp_ranks[i]
|
meta_info["dp_rank"] = recv_obj.dp_ranks[i]
|
||||||
|
|
||||||
|
state.finished = recv_obj.finished_reasons[i] is not None
|
||||||
if isinstance(recv_obj, BatchStrOutput):
|
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.
|
# Not all request types have `stream` (e.g., EmbeddingReqInput). Default to non-streaming.
|
||||||
is_stream = getattr(state.obj, "stream", False)
|
is_stream = getattr(state.obj, "stream", False)
|
||||||
if self.server_args.incremental_streaming_output and is_stream:
|
incremental = (
|
||||||
output_offset = state.last_output_offset
|
self.server_args.incremental_streaming_output and is_stream
|
||||||
state.output_ids.extend(recv_obj.output_ids[i])
|
)
|
||||||
output_token_ids = state.output_ids[output_offset:]
|
output_offset = state.last_output_offset
|
||||||
_slice_streaming_output_meta_info(meta_info, output_offset)
|
state.append_text(recv_obj.output_strs[i])
|
||||||
state.last_output_offset = len(state.output_ids)
|
state.output_ids.extend(recv_obj.output_ids[i])
|
||||||
output_text = state.text[state.last_text_offset :]
|
|
||||||
state.last_text_offset = len(state.text)
|
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)
|
||||||
|
text = state.get_text()
|
||||||
|
output_text = text[state.last_text_offset :]
|
||||||
|
state.last_text_offset = len(text)
|
||||||
|
else:
|
||||||
|
output_token_ids = state.output_ids.copy()
|
||||||
|
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:
|
else:
|
||||||
state.output_ids.extend(recv_obj.output_ids[i])
|
out_dict = None
|
||||||
output_token_ids = state.output_ids.copy()
|
|
||||||
output_text = state.text
|
|
||||||
|
|
||||||
out_dict = {
|
|
||||||
"text": output_text,
|
|
||||||
"output_ids": output_token_ids,
|
|
||||||
"meta_info": meta_info,
|
|
||||||
}
|
|
||||||
|
|
||||||
elif isinstance(recv_obj, BatchTokenIDOutput):
|
elif isinstance(recv_obj, BatchTokenIDOutput):
|
||||||
is_stream = getattr(state.obj, "stream", False)
|
is_stream = getattr(state.obj, "stream", False)
|
||||||
if self.server_args.incremental_streaming_output and is_stream:
|
incremental = (
|
||||||
output_offset = state.last_output_offset
|
self.server_args.incremental_streaming_output and is_stream
|
||||||
state.output_ids.extend(recv_obj.output_ids[i])
|
)
|
||||||
output_token_ids = state.output_ids[output_offset:]
|
output_offset = state.last_output_offset
|
||||||
_slice_streaming_output_meta_info(meta_info, output_offset)
|
state.output_ids.extend(recv_obj.output_ids[i])
|
||||||
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 = {
|
if is_stream:
|
||||||
"output_ids": output_token_ids,
|
if incremental:
|
||||||
"meta_info": meta_info,
|
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:
|
||||||
|
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:
|
else:
|
||||||
assert isinstance(recv_obj, BatchEmbeddingOutput)
|
assert isinstance(recv_obj, BatchEmbeddingOutput)
|
||||||
out_dict = {
|
out_dict = {
|
||||||
@@ -1715,8 +1778,6 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
|
|||||||
"meta_info": meta_info,
|
"meta_info": meta_info,
|
||||||
}
|
}
|
||||||
|
|
||||||
state.finished = recv_obj.finished_reasons[i] is not None
|
|
||||||
|
|
||||||
# Set first_token_time on the first output batch.
|
# Set first_token_time on the first output batch.
|
||||||
# This is the single write point for first_token_time.
|
# This is the single write point for first_token_time.
|
||||||
if state.time_stats.first_token_time == 0.0:
|
if state.time_stats.first_token_time == 0.0:
|
||||||
@@ -1754,8 +1815,9 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
|
|||||||
if self.server_args.enable_lora and state.obj.lora_path:
|
if self.server_args.enable_lora and state.obj.lora_path:
|
||||||
asyncio.create_task(self.lora_registry.release(state.obj.lora_id))
|
asyncio.create_task(self.lora_registry.release(state.obj.lora_id))
|
||||||
|
|
||||||
state.out_list.append(out_dict)
|
if out_dict is not None:
|
||||||
state.event.set()
|
state.out_list.append(out_dict)
|
||||||
|
state.event.set()
|
||||||
|
|
||||||
# Log metrics and dump
|
# Log metrics and dump
|
||||||
if self.enable_metrics and state.obj.log_metrics:
|
if self.enable_metrics and state.obj.log_metrics:
|
||||||
@@ -2167,7 +2229,11 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
|
|||||||
unfinished_requests.append(
|
unfinished_requests.append(
|
||||||
(
|
(
|
||||||
state.obj,
|
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.created_time),
|
||||||
convert_time_to_realtime(state.time_stats.finished_time),
|
convert_time_to_realtime(state.time_stats.finished_time),
|
||||||
)
|
)
|
||||||
@@ -2277,7 +2343,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
|
|||||||
if is_stream:
|
if is_stream:
|
||||||
output_ids = [output_ids[-1]] if len(output_ids) > 0 else []
|
output_ids = [output_ids[-1]] if len(output_ids) > 0 else []
|
||||||
out = {
|
out = {
|
||||||
"text": state.text,
|
"text": state.get_text(),
|
||||||
"output_ids": output_ids,
|
"output_ids": output_ids,
|
||||||
"meta_info": meta_info,
|
"meta_info": meta_info,
|
||||||
}
|
}
|
||||||
@@ -2397,7 +2463,13 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
|
|||||||
|
|
||||||
if not hasattr(obj, "is_single") or obj.is_single:
|
if not hasattr(obj, "is_single") or obj.is_single:
|
||||||
time_stats = APIServerReqTimeStats(disagg_mode=self.disaggregation_mode)
|
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
|
self.rid_to_state[obj.rid] = state
|
||||||
|
|
||||||
if self.server_args.enable_trace:
|
if self.server_args.enable_trace:
|
||||||
@@ -2413,7 +2485,13 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
|
|||||||
else:
|
else:
|
||||||
for i in range(len(obj.rid)):
|
for i in range(len(obj.rid)):
|
||||||
time_stats = APIServerReqTimeStats(disagg_mode=self.disaggregation_mode)
|
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
|
self.rid_to_state[obj.rid[i]] = state
|
||||||
|
|
||||||
if self.server_args.enable_trace:
|
if self.server_args.enable_trace:
|
||||||
|
|||||||
@@ -2,19 +2,30 @@
|
|||||||
Unit tests for TokenizerManager helper methods.
|
Unit tests for TokenizerManager helper methods.
|
||||||
|
|
||||||
This tests the refactored tokenization functionality including input format detection,
|
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:
|
Usage:
|
||||||
python3 -m unittest test_tokenizer_manager.TestInputFormatDetection
|
python3 -m unittest test_tokenizer_manager.TestInputFormatDetection
|
||||||
python3 -m unittest test_tokenizer_manager.TestTokenizerInputPreparation
|
python3 -m unittest test_tokenizer_manager.TestTokenizerInputPreparation
|
||||||
python3 -m unittest test_tokenizer_manager.TestTokenizerResultExtraction
|
python3 -m unittest test_tokenizer_manager.TestTokenizerResultExtraction
|
||||||
python3 -m unittest test_tokenizer_manager.TestTokenizerManagerIntegration
|
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
|
import unittest
|
||||||
from unittest.mock import Mock, patch
|
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.srt.server_args import PortArgs, ServerArgs
|
||||||
from sglang.test.test_utils import DEFAULT_SMALL_MODEL_NAME_FOR_TEST
|
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)
|
self.port_args = PortArgs.init_new(self.server_args)
|
||||||
|
|
||||||
with patch("zmq.asyncio.Context"), patch(
|
with patch("zmq.asyncio.Context"), patch(
|
||||||
"sglang.srt.utils.get_zmq_socket"
|
"sglang.srt.utils.network.get_zmq_socket"
|
||||||
), patch(
|
), patch(
|
||||||
"sglang.srt.utils.hf_transformers_utils.get_tokenizer"
|
"sglang.srt.utils.hf_transformers_utils.get_tokenizer"
|
||||||
) as mock_tokenizer:
|
) as mock_tokenizer:
|
||||||
@@ -125,7 +136,7 @@ class TestTokenizerInputPreparation(unittest.TestCase):
|
|||||||
self.port_args = PortArgs.init_new(self.server_args)
|
self.port_args = PortArgs.init_new(self.server_args)
|
||||||
|
|
||||||
with patch("zmq.asyncio.Context"), patch(
|
with patch("zmq.asyncio.Context"), patch(
|
||||||
"sglang.srt.utils.get_zmq_socket"
|
"sglang.srt.utils.network.get_zmq_socket"
|
||||||
), patch(
|
), patch(
|
||||||
"sglang.srt.utils.hf_transformers_utils.get_tokenizer"
|
"sglang.srt.utils.hf_transformers_utils.get_tokenizer"
|
||||||
) as mock_tokenizer:
|
) as mock_tokenizer:
|
||||||
@@ -183,7 +194,7 @@ class TestTokenizerResultExtraction(unittest.TestCase):
|
|||||||
self.port_args = PortArgs.init_new(self.server_args)
|
self.port_args = PortArgs.init_new(self.server_args)
|
||||||
|
|
||||||
with patch("zmq.asyncio.Context"), patch(
|
with patch("zmq.asyncio.Context"), patch(
|
||||||
"sglang.srt.utils.get_zmq_socket"
|
"sglang.srt.utils.network.get_zmq_socket"
|
||||||
), patch(
|
), patch(
|
||||||
"sglang.srt.utils.hf_transformers_utils.get_tokenizer"
|
"sglang.srt.utils.hf_transformers_utils.get_tokenizer"
|
||||||
) as mock_tokenizer:
|
) as mock_tokenizer:
|
||||||
@@ -305,7 +316,7 @@ class TestTokenizerManagerIntegration(unittest.TestCase):
|
|||||||
self.port_args = PortArgs.init_new(self.server_args)
|
self.port_args = PortArgs.init_new(self.server_args)
|
||||||
|
|
||||||
with patch("zmq.asyncio.Context"), patch(
|
with patch("zmq.asyncio.Context"), patch(
|
||||||
"sglang.srt.utils.get_zmq_socket"
|
"sglang.srt.utils.network.get_zmq_socket"
|
||||||
), patch(
|
), patch(
|
||||||
"sglang.srt.utils.hf_transformers_utils.get_tokenizer"
|
"sglang.srt.utils.hf_transformers_utils.get_tokenizer"
|
||||||
) as mock_tokenizer:
|
) as mock_tokenizer:
|
||||||
@@ -404,5 +415,94 @@ class TestTokenizerManagerIntegration(unittest.TestCase):
|
|||||||
self.assertIsNone(result_token_type_ids)
|
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__":
|
if __name__ == "__main__":
|
||||||
unittest.main(verbosity=2)
|
unittest.main(verbosity=2)
|
||||||
|
|||||||
Reference in New Issue
Block a user