[HiCache] return cached_tokens_details in sglext for streaming responses (#22055)
Signed-off-by: Vladislav Nosivskoy <vladnosiv@gmail.com>
This commit is contained in:
@@ -775,6 +775,126 @@ class ServingChatTestCase(unittest.TestCase):
|
||||
self.assertEqual(len(chunks), 2)
|
||||
self.assertIn("error", chunks[0])
|
||||
|
||||
def test_non_streaming_cached_tokens_details_emits_sglext(self):
|
||||
"""Test that non-streaming chat responses emit cached token details in sglext."""
|
||||
|
||||
req = ChatCompletionRequest(
|
||||
model="x",
|
||||
messages=[{"role": "user", "content": "Hi?"}],
|
||||
max_tokens=100,
|
||||
return_cached_tokens_details=True,
|
||||
)
|
||||
ret = [
|
||||
{
|
||||
"text": "Cached response",
|
||||
"meta_info": {
|
||||
"id": "chatcmpl-cache-test",
|
||||
"prompt_tokens": 10,
|
||||
"completion_tokens": 2,
|
||||
"cached_tokens": 6,
|
||||
"cached_tokens_details": {
|
||||
"device": 4,
|
||||
"host": 1,
|
||||
"storage": 1,
|
||||
"storage_backend": "file",
|
||||
},
|
||||
"finish_reason": {"type": "stop", "matched": None},
|
||||
"weight_version": "default",
|
||||
},
|
||||
}
|
||||
]
|
||||
|
||||
response = self.chat._build_chat_response(req, ret, 1234567890)
|
||||
|
||||
self.assertIsNotNone(response.sglext)
|
||||
self.assertEqual(
|
||||
response.sglext.cached_tokens_details.model_dump(exclude_none=True),
|
||||
{
|
||||
"device": 4,
|
||||
"host": 1,
|
||||
"storage": 1,
|
||||
"storage_backend": "file",
|
||||
},
|
||||
)
|
||||
|
||||
def test_streaming_cached_tokens_details_emits_sglext(self):
|
||||
"""Test that streaming chat responses emit cached token details in sglext."""
|
||||
|
||||
async def _mock_generate_with_cached_tokens_details():
|
||||
yield {
|
||||
"text": "Cached response",
|
||||
"meta_info": {
|
||||
"id": "chatcmpl-cache-test",
|
||||
"prompt_tokens": 10,
|
||||
"completion_tokens": 2,
|
||||
"cached_tokens": 6,
|
||||
"cached_tokens_details": {
|
||||
"device": 4,
|
||||
"host": 1,
|
||||
"storage": 1,
|
||||
"storage_backend": "file",
|
||||
},
|
||||
"finish_reason": {"type": "stop", "matched": None},
|
||||
"output_token_logprobs": None,
|
||||
"output_top_logprobs": None,
|
||||
},
|
||||
"index": 0,
|
||||
}
|
||||
|
||||
self.tm.generate_request.return_value = (
|
||||
_mock_generate_with_cached_tokens_details()
|
||||
)
|
||||
|
||||
req = ChatCompletionRequest(
|
||||
model="x",
|
||||
messages=[{"role": "user", "content": "Hi?"}],
|
||||
max_tokens=100,
|
||||
stream=True,
|
||||
return_cached_tokens_details=True,
|
||||
)
|
||||
|
||||
with patch(
|
||||
"sglang.srt.entrypoints.openai.serving_chat.generate_chat_conv"
|
||||
) as conv_mock:
|
||||
conv_ins = Mock()
|
||||
conv_ins.get_prompt.return_value = "Test prompt"
|
||||
conv_mock.return_value = conv_ins
|
||||
|
||||
adapted_request, _ = self.chat._convert_to_internal_request(
|
||||
req, self.fastapi_request
|
||||
)
|
||||
|
||||
async def run_stream():
|
||||
chunks = []
|
||||
async for chunk in self.chat._generate_chat_stream(
|
||||
adapted_request, req, self.fastapi_request
|
||||
):
|
||||
chunks.append(chunk)
|
||||
return chunks
|
||||
|
||||
loop = get_or_create_event_loop()
|
||||
chunks = loop.run_until_complete(run_stream())
|
||||
|
||||
sglext_chunks = []
|
||||
for chunk in chunks:
|
||||
if not chunk.startswith("data: ") or chunk.strip() == "data: [DONE]":
|
||||
continue
|
||||
data = json.loads(chunk[len("data: ") :])
|
||||
if "sglext" in data:
|
||||
sglext_chunks.append(data)
|
||||
|
||||
self.assertEqual(len(sglext_chunks), 1)
|
||||
self.assertEqual(sglext_chunks[0]["choices"], [])
|
||||
self.assertEqual(
|
||||
sglext_chunks[0]["sglext"]["cached_tokens_details"],
|
||||
{
|
||||
"device": 4,
|
||||
"host": 1,
|
||||
"storage": 1,
|
||||
"storage_backend": "file",
|
||||
},
|
||||
)
|
||||
|
||||
# ------------- incremental streaming output tests -------------
|
||||
def test_incremental_streaming_output_delta(self):
|
||||
"""Test that streaming with incremental_streaming_output produces correct deltas.
|
||||
|
||||
@@ -256,6 +256,117 @@ class ServingCompletionTestCase(unittest.TestCase):
|
||||
self.assertGreaterEqual(len(chunks), 2)
|
||||
self.assertIn("error", chunks[0])
|
||||
|
||||
def test_non_streaming_cached_tokens_details_emits_sglext(self):
|
||||
"""Test that non-streaming completion responses emit cached token details in sglext."""
|
||||
|
||||
req = CompletionRequest(
|
||||
model="x",
|
||||
prompt="Hello world",
|
||||
max_tokens=100,
|
||||
return_cached_tokens_details=True,
|
||||
)
|
||||
ret = [
|
||||
{
|
||||
"text": "Cached response",
|
||||
"meta_info": {
|
||||
"id": "cmpl-cache-test",
|
||||
"prompt_tokens": 10,
|
||||
"completion_tokens": 2,
|
||||
"cached_tokens": 6,
|
||||
"cached_tokens_details": {
|
||||
"device": 4,
|
||||
"host": 1,
|
||||
"storage": 1,
|
||||
"storage_backend": "file",
|
||||
},
|
||||
"finish_reason": {"type": "stop", "matched": None},
|
||||
"weight_version": "default",
|
||||
},
|
||||
}
|
||||
]
|
||||
|
||||
response = self.sc._build_completion_response(req, ret, 1234567890)
|
||||
|
||||
self.assertIsNotNone(response.sglext)
|
||||
self.assertEqual(
|
||||
response.sglext.cached_tokens_details.model_dump(exclude_none=True),
|
||||
{
|
||||
"device": 4,
|
||||
"host": 1,
|
||||
"storage": 1,
|
||||
"storage_backend": "file",
|
||||
},
|
||||
)
|
||||
|
||||
def test_streaming_cached_tokens_details_emits_sglext(self):
|
||||
"""Test that streaming completion responses emit cached token details in sglext."""
|
||||
|
||||
async def _mock_generate_with_cached_tokens_details(*args, **kwargs):
|
||||
yield {
|
||||
"text": "Cached response",
|
||||
"meta_info": {
|
||||
"id": "cmpl-cache-test",
|
||||
"prompt_tokens": 10,
|
||||
"completion_tokens": 2,
|
||||
"cached_tokens": 6,
|
||||
"cached_tokens_details": {
|
||||
"device": 4,
|
||||
"host": 1,
|
||||
"storage": 1,
|
||||
"storage_backend": "file",
|
||||
},
|
||||
"finish_reason": {"type": "stop", "matched": None},
|
||||
"output_token_logprobs": None,
|
||||
"output_top_logprobs": None,
|
||||
},
|
||||
"index": 0,
|
||||
}
|
||||
|
||||
self.sc.tokenizer_manager.generate_request = (
|
||||
_mock_generate_with_cached_tokens_details
|
||||
)
|
||||
|
||||
req = CompletionRequest(
|
||||
model="x",
|
||||
prompt="Hello world",
|
||||
max_tokens=100,
|
||||
stream=True,
|
||||
return_cached_tokens_details=True,
|
||||
)
|
||||
|
||||
adapted_request, _ = self.sc._convert_to_internal_request(req)
|
||||
|
||||
async def run_stream():
|
||||
chunks = []
|
||||
async for chunk in self.sc._generate_completion_stream(
|
||||
adapted_request, req, self.fastapi_request
|
||||
):
|
||||
chunks.append(chunk)
|
||||
return chunks
|
||||
|
||||
loop = get_or_create_event_loop()
|
||||
chunks = loop.run_until_complete(run_stream())
|
||||
|
||||
sglext_chunks = []
|
||||
for chunk in chunks:
|
||||
if not chunk.startswith("data: ") or chunk.strip() == "data: [DONE]":
|
||||
continue
|
||||
data = json.loads(chunk[len("data: ") :])
|
||||
if "sglext" in data:
|
||||
sglext_chunks.append(data)
|
||||
|
||||
self.assertEqual(len(sglext_chunks), 1)
|
||||
self.assertEqual(sglext_chunks[0]["choices"], [])
|
||||
self.assertEqual(
|
||||
sglext_chunks[0]["sglext"]["cached_tokens_details"],
|
||||
{
|
||||
"device": 4,
|
||||
"host": 1,
|
||||
"storage": 1,
|
||||
"storage_backend": "file",
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main(verbosity=2)
|
||||
|
||||
Reference in New Issue
Block a user