Pack scattered output-streamer state into a dedicated accumulator (#25705)

This commit is contained in:
fzyzcjy
2026-05-19 09:14:59 +08:00
committed by GitHub
parent e8e55bb19b
commit d8f190dfba
@@ -1,7 +1,7 @@
from __future__ import annotations
import logging
from dataclasses import dataclass
from dataclasses import dataclass, field
from typing import (
Any,
Callable,
@@ -123,58 +123,15 @@ class SchedulerOutputStreamer:
skip_req: Optional[Req] = None,
is_idle_batch: bool = False,
):
rids = []
http_worker_ipcs = []
finished_reasons: List[BaseFinishReason] = []
decoded_texts = []
decode_ids_list = []
read_offsets = []
output_ids = []
skip_special_tokens = []
spaces_between_special_tokens = []
no_stop_trim = []
prompt_tokens = []
reasoning_tokens = []
completion_tokens = []
cached_tokens = []
cached_tokens_details = [] # Detailed breakdown by cache source
spec_verify_ct = []
spec_num_correct_drafts = []
spec_correct_drafts_histogram = []
retraction_counts = []
output_hidden_states = None
acc = _GenerationStreamAccumulator(
return_logprob=return_logprob,
spec_algorithm=self.spec_algorithm,
disaggregation_mode=self.disaggregation_mode,
default_stream_interval=self.server_args.stream_interval,
default_force_stream_interval=DEFAULT_FORCE_STREAM_INTERVAL,
get_cached_tokens_details=self.get_cached_tokens_details,
)
load = self.load_inquirer_get_loads(GetLoadsReqInput(include=["core"]))
routed_experts = None
indexer_topk = None
customized_info = {}
time_stats = []
if return_logprob:
input_token_logprobs_val = []
input_token_logprobs_idx = []
output_token_logprobs_val = []
output_token_logprobs_idx = []
input_top_logprobs_val = []
input_top_logprobs_idx = []
output_top_logprobs_val = []
output_top_logprobs_idx = []
input_token_ids_logprobs_val = []
input_token_ids_logprobs_idx = []
output_token_ids_logprobs_val = []
output_token_ids_logprobs_idx = []
else:
input_token_logprobs_val = input_token_logprobs_idx = (
output_token_logprobs_val
) = output_token_logprobs_idx = input_top_logprobs_val = (
input_top_logprobs_idx
) = output_top_logprobs_val = output_top_logprobs_idx = (
input_token_ids_logprobs_val
) = input_token_ids_logprobs_idx = output_token_ids_logprobs_val = (
output_token_ids_logprobs_idx
) = None
for req in reqs:
if req is skip_req:
@@ -216,44 +173,44 @@ class SchedulerOutputStreamer:
send_output_token_logprobs_offset = (
req.send_output_token_logprobs_offset
)
rids.append(req.rid)
http_worker_ipcs.append(req.http_worker_ipc)
finished_reasons.append(
acc.rids.append(req.rid)
acc.http_worker_ipcs.append(req.http_worker_ipc)
acc.finished_reasons.append(
req.finished_reason.to_json() if req.finished_reason else None
)
decoded_texts.append(req.decoded_text)
acc.decoded_texts.append(req.decoded_text)
decode_ids, read_offset = req.init_incremental_detokenize()
decode_ids_list.append(decode_ids[req.send_decode_id_offset :])
acc.decode_ids_list.append(decode_ids[req.send_decode_id_offset :])
# Exclude the tokens after stop condition
output_ids_ = req.output_ids_through_stop
req.send_decode_id_offset = len(decode_ids)
read_offsets.append(read_offset)
output_ids.append(output_ids_[send_token_offset:])
acc.read_offsets.append(read_offset)
acc.output_ids.append(output_ids_[send_token_offset:])
req.send_token_offset = len(output_ids_)
skip_special_tokens.append(req.sampling_params.skip_special_tokens)
spaces_between_special_tokens.append(
acc.skip_special_tokens.append(req.sampling_params.skip_special_tokens)
acc.spaces_between_special_tokens.append(
req.sampling_params.spaces_between_special_tokens
)
no_stop_trim.append(req.sampling_params.no_stop_trim)
prompt_tokens.append(len(req.origin_input_ids))
reasoning_tokens.append(req.reasoning_tokens)
completion_tokens.append(len(output_ids_))
cached_tokens.append(req.cached_tokens)
acc.no_stop_trim.append(req.sampling_params.no_stop_trim)
acc.prompt_tokens.append(len(req.origin_input_ids))
acc.reasoning_tokens.append(req.reasoning_tokens)
acc.completion_tokens.append(len(output_ids_))
acc.cached_tokens.append(req.cached_tokens)
# Collect detailed cache breakdown if available
cached_tokens_details.append(self.get_cached_tokens_details(req))
acc.cached_tokens_details.append(self.get_cached_tokens_details(req))
retraction_counts.append(req.retraction_count)
acc.retraction_counts.append(req.retraction_count)
time_stats.append(req.time_stats)
acc.time_stats.append(req.time_stats)
if not self.spec_algorithm.is_none():
spec_verify_ct.append(req.spec_verify_ct)
spec_num_correct_drafts.append(req.spec_num_correct_drafts)
spec_correct_drafts_histogram.append(
acc.spec_verify_ct.append(req.spec_verify_ct)
acc.spec_num_correct_drafts.append(req.spec_num_correct_drafts)
acc.spec_correct_drafts_histogram.append(
req.spec_correct_drafts_histogram
)
@@ -266,84 +223,82 @@ class SchedulerOutputStreamer:
# Only send when input logprobs have been computed (after prefill)
and req.input_token_logprobs_val is not None
):
input_token_logprobs_val.append(req.input_token_logprobs_val)
input_token_logprobs_idx.append(req.input_token_logprobs_idx)
input_top_logprobs_val.append(req.input_top_logprobs_val)
input_top_logprobs_idx.append(req.input_top_logprobs_idx)
input_token_ids_logprobs_val.append(
acc.input_token_logprobs_val.append(
req.input_token_logprobs_val
)
acc.input_token_logprobs_idx.append(
req.input_token_logprobs_idx
)
acc.input_top_logprobs_val.append(req.input_top_logprobs_val)
acc.input_top_logprobs_idx.append(req.input_top_logprobs_idx)
acc.input_token_ids_logprobs_val.append(
req.input_token_ids_logprobs_val
)
input_token_ids_logprobs_idx.append(
acc.input_token_ids_logprobs_idx.append(
req.input_token_ids_logprobs_idx
)
req.input_logprob_sent = True
else:
input_token_logprobs_val.append([])
input_token_logprobs_idx.append([])
input_top_logprobs_val.append([])
input_top_logprobs_idx.append([])
input_token_ids_logprobs_val.append([])
input_token_ids_logprobs_idx.append([])
acc.input_token_logprobs_val.append([])
acc.input_token_logprobs_idx.append([])
acc.input_top_logprobs_val.append([])
acc.input_top_logprobs_idx.append([])
acc.input_token_ids_logprobs_val.append([])
acc.input_token_ids_logprobs_idx.append([])
if req.return_logprob:
logprob_end = max(len(output_ids_), 1)
output_token_logprobs_val.append(
acc.output_token_logprobs_val.append(
req.output_token_logprobs_val[
send_output_token_logprobs_offset:logprob_end
]
)
output_token_logprobs_idx.append(
acc.output_token_logprobs_idx.append(
req.output_token_logprobs_idx[
send_output_token_logprobs_offset:logprob_end
]
)
output_top_logprobs_val.append(
acc.output_top_logprobs_val.append(
req.output_top_logprobs_val[
send_output_token_logprobs_offset:logprob_end
]
)
output_top_logprobs_idx.append(
acc.output_top_logprobs_idx.append(
req.output_top_logprobs_idx[
send_output_token_logprobs_offset:logprob_end
]
)
output_token_ids_logprobs_val.append(
acc.output_token_ids_logprobs_val.append(
req.output_token_ids_logprobs_val[
send_output_token_logprobs_offset:logprob_end
]
)
output_token_ids_logprobs_idx.append(
acc.output_token_ids_logprobs_idx.append(
req.output_token_ids_logprobs_idx[
send_output_token_logprobs_offset:logprob_end
]
)
req.send_output_token_logprobs_offset = logprob_end
else:
output_token_logprobs_val.append([])
output_token_logprobs_idx.append([])
output_top_logprobs_val.append([])
output_top_logprobs_idx.append([])
output_token_ids_logprobs_val.append([])
output_token_ids_logprobs_idx.append([])
acc.output_token_logprobs_val.append([])
acc.output_token_logprobs_idx.append([])
acc.output_top_logprobs_val.append([])
acc.output_top_logprobs_idx.append([])
acc.output_token_ids_logprobs_val.append([])
acc.output_token_ids_logprobs_idx.append([])
if req.return_hidden_states:
if output_hidden_states is None:
output_hidden_states = []
output_hidden_states.append(req.hidden_states)
acc.output_hidden_states.append(req.hidden_states)
if req.return_routed_experts:
if routed_experts is None:
routed_experts = []
routed_experts.append(req.routed_experts)
acc.routed_experts.append(req.routed_experts)
if req.return_indexer_topk:
if indexer_topk is None:
indexer_topk = []
indexer_topk.append(req.indexer_topk)
acc.indexer_topk.append(req.indexer_topk)
if req.customized_info is not None:
for k, v in req.customized_info.items():
if k not in customized_info:
customized_info[k] = []
customized_info[k].append(
if k not in acc.customized_info:
acc.customized_info[k] = []
acc.customized_info[k].append(
v[send_token_offset : len(output_ids_)]
)
@@ -354,51 +309,51 @@ class SchedulerOutputStreamer:
):
req.log_time_stats()
dp_ranks = [self.ps.dp_rank] * len(rids) if rids else None
dp_ranks = [self.ps.dp_rank] * len(acc.rids) if acc.rids else None
# Send to detokenizer
if reqs or is_idle_batch:
self.send_to_detokenizer.send_output(
BatchTokenIDOutput(
rids=rids,
http_worker_ipcs=http_worker_ipcs,
spec_verify_ct=spec_verify_ct,
spec_num_correct_drafts=spec_num_correct_drafts,
spec_correct_drafts_histogram=spec_correct_drafts_histogram,
time_stats=time_stats,
finished_reasons=finished_reasons,
decoded_texts=decoded_texts,
decode_ids=decode_ids_list,
read_offsets=read_offsets,
output_ids=output_ids,
skip_special_tokens=skip_special_tokens,
spaces_between_special_tokens=spaces_between_special_tokens,
no_stop_trim=no_stop_trim,
prompt_tokens=prompt_tokens,
reasoning_tokens=reasoning_tokens,
completion_tokens=completion_tokens,
cached_tokens=cached_tokens,
cached_tokens_details=cached_tokens_details,
input_token_logprobs_val=input_token_logprobs_val,
input_token_logprobs_idx=input_token_logprobs_idx,
output_token_logprobs_val=output_token_logprobs_val,
output_token_logprobs_idx=output_token_logprobs_idx,
input_top_logprobs_val=input_top_logprobs_val,
input_top_logprobs_idx=input_top_logprobs_idx,
output_top_logprobs_val=output_top_logprobs_val,
output_top_logprobs_idx=output_top_logprobs_idx,
input_token_ids_logprobs_val=input_token_ids_logprobs_val,
input_token_ids_logprobs_idx=input_token_ids_logprobs_idx,
output_token_ids_logprobs_val=output_token_ids_logprobs_val,
output_token_ids_logprobs_idx=output_token_ids_logprobs_idx,
rids=acc.rids,
http_worker_ipcs=acc.http_worker_ipcs,
spec_verify_ct=acc.spec_verify_ct,
spec_num_correct_drafts=acc.spec_num_correct_drafts,
spec_correct_drafts_histogram=acc.spec_correct_drafts_histogram,
time_stats=acc.time_stats,
finished_reasons=acc.finished_reasons,
decoded_texts=acc.decoded_texts,
decode_ids=acc.decode_ids_list,
read_offsets=acc.read_offsets,
output_ids=acc.output_ids,
skip_special_tokens=acc.skip_special_tokens,
spaces_between_special_tokens=acc.spaces_between_special_tokens,
no_stop_trim=acc.no_stop_trim,
prompt_tokens=acc.prompt_tokens,
reasoning_tokens=acc.reasoning_tokens,
completion_tokens=acc.completion_tokens,
cached_tokens=acc.cached_tokens,
cached_tokens_details=acc.cached_tokens_details,
input_token_logprobs_val=acc.input_token_logprobs_val,
input_token_logprobs_idx=acc.input_token_logprobs_idx,
output_token_logprobs_val=acc.output_token_logprobs_val,
output_token_logprobs_idx=acc.output_token_logprobs_idx,
input_top_logprobs_val=acc.input_top_logprobs_val,
input_top_logprobs_idx=acc.input_top_logprobs_idx,
output_top_logprobs_val=acc.output_top_logprobs_val,
output_top_logprobs_idx=acc.output_top_logprobs_idx,
input_token_ids_logprobs_val=acc.input_token_ids_logprobs_val,
input_token_ids_logprobs_idx=acc.input_token_ids_logprobs_idx,
output_token_ids_logprobs_val=acc.output_token_ids_logprobs_val,
output_token_ids_logprobs_idx=acc.output_token_ids_logprobs_idx,
output_token_entropy_val=None,
output_hidden_states=output_hidden_states,
routed_experts=routed_experts,
indexer_topk=indexer_topk,
customized_info=customized_info,
output_hidden_states=acc.output_hidden_states or None,
routed_experts=acc.routed_experts or None,
indexer_topk=acc.indexer_topk or None,
customized_info=acc.customized_info,
placeholder_tokens_idx=None,
placeholder_tokens_val=None,
retraction_counts=retraction_counts,
retraction_counts=acc.retraction_counts,
load=load,
dp_ranks=dp_ranks,
)
@@ -468,3 +423,75 @@ class SchedulerOutputStreamer:
pooled_hidden_states=stacked_phs,
)
)
@dataclass(slots=True, kw_only=True)
class _GenerationStreamAccumulator:
return_logprob: bool
spec_algorithm: Any
disaggregation_mode: DisaggregationMode
default_stream_interval: int
default_force_stream_interval: int
get_cached_tokens_details: Callable[[Req], Optional[dict]]
rids: list = field(default_factory=list)
http_worker_ipcs: list = field(default_factory=list)
finished_reasons: list = field(default_factory=list)
decoded_texts: list = field(default_factory=list)
decode_ids_list: list = field(default_factory=list)
read_offsets: list = field(default_factory=list)
output_ids: list = field(default_factory=list)
skip_special_tokens: list = field(default_factory=list)
spaces_between_special_tokens: list = field(default_factory=list)
no_stop_trim: list = field(default_factory=list)
prompt_tokens: list = field(default_factory=list)
reasoning_tokens: list = field(default_factory=list)
completion_tokens: list = field(default_factory=list)
cached_tokens: list = field(default_factory=list)
cached_tokens_details: list = field(
default_factory=list
) # Detailed breakdown by cache source
spec_verify_ct: list = field(default_factory=list)
spec_num_correct_drafts: list = field(default_factory=list)
spec_correct_drafts_histogram: list = field(default_factory=list)
retraction_counts: list = field(default_factory=list)
output_hidden_states: list = field(default_factory=list)
routed_experts: list = field(default_factory=list)
indexer_topk: list = field(default_factory=list)
customized_info: dict = field(default_factory=dict)
time_stats: list = field(default_factory=list)
input_token_logprobs_val: Optional[list] = None
input_token_logprobs_idx: Optional[list] = None
output_token_logprobs_val: Optional[list] = None
output_token_logprobs_idx: Optional[list] = None
input_top_logprobs_val: Optional[list] = None
input_top_logprobs_idx: Optional[list] = None
output_top_logprobs_val: Optional[list] = None
output_top_logprobs_idx: Optional[list] = None
input_token_ids_logprobs_val: Optional[list] = None
input_token_ids_logprobs_idx: Optional[list] = None
output_token_ids_logprobs_val: Optional[list] = None
output_token_ids_logprobs_idx: Optional[list] = None
def __post_init__(self) -> None:
if self.return_logprob:
self.input_token_logprobs_val = []
self.input_token_logprobs_idx = []
self.output_token_logprobs_val = []
self.output_token_logprobs_idx = []
self.input_top_logprobs_val = []
self.input_top_logprobs_idx = []
self.output_top_logprobs_val = []
self.output_top_logprobs_idx = []
self.input_token_ids_logprobs_val = []
self.input_token_ids_logprobs_idx = []
self.output_token_ids_logprobs_val = []
self.output_token_ids_logprobs_idx = []
def accept(self, *, req: Req) -> None:
raise NotImplementedError
def to_payload(
self, *, load, dp_rank: int, is_idle_batch: bool, has_reqs: bool
) -> Optional[BatchTokenIDOutput]:
raise NotImplementedError