Fix customized_info incremental streaming (#27205)

This commit is contained in:
Aurick Qiao
2026-06-05 21:55:01 +08:00
committed by GitHub
parent faa6286946
commit c06802dc16
2 changed files with 207 additions and 9 deletions
@@ -33,7 +33,7 @@ from contextlib import nullcontext
from datetime import datetime
from enum import Enum
from http import HTTPStatus
from typing import Any, Awaitable, Dict, List, Optional, Tuple, Union
from typing import Any, Awaitable, Dict, Iterable, List, Optional, Tuple, Union
import fastapi
import pybase64
@@ -205,6 +205,9 @@ class ReqState:
output_top_logprobs: List[Any] = dataclasses.field(default_factory=list)
input_token_ids_logprobs: List[Any] = dataclasses.field(default_factory=list)
output_token_ids_logprobs: List[Any] = dataclasses.field(default_factory=list)
customized_info_accumulated: Dict[str, List[Any]] = dataclasses.field(
default_factory=dict
)
# For return_prompt_token_ids: stores prompt token IDs captured after tokenization
prompt_token_ids: Optional[List[int]] = None
@@ -213,9 +216,13 @@ class ReqState:
def _slice_streaming_output_meta_info(
meta_info: Dict[Any, Any],
last_output_offset: int,
customized_info_keys: Optional[Iterable[str]] = None,
) -> None:
"""Align output-side metadata with the current incremental streaming chunk."""
for key in meta_info.keys() & set(_INCREMENTAL_STREAMING_META_INFO_KEYS):
streaming_meta_info_keys = set(_INCREMENTAL_STREAMING_META_INFO_KEYS)
if customized_info_keys is not None:
streaming_meta_info_keys.update(customized_info_keys)
for key in meta_info.keys() & streaming_meta_info_keys:
meta_info[key] = meta_info[key][last_output_offset:]
@@ -1288,6 +1295,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
self,
out_list: list,
rid: str,
customized_info_keys: Optional[Iterable[str]] = None,
) -> dict:
"""Coalesce multiple incremental streaming chunks into one.
@@ -1309,7 +1317,10 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
if "meta_info" in out:
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:
incremental_streaming_keys = set(_INCREMENTAL_STREAMING_META_INFO_KEYS)
if customized_info_keys is not None:
incremental_streaming_keys.update(customized_info_keys)
for key in incremental_streaming_keys:
if any(key in m for m in meta_info_list):
meta_info[key] = [
item for m in meta_info_list for item in m.get(key, [])
@@ -1401,7 +1412,11 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
is_stream and self.server_args.incremental_streaming_output
)
if incremental_stream and len(out_list) > 1:
out = self._coalesce_streaming_chunks(out_list, obj.rid)
out = self._coalesce_streaming_chunks(
out_list,
obj.rid,
state.customized_info_accumulated.keys(),
)
else:
out = out_list[-1]
@@ -1847,6 +1862,12 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
meta_info["cached_tokens_details"] = recv_obj.cached_tokens_details[
i
]
if recv_obj.customized_info is not None:
for k, v in recv_obj.customized_info.items():
if k not in state.customized_info_accumulated:
state.customized_info_accumulated[k] = []
state.customized_info_accumulated[k].extend(v[i])
meta_info[k] = state.customized_info_accumulated[k]
if getattr(recv_obj, "output_hidden_states", None):
hidden_states = recv_obj.output_hidden_states[i]
@@ -1866,9 +1887,6 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
if isinstance(val, torch.Tensor):
val = pybase64.b64encode(val.numpy().tobytes()).decode("utf-8")
meta_info["indexer_topk"] = val
if getattr(recv_obj, "customized_info", None):
for k, v in recv_obj.customized_info.items():
meta_info[k] = v[i]
if getattr(recv_obj, "dp_ranks", None):
meta_info["dp_rank"] = recv_obj.dp_ranks[i]
@@ -1888,7 +1906,11 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
if is_stream:
if incremental:
output_token_ids = delta_output_ids
_slice_streaming_output_meta_info(meta_info, output_offset)
_slice_streaming_output_meta_info(
meta_info,
output_offset,
state.customized_info_accumulated.keys(),
)
state.last_output_offset = len(state.output_ids)
out_dict = {
"text": delta_text,
@@ -1932,7 +1954,11 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
if is_stream:
if incremental:
output_token_ids = delta_output_ids
_slice_streaming_output_meta_info(meta_info, output_offset)
_slice_streaming_output_meta_info(
meta_info,
output_offset,
state.customized_info_accumulated.keys(),
)
state.last_output_offset = len(state.output_ids)
out_dict = {
"output_ids": output_token_ids,