Pack scattered output-streamer state into a dedicated accumulator (#25705)
This commit is contained in:
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user