From 18a7eb9e58aa3016c9e0a4cfa61af484eb40ac03 Mon Sep 17 00:00:00 2001 From: fzyzcjy <5236035+fzyzcjy@users.noreply.github.com> Date: Mon, 18 May 2026 18:44:07 +0800 Subject: [PATCH] Move output streaming to SchedulerOutputStreamer (#25635) --- python/sglang/srt/disaggregation/decode.py | 17 +- python/sglang/srt/disaggregation/prefill.py | 11 +- python/sglang/srt/dllm/mixin/scheduler.py | 2 +- python/sglang/srt/managers/scheduler.py | 6 +- .../scheduler_components/output_streamer.py | 439 ++++++++++++++++- .../scheduler_output_processor_mixin.py | 459 +----------------- ...test_priority_scheduling_disaggregation.py | 4 +- .../mem_cache/test_decode_radix_lock_ref.py | 2 +- 8 files changed, 459 insertions(+), 481 deletions(-) diff --git a/python/sglang/srt/disaggregation/decode.py b/python/sglang/srt/disaggregation/decode.py index 2f9218388..9ee27fdc6 100644 --- a/python/sglang/srt/disaggregation/decode.py +++ b/python/sglang/srt/disaggregation/decode.py @@ -532,9 +532,7 @@ class DecodePreallocQueue: message = f"Request {req.rid} exceeds the maximum number of tokens: {len(req.origin_input_ids)} > {self.max_total_num_tokens}" logger.error(message) prepare_abort(req, message, status_code=HTTPStatus.BAD_REQUEST) - self.scheduler.stream_output( - self.scheduler.output_streamer, [req], req.return_logprob - ) + self.scheduler.output_streamer.stream_output([req], req.return_logprob) return True if self._uses_swa_tail_prealloc(): _, swa_required = self._prealloc_required_tokens(req) @@ -546,9 +544,7 @@ class DecodePreallocQueue: ) logger.error(message) prepare_abort(req, message, status_code=HTTPStatus.BAD_REQUEST) - self.scheduler.stream_output( - self.scheduler.output_streamer, [req], req.return_logprob - ) + self.scheduler.output_streamer.stream_output([req], req.return_logprob) return True return False @@ -782,8 +778,7 @@ class DecodePreallocQueue: if rids_to_check is not None and decode_req.req.rid not in rids_to_check: continue if isinstance(decode_req.req.finished_reason, FINISH_ABORT): - self.scheduler.stream_output( - self.scheduler.output_streamer, + self.scheduler.output_streamer.stream_output( [decode_req.req], decode_req.req.return_logprob, ) @@ -1514,8 +1509,7 @@ class DecodeTransferQueue: error_message, status_code=HTTPStatus.INTERNAL_SERVER_ERROR, ) - self.scheduler.stream_output( - self.scheduler.output_streamer, + self.scheduler.output_streamer.stream_output( [decode_req.req], decode_req.req.return_logprob, ) @@ -1533,8 +1527,7 @@ class DecodeTransferQueue: indices_to_remove.add(i) # Check if request was aborted due to corruption if isinstance(decode_req.req.finished_reason, FINISH_ABORT): - self.scheduler.stream_output( - self.scheduler.output_streamer, + self.scheduler.output_streamer.stream_output( [decode_req.req], decode_req.req.return_logprob, ) diff --git a/python/sglang/srt/disaggregation/prefill.py b/python/sglang/srt/disaggregation/prefill.py index d98a44134..1d359979a 100644 --- a/python/sglang/srt/disaggregation/prefill.py +++ b/python/sglang/srt/disaggregation/prefill.py @@ -252,9 +252,7 @@ class PrefillBootstrapQueue: logger.error(message) req.time_stats.trace_ctx.abort(abort_info={"reason": message}) prepare_abort(req, message, status_code=HTTPStatus.BAD_REQUEST) - self.scheduler.stream_output( - self.scheduler.output_streamer, [req], req.return_logprob - ) + self.scheduler.output_streamer.stream_output([req], req.return_logprob) return True return False @@ -311,9 +309,7 @@ class PrefillBootstrapQueue: prepare_abort( req, error_message, status_code=HTTPStatus.INTERNAL_SERVER_ERROR ) - self.scheduler.stream_output( - self.scheduler.output_streamer, [req], req.return_logprob - ) + self.scheduler.output_streamer.stream_output([req], req.return_logprob) indices_to_remove.add(i) failed_reqs.append(req) if self.scheduler.metrics_reporter.enable_metrics: @@ -694,8 +690,7 @@ class SchedulerDisaggregationPrefillMixin: self.metrics_reporter.kv_transfer_speed_gb_s = metrics["speed_gb_s"] # Stream requests which have finished transfer - self.stream_output( - self.output_streamer, + self.output_streamer.stream_output( done_reqs, any(req.return_logprob for req in done_reqs), None, diff --git a/python/sglang/srt/dllm/mixin/scheduler.py b/python/sglang/srt/dllm/mixin/scheduler.py index 703bf4988..6d35864ab 100644 --- a/python/sglang/srt/dllm/mixin/scheduler.py +++ b/python/sglang/srt/dllm/mixin/scheduler.py @@ -89,7 +89,7 @@ class SchedulerDllmMixin: release_kv_cache(req, self.tree_cache) req.time_stats.set_completion_time() - self.stream_output(self.output_streamer, batch.reqs, batch.return_logprob) + self.output_streamer.stream_output(batch.reqs, batch.return_logprob) self.token_to_kv_pool_allocator.free_group_end() can_run_cuda_graph = getattr(result, "can_run_cuda_graph", False) diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 6e36b3398..96faca466 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -651,9 +651,7 @@ class Scheduler( server_args=self.server_args, model_config=self.model_config, max_recv_per_poll=self.max_recv_per_poll, - stream_output=lambda *a, **kw: self.stream_output( - self.output_streamer, *a, **kw - ), + stream_output=lambda *a, **kw: self.output_streamer.stream_output(*a, **kw), get_last_forward_mode=lambda: ( self.last_batch.forward_mode if self.last_batch is not None else None ), @@ -1930,7 +1928,7 @@ class Scheduler( abort_info={"reason": error_msg} ) prepare_abort(req, error_msg, status_code=HTTPStatus.BAD_REQUEST) - self.stream_output(self.output_streamer, [req], req.return_logprob) + self.output_streamer.stream_output([req], req.return_logprob) return elif ( diff --git a/python/sglang/srt/managers/scheduler_components/output_streamer.py b/python/sglang/srt/managers/scheduler_components/output_streamer.py index eba315795..3569a4975 100644 --- a/python/sglang/srt/managers/scheduler_components/output_streamer.py +++ b/python/sglang/srt/managers/scheduler_components/output_streamer.py @@ -2,13 +2,28 @@ from __future__ import annotations import logging from dataclasses import dataclass -from typing import Any, Callable +from typing import ( + Any, + Callable, + List, + Optional, +) +import torch import zmq from sglang.srt.disaggregation.utils import DisaggregationMode from sglang.srt.distributed.parallel_state_wrapper import ParallelState from sglang.srt.environ import envs +from sglang.srt.managers.io_struct import ( + BatchEmbeddingOutput, + BatchTokenIDOutput, + GetLoadsReqInput, +) +from sglang.srt.managers.schedule_batch import ( + BaseFinishReason, + Req, +) from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache from sglang.srt.server_args import ServerArgs from sglang.srt.speculative.spec_info import SpeculativeAlgorithm @@ -31,3 +46,425 @@ class SchedulerOutputStreamer: enable_hicache_storage: Callable[[], bool] load_inquirer_get_loads: Callable[..., Any] _test_stream_output_count: int = 0 + + def _get_storage_backend_type(self) -> str: + """Get storage backend type from tree_cache.""" + storage_backend_type = "none" + cache_controller = getattr(self.tree_cache, "cache_controller", None) + if cache_controller and hasattr(cache_controller, "storage_backend"): + storage_backend = cache_controller.storage_backend + if storage_backend is not None: + storage_backend_type = type(storage_backend).__name__ + return storage_backend_type + + def get_cached_tokens_details(self, req: Req) -> Optional[dict]: + """Get detailed cache breakdown for a request, if available. + + Returns: + - None if no cached tokens at all + - {"device": X, "host": Y} without storage breakdown + - {"device": X, "host": Y, "storage": Z} with storage breakdown + """ + if ( + req.cached_tokens_device > 0 + or req.cached_tokens_host > 0 + or req.cached_tokens_storage > 0 + ): + details = { + "device": req.cached_tokens_device, + "host": req.cached_tokens_host, + } + # Only include storage fields if L3 storage is enabled + if self.enable_hicache_storage(): + details["storage"] = req.cached_tokens_storage + details["storage_backend"] = self._get_storage_backend_type() + return details + + if req.cached_tokens > 0: + return { + "device": req.cached_tokens, + "host": 0, + } + + return None + + def stream_output( + self, + reqs: List[Req], + return_logprob: bool, + skip_req: Optional[Req] = None, + ): + """Stream the output to detokenizer.""" + if self.is_generation: + self._stream_output_generation(reqs, return_logprob, skip_req) + else: # embedding or reward model + self._stream_output_embedding(reqs) + + if envs.SGLANG_TEST_CRASH_AFTER_STREAM_OUTPUTS.get() > 0: + self._trigger_crash_for_tests( + envs.SGLANG_TEST_CRASH_AFTER_STREAM_OUTPUTS.get() + ) + + def _trigger_crash_for_tests(self, crash_threshold: int): + # Crash trigger: crash after stream_output is called N times + # This is used for testing purposes. + if not hasattr(self, "_test_stream_output_count"): + self._test_stream_output_count = 0 + self._test_stream_output_count += 1 + if self._test_stream_output_count >= crash_threshold: + raise RuntimeError( + f"Test crash after stream_output called {self._test_stream_output_count} times" + ) + + def _stream_output_generation( + self, + reqs: List[Req], + return_logprob: bool, + 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 + 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: + continue + + if req.finished(): + if req.finished_output: + # With the overlap schedule, a request will try to output twice and hit this line twice + # because of the one additional delayed token. This "continue" prevented the dummy output. + continue + req.finished_output = True + if req.finished_len is None: + req.finished_len = len(req.output_ids) + should_output = True + else: + if req.stream: + stream_interval = ( + req.sampling_params.stream_interval + or self.server_args.stream_interval + ) + + # origin stream_interval logic + should_output = ( + len(req.output_ids) % stream_interval == 1 + if stream_interval > 1 + else len(req.output_ids) % stream_interval == 0 + ) + + if should_output: + # check_match_stop_str_prefix if tail_str's suffix match stop_str prefix + should_output &= not req.check_match_stop_str_prefix() + else: + should_output = ( + len(req.output_ids) % DEFAULT_FORCE_STREAM_INTERVAL == 0 + ) + + if should_output: + send_token_offset = req.send_token_offset + 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( + req.finished_reason.to_json() if req.finished_reason else None + ) + 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 :]) + + # 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:]) + req.send_token_offset = len(output_ids_) + skip_special_tokens.append(req.sampling_params.skip_special_tokens) + 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) + + # Collect detailed cache breakdown if available + cached_tokens_details.append(self.get_cached_tokens_details(req)) + + retraction_counts.append(req.retraction_count) + + 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( + req.spec_correct_drafts_histogram + ) + + if return_logprob: + if ( + req.return_logprob + and not req.input_logprob_sent + # Decode server does not send input logprobs + and self.disaggregation_mode != DisaggregationMode.DECODE + # 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( + req.input_token_ids_logprobs_val + ) + 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([]) + + if req.return_logprob: + logprob_end = max(len(output_ids_), 1) + output_token_logprobs_val.append( + req.output_token_logprobs_val[ + send_output_token_logprobs_offset:logprob_end + ] + ) + output_token_logprobs_idx.append( + req.output_token_logprobs_idx[ + send_output_token_logprobs_offset:logprob_end + ] + ) + output_top_logprobs_val.append( + req.output_top_logprobs_val[ + send_output_token_logprobs_offset:logprob_end + ] + ) + output_top_logprobs_idx.append( + req.output_top_logprobs_idx[ + send_output_token_logprobs_offset:logprob_end + ] + ) + 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( + 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([]) + + if req.return_hidden_states: + if output_hidden_states is None: + output_hidden_states = [] + 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) + if req.return_indexer_topk: + if indexer_topk is None: + indexer_topk = [] + 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( + v[send_token_offset : len(output_ids_)] + ) + + if ( + req.finished() + and self.ps.attn_tp_rank == 0 + and self.server_args.enable_request_time_stats_logging + ): + req.log_time_stats() + + dp_ranks = [self.ps.dp_rank] * len(rids) if 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, + output_token_entropy_val=None, + output_hidden_states=output_hidden_states, + routed_experts=routed_experts, + indexer_topk=indexer_topk, + customized_info=customized_info, + placeholder_tokens_idx=None, + placeholder_tokens_val=None, + retraction_counts=retraction_counts, + load=load, + dp_ranks=dp_ranks, + ) + ) + + def _stream_output_embedding(self, reqs: List[Req]): + rids = [] + http_worker_ipcs = [] + finished_reasons: List[BaseFinishReason] = [] + + embeddings = [] + prompt_tokens = [] + cached_tokens = [] + cached_tokens_details = [] # Detailed breakdown by cache source + time_stats = [] + retraction_counts = [] + phs_list = [] + has_phs = False + for req in reqs: + if req.finished(): + rids.append(req.rid) + http_worker_ipcs.append(req.http_worker_ipc) + finished_reasons.append(req.finished_reason.to_json()) + embeddings.append(req.embedding) + prompt_tokens.append(len(req.origin_input_ids)) + cached_tokens.append(req.cached_tokens) + + # Collect detailed cache breakdown if available + cached_tokens_details.append(self.get_cached_tokens_details(req)) + time_stats.append(req.time_stats) + retraction_counts.append(req.retraction_count) + + phs = req.pooled_hidden_state + phs_list.append(phs) + if phs is not None: + has_phs = True + + # Optimize PHS for pickle: torch.stack reduces N __reduce_ex__ + # calls to 1 across the ZMQ IPC boundary. We can only stack when + # *every* entry is non-None (homogeneous batch); mixed batches + # (some requests want PHS, others don't) keep the raw list so + # positional indexing on the receiver side stays correct. + stacked_phs = None + if has_phs: + all_have_phs = all(t is not None for t in phs_list) + if all_have_phs: + if all(t.shape == phs_list[0].shape for t in phs_list): + stacked_phs = torch.stack(phs_list) + else: + stacked_phs = phs_list + else: + stacked_phs = phs_list + + self.send_to_detokenizer.send_output( + BatchEmbeddingOutput( + rids=rids, + http_worker_ipcs=http_worker_ipcs, + time_stats=time_stats, + finished_reasons=finished_reasons, + embeddings=embeddings, + prompt_tokens=prompt_tokens, + cached_tokens=cached_tokens, + cached_tokens_details=cached_tokens_details, + placeholder_tokens_idx=None, + placeholder_tokens_val=None, + retraction_counts=retraction_counts, + pooled_hidden_states=stacked_phs, + ) + ) diff --git a/python/sglang/srt/managers/scheduler_output_processor_mixin.py b/python/sglang/srt/managers/scheduler_output_processor_mixin.py index 848de4c42..adb9d9377 100644 --- a/python/sglang/srt/managers/scheduler_output_processor_mixin.py +++ b/python/sglang/srt/managers/scheduler_output_processor_mixin.py @@ -1,7 +1,7 @@ from __future__ import annotations import logging -from typing import TYPE_CHECKING, List, Optional, Union +from typing import TYPE_CHECKING, List, Union import torch @@ -10,12 +10,8 @@ from sglang.srt.environ import envs from sglang.srt.layers.logits_processor import LogitsProcessorOutput from sglang.srt.managers.io_struct import ( AbortReq, - BatchEmbeddingOutput, - BatchTokenIDOutput, - GetLoadsReqInput, ) from sglang.srt.managers.schedule_batch import ( - BaseFinishReason, Req, ScheduleBatch, ) @@ -33,9 +29,6 @@ if TYPE_CHECKING: ScheduleBatch, Scheduler, ) - from sglang.srt.managers.scheduler_components.output_streamer import ( - SchedulerOutputStreamer, - ) logger = logging.getLogger(__name__) @@ -50,53 +43,6 @@ class SchedulerOutputProcessorMixin: We put them into a separate file to make the `scheduler.py` shorter. """ - @staticmethod - def _get_storage_backend_type(self: "SchedulerOutputStreamer") -> str: - """Get storage backend type from tree_cache.""" - storage_backend_type = "none" - cache_controller = getattr(self.tree_cache, "cache_controller", None) - if cache_controller and hasattr(cache_controller, "storage_backend"): - storage_backend = cache_controller.storage_backend - if storage_backend is not None: - storage_backend_type = type(storage_backend).__name__ - return storage_backend_type - - @staticmethod - def get_cached_tokens_details( - self: "SchedulerOutputStreamer", req: Req - ) -> Optional[dict]: - """Get detailed cache breakdown for a request, if available. - - Returns: - - None if no cached tokens at all - - {"device": X, "host": Y} without storage breakdown - - {"device": X, "host": Y, "storage": Z} with storage breakdown - """ - if ( - req.cached_tokens_device > 0 - or req.cached_tokens_host > 0 - or req.cached_tokens_storage > 0 - ): - details = { - "device": req.cached_tokens_device, - "host": req.cached_tokens_host, - } - # Only include storage fields if L3 storage is enabled - if self.enable_hicache_storage(): - details["storage"] = req.cached_tokens_storage - details["storage_backend"] = ( - SchedulerOutputProcessorMixin._get_storage_backend_type(self) - ) - return details - - if req.cached_tokens > 0: - return { - "device": req.cached_tokens, - "host": 0, - } - - return None - def process_batch_result_prebuilt(self: Scheduler, batch: ScheduleBatch): assert self.disaggregation_mode == DisaggregationMode.DECODE use_free_group = self.server_args.disaggregation_decode_enable_radix_cache @@ -112,7 +58,7 @@ class SchedulerOutputProcessorMixin: release_kv_cache(req, self.tree_cache) # Note: Logprobs should be handled on the prefill engine. - self.stream_output(self.output_streamer, batch.reqs, batch.return_logprob) + self.output_streamer.stream_output(batch.reqs, batch.return_logprob) if use_free_group: self.token_to_kv_pool_allocator.free_group_end() @@ -411,8 +357,8 @@ class SchedulerOutputProcessorMixin: req.is_chunked -= 1 req.time_stats.set_last_chunked_prefill_finish_time() - self.stream_output( - self.output_streamer, batch.reqs, batch.return_logprob, skip_stream_req + self.output_streamer.stream_output( + batch.reqs, batch.return_logprob, skip_stream_req ) can_run_cuda_graph = getattr(result, "can_run_cuda_graph", False) @@ -478,8 +424,8 @@ class SchedulerOutputProcessorMixin: if result.copy_done is not None: result.copy_done.synchronize() - self._stream_output_generation( - self.output_streamer, batch.reqs, batch.return_logprob, is_idle_batch=True + self.output_streamer._stream_output_generation( + batch.reqs, batch.return_logprob, is_idle_batch=True ) def process_batch_result_decode( @@ -640,7 +586,7 @@ class SchedulerOutputProcessorMixin: self.abort_request(AbortReq(rid=req.rid)) req.grammar.finished = req.finished() - self.stream_output(self.output_streamer, batch.reqs, batch.return_logprob) + self.output_streamer.stream_output(batch.reqs, batch.return_logprob) self.token_to_kv_pool_allocator.free_group_end() self.metrics_reporter.forward_ct_decode = ( @@ -725,394 +671,3 @@ class SchedulerOutputProcessorMixin: req.mamba_last_track_seqlen = ( actual_seq_len // mamba_track_interval * mamba_track_interval ) - - @staticmethod - def stream_output( - self: "SchedulerOutputStreamer", - reqs: List[Req], - return_logprob: bool, - skip_req: Optional[Req] = None, - ): - """Stream the output to detokenizer.""" - if self.is_generation: - SchedulerOutputProcessorMixin._stream_output_generation( - self, reqs, return_logprob, skip_req - ) - else: # embedding or reward model - SchedulerOutputProcessorMixin._stream_output_embedding(self, reqs) - - if envs.SGLANG_TEST_CRASH_AFTER_STREAM_OUTPUTS.get() > 0: - SchedulerOutputProcessorMixin._trigger_crash_for_tests( - self, envs.SGLANG_TEST_CRASH_AFTER_STREAM_OUTPUTS.get() - ) - - @staticmethod - def _trigger_crash_for_tests(self: "SchedulerOutputStreamer", crash_threshold: int): - # Crash trigger: crash after stream_output is called N times - # This is used for testing purposes. - if not hasattr(self, "_test_stream_output_count"): - self._test_stream_output_count = 0 - self._test_stream_output_count += 1 - if self._test_stream_output_count >= crash_threshold: - raise RuntimeError( - f"Test crash after stream_output called {self._test_stream_output_count} times" - ) - - @staticmethod - def _stream_output_generation( - self: "SchedulerOutputStreamer", - reqs: List[Req], - return_logprob: bool, - 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 - 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: - continue - - if req.finished(): - if req.finished_output: - # With the overlap schedule, a request will try to output twice and hit this line twice - # because of the one additional delayed token. This "continue" prevented the dummy output. - continue - req.finished_output = True - if req.finished_len is None: - req.finished_len = len(req.output_ids) - should_output = True - else: - if req.stream: - stream_interval = ( - req.sampling_params.stream_interval - or self.server_args.stream_interval - ) - - # origin stream_interval logic - should_output = ( - len(req.output_ids) % stream_interval == 1 - if stream_interval > 1 - else len(req.output_ids) % stream_interval == 0 - ) - - if should_output: - # check_match_stop_str_prefix if tail_str's suffix match stop_str prefix - should_output &= not req.check_match_stop_str_prefix() - else: - should_output = ( - len(req.output_ids) % DEFAULT_FORCE_STREAM_INTERVAL == 0 - ) - - if should_output: - send_token_offset = req.send_token_offset - 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( - req.finished_reason.to_json() if req.finished_reason else None - ) - 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 :]) - - # 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:]) - req.send_token_offset = len(output_ids_) - skip_special_tokens.append(req.sampling_params.skip_special_tokens) - 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) - - # Collect detailed cache breakdown if available - cached_tokens_details.append( - SchedulerOutputProcessorMixin.get_cached_tokens_details(self, req) - ) - - retraction_counts.append(req.retraction_count) - - 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( - req.spec_correct_drafts_histogram - ) - - if return_logprob: - if ( - req.return_logprob - and not req.input_logprob_sent - # Decode server does not send input logprobs - and self.disaggregation_mode != DisaggregationMode.DECODE - # 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( - req.input_token_ids_logprobs_val - ) - 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([]) - - if req.return_logprob: - logprob_end = max(len(output_ids_), 1) - output_token_logprobs_val.append( - req.output_token_logprobs_val[ - send_output_token_logprobs_offset:logprob_end - ] - ) - output_token_logprobs_idx.append( - req.output_token_logprobs_idx[ - send_output_token_logprobs_offset:logprob_end - ] - ) - output_top_logprobs_val.append( - req.output_top_logprobs_val[ - send_output_token_logprobs_offset:logprob_end - ] - ) - output_top_logprobs_idx.append( - req.output_top_logprobs_idx[ - send_output_token_logprobs_offset:logprob_end - ] - ) - 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( - 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([]) - - if req.return_hidden_states: - if output_hidden_states is None: - output_hidden_states = [] - 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) - if req.return_indexer_topk: - if indexer_topk is None: - indexer_topk = [] - 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( - v[send_token_offset : len(output_ids_)] - ) - - if ( - req.finished() - and self.ps.attn_tp_rank == 0 - and self.server_args.enable_request_time_stats_logging - ): - req.log_time_stats() - - dp_ranks = [self.ps.dp_rank] * len(rids) if 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, - output_token_entropy_val=None, - output_hidden_states=output_hidden_states, - routed_experts=routed_experts, - indexer_topk=indexer_topk, - customized_info=customized_info, - placeholder_tokens_idx=None, - placeholder_tokens_val=None, - retraction_counts=retraction_counts, - load=load, - dp_ranks=dp_ranks, - ) - ) - - @staticmethod - def _stream_output_embedding(self: "SchedulerOutputStreamer", reqs: List[Req]): - rids = [] - http_worker_ipcs = [] - finished_reasons: List[BaseFinishReason] = [] - - embeddings = [] - prompt_tokens = [] - cached_tokens = [] - cached_tokens_details = [] # Detailed breakdown by cache source - time_stats = [] - retraction_counts = [] - phs_list = [] - has_phs = False - for req in reqs: - if req.finished(): - rids.append(req.rid) - http_worker_ipcs.append(req.http_worker_ipc) - finished_reasons.append(req.finished_reason.to_json()) - embeddings.append(req.embedding) - prompt_tokens.append(len(req.origin_input_ids)) - cached_tokens.append(req.cached_tokens) - - # Collect detailed cache breakdown if available - cached_tokens_details.append( - SchedulerOutputProcessorMixin.get_cached_tokens_details(self, req) - ) - time_stats.append(req.time_stats) - retraction_counts.append(req.retraction_count) - - phs = req.pooled_hidden_state - phs_list.append(phs) - if phs is not None: - has_phs = True - - # Optimize PHS for pickle: torch.stack reduces N __reduce_ex__ - # calls to 1 across the ZMQ IPC boundary. We can only stack when - # *every* entry is non-None (homogeneous batch); mixed batches - # (some requests want PHS, others don't) keep the raw list so - # positional indexing on the receiver side stays correct. - stacked_phs = None - if has_phs: - all_have_phs = all(t is not None for t in phs_list) - if all_have_phs: - if all(t.shape == phs_list[0].shape for t in phs_list): - stacked_phs = torch.stack(phs_list) - else: - stacked_phs = phs_list - else: - stacked_phs = phs_list - - self.send_to_detokenizer.send_output( - BatchEmbeddingOutput( - rids=rids, - http_worker_ipcs=http_worker_ipcs, - time_stats=time_stats, - finished_reasons=finished_reasons, - embeddings=embeddings, - prompt_tokens=prompt_tokens, - cached_tokens=cached_tokens, - cached_tokens_details=cached_tokens_details, - placeholder_tokens_idx=None, - placeholder_tokens_val=None, - retraction_counts=retraction_counts, - pooled_hidden_states=stacked_phs, - ) - ) diff --git a/test/registered/unit/managers/test_priority_scheduling_disaggregation.py b/test/registered/unit/managers/test_priority_scheduling_disaggregation.py index 172455598..3f287001c 100644 --- a/test/registered/unit/managers/test_priority_scheduling_disaggregation.py +++ b/test/registered/unit/managers/test_priority_scheduling_disaggregation.py @@ -138,7 +138,7 @@ class TestDecodePreallocQueuePriority(unittest.TestCase): scheduler.enable_hisparse = False scheduler.waiting_queue = [] scheduler.last_batch = None - scheduler.stream_output = MagicMock() + scheduler.output_streamer = MagicMock() queue.scheduler = scheduler return queue @@ -197,7 +197,7 @@ class TestDecodePreallocQueuePriority(unittest.TestCase): ) self.assertEqual([decode_req.req.rid for decode_req in failed], ["failed-low"]) self.assertEqual(queue.queue, []) - queue.scheduler.stream_output.assert_called_once_with( + queue.scheduler.output_streamer.stream_output.assert_called_once_with( [failed_low.req], failed_low.req.return_logprob ) diff --git a/test/registered/unit/mem_cache/test_decode_radix_lock_ref.py b/test/registered/unit/mem_cache/test_decode_radix_lock_ref.py index 863942000..ec67205d1 100644 --- a/test/registered/unit/mem_cache/test_decode_radix_lock_ref.py +++ b/test/registered/unit/mem_cache/test_decode_radix_lock_ref.py @@ -333,7 +333,7 @@ class TestDecodeLockRefScenarios(unittest.TestCase): scheduler.enable_hisparse = False scheduler.waiting_queue = [] scheduler.last_batch = None - scheduler.stream_output = MagicMock() + scheduler.output_streamer = MagicMock() queue.scheduler = scheduler # Initial budget says the request fits; post-lock budget says it does not.