Move output streaming to SchedulerOutputStreamer (#25635)

This commit is contained in:
fzyzcjy
2026-05-18 18:44:07 +08:00
committed by GitHub
parent dc88b4eeb4
commit 18a7eb9e58
8 changed files with 459 additions and 481 deletions
+5 -12
View File
@@ -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,
)
+3 -8
View File
@@ -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,
+1 -1
View File
@@ -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)
+2 -4
View File
@@ -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 (
@@ -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,
)
)
@@ -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,
)
)
@@ -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
)
@@ -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.