diff --git a/python/sglang/srt/managers/scheduler_components/output_streamer.py b/python/sglang/srt/managers/scheduler_components/output_streamer.py index 3569a4975..7a8ccf26a 100644 --- a/python/sglang/srt/managers/scheduler_components/output_streamer.py +++ b/python/sglang/srt/managers/scheduler_components/output_streamer.py @@ -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