Align incremental streaming logprobs with streamed output tokens (#21583)

This commit is contained in:
Aurick Qiao
2026-04-06 00:30:02 -07:00
committed by GitHub
parent 12272b6791
commit 3178f3959f
6 changed files with 156 additions and 18 deletions
+38 -1
View File
@@ -247,6 +247,30 @@ def cache_program(program, backend):
backend.cache_prefix(prefix) backend.cache_prefix(prefix)
_INCREMENTAL_STREAMING_META_INFO_KEYS = (
"output_token_logprobs",
"output_top_logprobs",
"output_token_ids_logprobs",
)
def _merge_stream_meta_info(
pending_meta_info: dict[str, Any] | None,
meta_info: dict[str, Any],
) -> dict[str, Any]:
if pending_meta_info is None:
return meta_info
merged_meta_info = dict(meta_info)
for key in _INCREMENTAL_STREAMING_META_INFO_KEYS:
if key not in meta_info and key not in pending_meta_info:
continue
merged_meta_info[key] = list(pending_meta_info.get(key, [])) + list(
meta_info.get(key, [])
)
return merged_meta_info
class StreamExecutor: class StreamExecutor:
"""A stream executor that executes SGL expressions in a background thread.""" """A stream executor that executes SGL expressions in a background thread."""
@@ -949,6 +973,7 @@ class ProgramState:
break break
else: else:
event = None event = None
pending_meta_info = None
while not event: while not event:
if var_name in self.stream_executor.stream_var_event: if var_name in self.stream_executor.stream_var_event:
event = self.stream_executor.stream_var_event[var_name] event = self.stream_executor.stream_var_event[var_name]
@@ -960,12 +985,24 @@ class ProgramState:
await loop.run_in_executor(None, event.wait) await loop.run_in_executor(None, event.wait)
event.clear() event.clear()
out = str(self.stream_executor.variables[var_name][prev:]) out = str(self.stream_executor.variables[var_name][prev:])
meta_info = self.stream_executor.meta_info.get(var_name)
prev += len(out) prev += len(out)
if out: if out:
if return_meta_data: if return_meta_data:
yield out, self.stream_executor.meta_info[var_name] assert meta_info is not None
merged_meta_info = _merge_stream_meta_info(
pending_meta_info,
meta_info,
)
pending_meta_info = None
yield out, merged_meta_info
else: else:
yield out yield out
elif return_meta_data and meta_info is not None:
pending_meta_info = _merge_stream_meta_info(
pending_meta_info,
meta_info,
)
if self.stream_executor.variable_event[var_name].is_set(): if self.stream_executor.variable_event[var_name].is_set():
break break
else: else:
@@ -1207,13 +1207,18 @@ class OpenAIServingChat(OpenAIServingBase):
total_output_logprobs: int, total_output_logprobs: int,
) -> ChoiceLogprobs: ) -> ChoiceLogprobs:
"""Process logprobs for streaming response""" """Process logprobs for streaming response"""
output_token_logprobs = content["meta_info"]["output_token_logprobs"]
output_top_logprobs = content["meta_info"].get("output_top_logprobs", [])
if not self.tokenizer_manager.server_args.incremental_streaming_output:
output_token_logprobs = output_token_logprobs[
n_prev_token:total_output_logprobs
]
output_top_logprobs = output_top_logprobs[
n_prev_token:total_output_logprobs
]
logprobs = to_openai_style_logprobs( logprobs = to_openai_style_logprobs(
output_token_logprobs=content["meta_info"]["output_token_logprobs"][ output_token_logprobs=output_token_logprobs,
n_prev_token:total_output_logprobs output_top_logprobs=output_top_logprobs,
],
output_top_logprobs=content["meta_info"].get("output_top_logprobs", [])[
n_prev_token:total_output_logprobs
],
) )
token_logprobs = self._process_logprobs_tokens(logprobs, use_token_index=False) token_logprobs = self._process_logprobs_tokens(logprobs, use_token_index=False)
@@ -277,15 +277,26 @@ class OpenAIServingCompletion(OpenAIServingBase):
n_prev_token < total_output_logprobs n_prev_token < total_output_logprobs
or input_token_logprobs is not None or input_token_logprobs is not None
): ):
output_token_logprobs = content["meta_info"][
"output_token_logprobs"
]
output_top_logprobs = content["meta_info"].get(
"output_top_logprobs", []
)
if (
not self.tokenizer_manager.server_args.incremental_streaming_output
):
output_token_logprobs = output_token_logprobs[
n_prev_token:total_output_logprobs
]
output_top_logprobs = output_top_logprobs[
n_prev_token:total_output_logprobs
]
logprobs = to_openai_style_logprobs( logprobs = to_openai_style_logprobs(
input_token_logprobs=input_token_logprobs, input_token_logprobs=input_token_logprobs,
input_top_logprobs=input_top_logprobs, input_top_logprobs=input_top_logprobs,
output_token_logprobs=content["meta_info"][ output_token_logprobs=output_token_logprobs,
"output_token_logprobs" output_top_logprobs=output_top_logprobs,
][n_prev_token:total_output_logprobs],
output_top_logprobs=content["meta_info"].get(
"output_top_logprobs", []
)[n_prev_token:total_output_logprobs],
) )
n_prev_tokens[index] = total_output_logprobs n_prev_tokens[index] = total_output_logprobs
@@ -122,6 +122,12 @@ _REQUEST_STATE_WAIT_TIMEOUT = envs.SGLANG_REQUEST_STATE_WAIT_TIMEOUT.get()
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
_INCREMENTAL_STREAMING_META_INFO_KEYS = (
"output_token_logprobs",
"output_top_logprobs",
"output_token_ids_logprobs",
)
@dataclasses.dataclass @dataclasses.dataclass
class ReqState: class ReqState:
@@ -167,6 +173,31 @@ 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 _slice_streaming_output_meta_info(
meta_info: Dict[Any, Any],
last_output_offset: int,
) -> None:
"""Align output-side metadata with the current incremental streaming chunk."""
for key in meta_info.keys() & set(_INCREMENTAL_STREAMING_META_INFO_KEYS):
meta_info[key] = meta_info[key][last_output_offset:]
def _merge_incremental_stream_meta_info(
out_list: list[dict[str, Any]],
) -> dict[str, Any]:
"""Preserve delta-style output metadata when queued chunks are coalesced."""
meta_info_list = [chunk["meta_info"] for chunk in out_list]
meta_info = dict(meta_info_list[-1])
for key in _INCREMENTAL_STREAMING_META_INFO_KEYS:
if any(key in chunk_meta_info for chunk_meta_info in meta_info_list):
meta_info[key] = [
item
for chunk_meta_info in meta_info_list
for item in chunk_meta_info.get(key, [])
]
return meta_info
class InputFormat(Enum): class InputFormat(Enum):
"""Input format types for tokenization handling.""" """Input format types for tokenization handling."""
@@ -1167,9 +1198,8 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
obj.rid, obj.rid,
len(out_list), len(out_list),
) )
# Coalesce all deltas into a single chunk. Both text and # Coalesce all deltas into a single chunk. Text, output_ids,
# output_ids are incremental, so we concatenate them; all # and output-side incremental metadata all need to be merged.
# other fields (meta_info, etc.) are taken from the last chunk.
out = dict(out_list[-1]) out = dict(out_list[-1])
if "output_ids" in out: if "output_ids" in out:
out["output_ids"] = [ out["output_ids"] = [
@@ -1177,6 +1207,8 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
] ]
if "text" in out: if "text" in out:
out["text"] = "".join(chunk["text"] for chunk in out_list) out["text"] = "".join(chunk["text"] for chunk in out_list)
if "meta_info" in out:
out["meta_info"] = _merge_incremental_stream_meta_info(out_list)
else: else:
out = out_list[-1] out = out_list[-1]
@@ -1607,8 +1639,10 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
# 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: if self.server_args.incremental_streaming_output and is_stream:
output_offset = state.last_output_offset
state.output_ids.extend(recv_obj.output_ids[i]) state.output_ids.extend(recv_obj.output_ids[i])
output_token_ids = state.output_ids[state.last_output_offset :] 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) state.last_output_offset = len(state.output_ids)
output_text = state.text[state.last_text_offset :] output_text = state.text[state.last_text_offset :]
state.last_text_offset = len(state.text) state.last_text_offset = len(state.text)
@@ -1626,8 +1660,10 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
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: if self.server_args.incremental_streaming_output and is_stream:
output_offset = state.last_output_offset
state.output_ids.extend(recv_obj.output_ids[i]) state.output_ids.extend(recv_obj.output_ids[i])
output_token_ids = state.output_ids[state.last_output_offset :] 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) state.last_output_offset = len(state.output_ids)
else: else:
state.output_ids.extend(recv_obj.output_ids[i]) state.output_ids.extend(recv_obj.output_ids[i])
+44
View File
@@ -1,5 +1,6 @@
"""This file contains the SGL programs used for unit testing.""" """This file contains the SGL programs used for unit testing."""
import asyncio
import json import json
import re import re
import time import time
@@ -352,6 +353,49 @@ def test_stream():
out += chunk out += chunk
def test_stream_logprobs():
@sgl.function
def qa(s, question):
s += sgl.system("You are a helpful assistant.")
s += sgl.user(question)
s += sgl.assistant(sgl.gen("answer", return_logprob=True))
async def collect_chunks():
ret = qa(
question="Compose an engaging travel blog post about a recent trip to Hawaii, highlighting cultural experiences and must-see attractions.",
stream=True,
temperature=0,
max_new_tokens=64,
)
chunks = []
async for chunk_text, meta_info in ret.text_async_iter(
"answer", return_meta_data=True
):
chunks.append((chunk_text, meta_info))
return chunks
chunks = asyncio.run(collect_chunks())
assert len(chunks) > 0
prev_completion_tokens = 0
prev_output_token_logprobs_length = 0
for chunk_text, meta_info in chunks:
assert chunk_text
assert "output_token_logprobs" in meta_info
assert "output_token_logprobs_length" in meta_info
completion_tokens = meta_info["completion_tokens"]
output_token_logprobs_length = meta_info["output_token_logprobs_length"]
chunk_output_token_logprobs = meta_info["output_token_logprobs"]
assert completion_tokens == output_token_logprobs_length
assert len(chunk_output_token_logprobs) == (
completion_tokens - prev_completion_tokens
)
assert len(chunk_output_token_logprobs) == (
output_token_logprobs_length - prev_output_token_logprobs_length
)
prev_completion_tokens = completion_tokens
prev_output_token_logprobs_length = output_token_logprobs_length
def test_regex(): def test_regex():
regex = r"((25[0-5]|2[0-4]\d|[01]?\d\d?).){3}(25[0-5]|2[0-4]\d|[01]?\d\d?)" regex = r"((25[0-5]|2[0-4]\d|[01]?\d\d?).){3}(25[0-5]|2[0-4]\d|[01]?\d\d?)"
@@ -15,6 +15,7 @@ from sglang.test.test_programs import (
test_regex, test_regex,
test_select, test_select,
test_stream, test_stream,
test_stream_logprobs,
test_tool_use, test_tool_use,
) )
from sglang.test.test_utils import DEFAULT_MODEL_NAME_FOR_TEST, CustomTestCase from sglang.test.test_utils import DEFAULT_MODEL_NAME_FOR_TEST, CustomTestCase
@@ -32,6 +33,7 @@ class TestSRTBackend(CustomTestCase):
model_path=DEFAULT_MODEL_NAME_FOR_TEST, model_path=DEFAULT_MODEL_NAME_FOR_TEST,
cuda_graph_max_bs=4, cuda_graph_max_bs=4,
mem_fraction_static=0.7, mem_fraction_static=0.7,
incremental_streaming_output=True,
log_level="info", log_level="info",
) )
sgl.set_default_backend(cls.backend) sgl.set_default_backend(cls.backend)
@@ -68,6 +70,9 @@ class TestSRTBackend(CustomTestCase):
def test_stream(self): def test_stream(self):
test_stream() test_stream()
def test_stream_logprobs(self):
test_stream_logprobs()
def test_regex(self): def test_regex(self):
test_regex() test_regex()