Move output streaming to SchedulerOutputStreamer (#25635)
This commit is contained in:
@@ -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}"
|
message = f"Request {req.rid} exceeds the maximum number of tokens: {len(req.origin_input_ids)} > {self.max_total_num_tokens}"
|
||||||
logger.error(message)
|
logger.error(message)
|
||||||
prepare_abort(req, message, status_code=HTTPStatus.BAD_REQUEST)
|
prepare_abort(req, message, status_code=HTTPStatus.BAD_REQUEST)
|
||||||
self.scheduler.stream_output(
|
self.scheduler.output_streamer.stream_output([req], req.return_logprob)
|
||||||
self.scheduler.output_streamer, [req], req.return_logprob
|
|
||||||
)
|
|
||||||
return True
|
return True
|
||||||
if self._uses_swa_tail_prealloc():
|
if self._uses_swa_tail_prealloc():
|
||||||
_, swa_required = self._prealloc_required_tokens(req)
|
_, swa_required = self._prealloc_required_tokens(req)
|
||||||
@@ -546,9 +544,7 @@ class DecodePreallocQueue:
|
|||||||
)
|
)
|
||||||
logger.error(message)
|
logger.error(message)
|
||||||
prepare_abort(req, message, status_code=HTTPStatus.BAD_REQUEST)
|
prepare_abort(req, message, status_code=HTTPStatus.BAD_REQUEST)
|
||||||
self.scheduler.stream_output(
|
self.scheduler.output_streamer.stream_output([req], req.return_logprob)
|
||||||
self.scheduler.output_streamer, [req], req.return_logprob
|
|
||||||
)
|
|
||||||
return True
|
return True
|
||||||
return False
|
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:
|
if rids_to_check is not None and decode_req.req.rid not in rids_to_check:
|
||||||
continue
|
continue
|
||||||
if isinstance(decode_req.req.finished_reason, FINISH_ABORT):
|
if isinstance(decode_req.req.finished_reason, FINISH_ABORT):
|
||||||
self.scheduler.stream_output(
|
self.scheduler.output_streamer.stream_output(
|
||||||
self.scheduler.output_streamer,
|
|
||||||
[decode_req.req],
|
[decode_req.req],
|
||||||
decode_req.req.return_logprob,
|
decode_req.req.return_logprob,
|
||||||
)
|
)
|
||||||
@@ -1514,8 +1509,7 @@ class DecodeTransferQueue:
|
|||||||
error_message,
|
error_message,
|
||||||
status_code=HTTPStatus.INTERNAL_SERVER_ERROR,
|
status_code=HTTPStatus.INTERNAL_SERVER_ERROR,
|
||||||
)
|
)
|
||||||
self.scheduler.stream_output(
|
self.scheduler.output_streamer.stream_output(
|
||||||
self.scheduler.output_streamer,
|
|
||||||
[decode_req.req],
|
[decode_req.req],
|
||||||
decode_req.req.return_logprob,
|
decode_req.req.return_logprob,
|
||||||
)
|
)
|
||||||
@@ -1533,8 +1527,7 @@ class DecodeTransferQueue:
|
|||||||
indices_to_remove.add(i)
|
indices_to_remove.add(i)
|
||||||
# Check if request was aborted due to corruption
|
# Check if request was aborted due to corruption
|
||||||
if isinstance(decode_req.req.finished_reason, FINISH_ABORT):
|
if isinstance(decode_req.req.finished_reason, FINISH_ABORT):
|
||||||
self.scheduler.stream_output(
|
self.scheduler.output_streamer.stream_output(
|
||||||
self.scheduler.output_streamer,
|
|
||||||
[decode_req.req],
|
[decode_req.req],
|
||||||
decode_req.req.return_logprob,
|
decode_req.req.return_logprob,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -252,9 +252,7 @@ class PrefillBootstrapQueue:
|
|||||||
logger.error(message)
|
logger.error(message)
|
||||||
req.time_stats.trace_ctx.abort(abort_info={"reason": message})
|
req.time_stats.trace_ctx.abort(abort_info={"reason": message})
|
||||||
prepare_abort(req, message, status_code=HTTPStatus.BAD_REQUEST)
|
prepare_abort(req, message, status_code=HTTPStatus.BAD_REQUEST)
|
||||||
self.scheduler.stream_output(
|
self.scheduler.output_streamer.stream_output([req], req.return_logprob)
|
||||||
self.scheduler.output_streamer, [req], req.return_logprob
|
|
||||||
)
|
|
||||||
return True
|
return True
|
||||||
return False
|
return False
|
||||||
|
|
||||||
@@ -311,9 +309,7 @@ class PrefillBootstrapQueue:
|
|||||||
prepare_abort(
|
prepare_abort(
|
||||||
req, error_message, status_code=HTTPStatus.INTERNAL_SERVER_ERROR
|
req, error_message, status_code=HTTPStatus.INTERNAL_SERVER_ERROR
|
||||||
)
|
)
|
||||||
self.scheduler.stream_output(
|
self.scheduler.output_streamer.stream_output([req], req.return_logprob)
|
||||||
self.scheduler.output_streamer, [req], req.return_logprob
|
|
||||||
)
|
|
||||||
indices_to_remove.add(i)
|
indices_to_remove.add(i)
|
||||||
failed_reqs.append(req)
|
failed_reqs.append(req)
|
||||||
if self.scheduler.metrics_reporter.enable_metrics:
|
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"]
|
self.metrics_reporter.kv_transfer_speed_gb_s = metrics["speed_gb_s"]
|
||||||
|
|
||||||
# Stream requests which have finished transfer
|
# Stream requests which have finished transfer
|
||||||
self.stream_output(
|
self.output_streamer.stream_output(
|
||||||
self.output_streamer,
|
|
||||||
done_reqs,
|
done_reqs,
|
||||||
any(req.return_logprob for req in done_reqs),
|
any(req.return_logprob for req in done_reqs),
|
||||||
None,
|
None,
|
||||||
|
|||||||
@@ -89,7 +89,7 @@ class SchedulerDllmMixin:
|
|||||||
release_kv_cache(req, self.tree_cache)
|
release_kv_cache(req, self.tree_cache)
|
||||||
req.time_stats.set_completion_time()
|
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()
|
self.token_to_kv_pool_allocator.free_group_end()
|
||||||
|
|
||||||
can_run_cuda_graph = getattr(result, "can_run_cuda_graph", False)
|
can_run_cuda_graph = getattr(result, "can_run_cuda_graph", False)
|
||||||
|
|||||||
@@ -651,9 +651,7 @@ class Scheduler(
|
|||||||
server_args=self.server_args,
|
server_args=self.server_args,
|
||||||
model_config=self.model_config,
|
model_config=self.model_config,
|
||||||
max_recv_per_poll=self.max_recv_per_poll,
|
max_recv_per_poll=self.max_recv_per_poll,
|
||||||
stream_output=lambda *a, **kw: self.stream_output(
|
stream_output=lambda *a, **kw: self.output_streamer.stream_output(*a, **kw),
|
||||||
self.output_streamer, *a, **kw
|
|
||||||
),
|
|
||||||
get_last_forward_mode=lambda: (
|
get_last_forward_mode=lambda: (
|
||||||
self.last_batch.forward_mode if self.last_batch is not None else None
|
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}
|
abort_info={"reason": error_msg}
|
||||||
)
|
)
|
||||||
prepare_abort(req, error_msg, status_code=HTTPStatus.BAD_REQUEST)
|
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
|
return
|
||||||
|
|
||||||
elif (
|
elif (
|
||||||
|
|||||||
@@ -2,13 +2,28 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import logging
|
import logging
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from typing import Any, Callable
|
from typing import (
|
||||||
|
Any,
|
||||||
|
Callable,
|
||||||
|
List,
|
||||||
|
Optional,
|
||||||
|
)
|
||||||
|
|
||||||
|
import torch
|
||||||
import zmq
|
import zmq
|
||||||
|
|
||||||
from sglang.srt.disaggregation.utils import DisaggregationMode
|
from sglang.srt.disaggregation.utils import DisaggregationMode
|
||||||
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
|
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
|
||||||
from sglang.srt.environ import envs
|
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.mem_cache.base_prefix_cache import BasePrefixCache
|
||||||
from sglang.srt.server_args import ServerArgs
|
from sglang.srt.server_args import ServerArgs
|
||||||
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
||||||
@@ -31,3 +46,425 @@ class SchedulerOutputStreamer:
|
|||||||
enable_hicache_storage: Callable[[], bool]
|
enable_hicache_storage: Callable[[], bool]
|
||||||
load_inquirer_get_loads: Callable[..., Any]
|
load_inquirer_get_loads: Callable[..., Any]
|
||||||
_test_stream_output_count: int = 0
|
_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
|
from __future__ import annotations
|
||||||
|
|
||||||
import logging
|
import logging
|
||||||
from typing import TYPE_CHECKING, List, Optional, Union
|
from typing import TYPE_CHECKING, List, Union
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
@@ -10,12 +10,8 @@ from sglang.srt.environ import envs
|
|||||||
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
|
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
|
||||||
from sglang.srt.managers.io_struct import (
|
from sglang.srt.managers.io_struct import (
|
||||||
AbortReq,
|
AbortReq,
|
||||||
BatchEmbeddingOutput,
|
|
||||||
BatchTokenIDOutput,
|
|
||||||
GetLoadsReqInput,
|
|
||||||
)
|
)
|
||||||
from sglang.srt.managers.schedule_batch import (
|
from sglang.srt.managers.schedule_batch import (
|
||||||
BaseFinishReason,
|
|
||||||
Req,
|
Req,
|
||||||
ScheduleBatch,
|
ScheduleBatch,
|
||||||
)
|
)
|
||||||
@@ -33,9 +29,6 @@ if TYPE_CHECKING:
|
|||||||
ScheduleBatch,
|
ScheduleBatch,
|
||||||
Scheduler,
|
Scheduler,
|
||||||
)
|
)
|
||||||
from sglang.srt.managers.scheduler_components.output_streamer import (
|
|
||||||
SchedulerOutputStreamer,
|
|
||||||
)
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -50,53 +43,6 @@ class SchedulerOutputProcessorMixin:
|
|||||||
We put them into a separate file to make the `scheduler.py` shorter.
|
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):
|
def process_batch_result_prebuilt(self: Scheduler, batch: ScheduleBatch):
|
||||||
assert self.disaggregation_mode == DisaggregationMode.DECODE
|
assert self.disaggregation_mode == DisaggregationMode.DECODE
|
||||||
use_free_group = self.server_args.disaggregation_decode_enable_radix_cache
|
use_free_group = self.server_args.disaggregation_decode_enable_radix_cache
|
||||||
@@ -112,7 +58,7 @@ class SchedulerOutputProcessorMixin:
|
|||||||
release_kv_cache(req, self.tree_cache)
|
release_kv_cache(req, self.tree_cache)
|
||||||
|
|
||||||
# Note: Logprobs should be handled on the prefill engine.
|
# 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:
|
if use_free_group:
|
||||||
self.token_to_kv_pool_allocator.free_group_end()
|
self.token_to_kv_pool_allocator.free_group_end()
|
||||||
|
|
||||||
@@ -411,8 +357,8 @@ class SchedulerOutputProcessorMixin:
|
|||||||
req.is_chunked -= 1
|
req.is_chunked -= 1
|
||||||
req.time_stats.set_last_chunked_prefill_finish_time()
|
req.time_stats.set_last_chunked_prefill_finish_time()
|
||||||
|
|
||||||
self.stream_output(
|
self.output_streamer.stream_output(
|
||||||
self.output_streamer, batch.reqs, batch.return_logprob, skip_stream_req
|
batch.reqs, batch.return_logprob, skip_stream_req
|
||||||
)
|
)
|
||||||
|
|
||||||
can_run_cuda_graph = getattr(result, "can_run_cuda_graph", False)
|
can_run_cuda_graph = getattr(result, "can_run_cuda_graph", False)
|
||||||
@@ -478,8 +424,8 @@ class SchedulerOutputProcessorMixin:
|
|||||||
if result.copy_done is not None:
|
if result.copy_done is not None:
|
||||||
result.copy_done.synchronize()
|
result.copy_done.synchronize()
|
||||||
|
|
||||||
self._stream_output_generation(
|
self.output_streamer._stream_output_generation(
|
||||||
self.output_streamer, batch.reqs, batch.return_logprob, is_idle_batch=True
|
batch.reqs, batch.return_logprob, is_idle_batch=True
|
||||||
)
|
)
|
||||||
|
|
||||||
def process_batch_result_decode(
|
def process_batch_result_decode(
|
||||||
@@ -640,7 +586,7 @@ class SchedulerOutputProcessorMixin:
|
|||||||
self.abort_request(AbortReq(rid=req.rid))
|
self.abort_request(AbortReq(rid=req.rid))
|
||||||
req.grammar.finished = req.finished()
|
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.token_to_kv_pool_allocator.free_group_end()
|
||||||
|
|
||||||
self.metrics_reporter.forward_ct_decode = (
|
self.metrics_reporter.forward_ct_decode = (
|
||||||
@@ -725,394 +671,3 @@ class SchedulerOutputProcessorMixin:
|
|||||||
req.mamba_last_track_seqlen = (
|
req.mamba_last_track_seqlen = (
|
||||||
actual_seq_len // mamba_track_interval * mamba_track_interval
|
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.enable_hisparse = False
|
||||||
scheduler.waiting_queue = []
|
scheduler.waiting_queue = []
|
||||||
scheduler.last_batch = None
|
scheduler.last_batch = None
|
||||||
scheduler.stream_output = MagicMock()
|
scheduler.output_streamer = MagicMock()
|
||||||
queue.scheduler = scheduler
|
queue.scheduler = scheduler
|
||||||
return queue
|
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([decode_req.req.rid for decode_req in failed], ["failed-low"])
|
||||||
self.assertEqual(queue.queue, [])
|
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
|
[failed_low.req], failed_low.req.return_logprob
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -333,7 +333,7 @@ class TestDecodeLockRefScenarios(unittest.TestCase):
|
|||||||
scheduler.enable_hisparse = False
|
scheduler.enable_hisparse = False
|
||||||
scheduler.waiting_queue = []
|
scheduler.waiting_queue = []
|
||||||
scheduler.last_batch = None
|
scheduler.last_batch = None
|
||||||
scheduler.stream_output = MagicMock()
|
scheduler.output_streamer = MagicMock()
|
||||||
queue.scheduler = scheduler
|
queue.scheduler = scheduler
|
||||||
|
|
||||||
# Initial budget says the request fits; post-lock budget says it does not.
|
# Initial budget says the request fits; post-lock budget says it does not.
|
||||||
|
|||||||
Reference in New Issue
Block a user