Carve out SchedulerBatchResultProcessor for batch-result state (#25636)

This commit is contained in:
fzyzcjy
2026-05-18 18:44:41 +08:00
committed by GitHub
parent 18a7eb9e58
commit 7d0b0b6991
5 changed files with 170 additions and 42 deletions
+3 -1
View File
@@ -1654,7 +1654,9 @@ class SchedulerDisaggregationDecodeMixin:
new_prebuilt_batch = self.get_new_prebuilt_batch()
if new_prebuilt_batch:
assert self.chunked_req is None
self.process_batch_result_prebuilt(new_prebuilt_batch)
self.process_batch_result_prebuilt(
self.batch_result_processor, new_prebuilt_batch
)
new_prebuilt_batch.filter_batch()
if not new_prebuilt_batch.is_empty():
if self.running_batch.is_empty():
+2 -2
View File
@@ -535,7 +535,7 @@ class SchedulerDisaggregationPrefillMixin:
extend_logprob_start_len = extend_logprob_start_len_per_req[i]
extend_input_len = extend_input_len_per_req[i]
num_input_logprobs = extend_input_len - extend_logprob_start_len
self.logprob_result_processor.add_logprob_return_values(
self.batch_result_processor.logprob_result_processor.add_logprob_return_values(
i,
req,
logprob_pt,
@@ -572,7 +572,7 @@ class SchedulerDisaggregationPrefillMixin:
if extend_logprob_start_len < extend_input_len:
# Update input logprobs.
num_input_logprobs = extend_input_len - extend_logprob_start_len
self.logprob_result_processor.add_input_logprob_return_values(
self.batch_result_processor.logprob_result_processor.add_input_logprob_return_values(
i,
req,
logits_output,
+32 -9
View File
@@ -164,6 +164,9 @@ from sglang.srt.managers.schedule_policy import (
PrefillAdder,
SchedulePolicy,
)
from sglang.srt.managers.scheduler_components.batch_result_processor import (
SchedulerBatchResultProcessor,
)
from sglang.srt.managers.scheduler_components.dp_attn import (
SchedulerDPAttnAdapter,
)
@@ -740,11 +743,6 @@ class Scheduler(
get_spec_total_num_forward_ct=lambda: self.metrics_reporter.spec_total_num_forward_ct,
)
self.logprob_result_processor = SchedulerLogprobResultProcessor(
server_args=self.server_args,
model_config=self.model_config,
)
self.output_streamer = SchedulerOutputStreamer(
send_to_detokenizer=self.send_to_detokenizer,
tree_cache=self.tree_cache,
@@ -757,6 +755,29 @@ class Scheduler(
load_inquirer_get_loads=lambda req: self.load_inquirer.get_loads(req),
)
self.batch_result_processor = SchedulerBatchResultProcessor(
is_generation=self.is_generation,
disaggregation_mode=self.disaggregation_mode,
enable_overlap=self.enable_overlap,
enable_overlap_mlx=self.enable_overlap_mlx,
server_args=self.server_args,
model_config=self.model_config,
token_to_kv_pool_allocator=self.token_to_kv_pool_allocator,
tree_cache=self.tree_cache,
hisparse_coordinator=self.hisparse_coordinator,
req_to_token_pool=self.req_to_token_pool,
decode_offload_manager=self.decode_offload_manager,
metrics_collector=self.metrics_collector,
metrics_reporter=self.metrics_reporter,
draft_worker=self.draft_worker,
model_worker=self.model_worker,
logprob_result_processor=SchedulerLogprobResultProcessor(
server_args=self.server_args, model_config=self.model_config
),
output_streamer=self.output_streamer,
abort_request=self.abort_request,
)
self.is_initializing = False
def init_zbal_on_npu(self):
@@ -3069,18 +3090,20 @@ class Scheduler(
result: Union[GenerationBatchResult, EmbeddingBatchResult],
):
if batch.forward_mode.is_decode():
self.process_batch_result_decode(batch, result)
self.process_batch_result_decode(self.batch_result_processor, batch, result)
elif batch.forward_mode.is_extend():
if batch.is_dllm():
self.process_batch_result_dllm(batch, result)
elif self.disaggregation_mode == DisaggregationMode.PREFILL:
self.process_batch_result_disagg_prefill(batch, result)
else:
self.process_batch_result_prefill(batch, result)
self.process_batch_result_prefill(
self.batch_result_processor, batch, result
)
elif batch.forward_mode.is_prebuilt():
self.process_batch_result_prebuilt(batch)
self.process_batch_result_prebuilt(self.batch_result_processor, batch)
elif batch.forward_mode.is_idle():
self.process_batch_result_idle(batch, result)
self.process_batch_result_idle(self.batch_result_processor, batch, result)
self.metrics_reporter.log_batch_result_stats(batch, result)
@@ -0,0 +1,54 @@
from __future__ import annotations
import logging
from dataclasses import dataclass
from typing import TYPE_CHECKING, Callable, Optional
from sglang.srt.disaggregation.utils import DisaggregationMode
if TYPE_CHECKING:
from sglang.srt.configs.model_config import ModelConfig
from sglang.srt.disaggregation.decode_kvcache_offload_manager import (
DecodeKVCacheOffloadManager,
)
from sglang.srt.managers.hisparse_coordinator import HiSparseCoordinator
from sglang.srt.managers.scheduler_components.logprob_result_processor import (
SchedulerLogprobResultProcessor,
)
from sglang.srt.managers.scheduler_components.metrics_reporter import (
SchedulerMetricsReporter,
)
from sglang.srt.managers.scheduler_components.output_streamer import (
SchedulerOutputStreamer,
)
from sglang.srt.managers.tp_worker import BaseTpWorker
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
from sglang.srt.observability.metrics_collector import SchedulerMetricsCollector
from sglang.srt.server_args import ServerArgs
logger = logging.getLogger(__name__)
@dataclass(kw_only=True, slots=True, frozen=True)
class SchedulerBatchResultProcessor:
is_generation: bool
disaggregation_mode: "DisaggregationMode"
enable_overlap: bool
enable_overlap_mlx: bool
server_args: "ServerArgs"
model_config: "ModelConfig"
token_to_kv_pool_allocator: "BaseTokenToKVPoolAllocator"
tree_cache: "BasePrefixCache"
hisparse_coordinator: Optional["HiSparseCoordinator"]
req_to_token_pool: "ReqToTokenPool"
decode_offload_manager: Optional["DecodeKVCacheOffloadManager"]
metrics_collector: "SchedulerMetricsCollector"
metrics_reporter: "SchedulerMetricsReporter"
draft_worker: "BaseTpWorker"
model_worker: "BaseTpWorker"
logprob_result_processor: "SchedulerLogprobResultProcessor"
output_streamer: "SchedulerOutputStreamer"
abort_request: Callable
@@ -27,7 +27,9 @@ if TYPE_CHECKING:
EmbeddingBatchResult,
GenerationBatchResult,
ScheduleBatch,
Scheduler,
)
from sglang.srt.managers.scheduler_components.batch_result_processor import (
SchedulerBatchResultProcessor,
)
logger = logging.getLogger(__name__)
@@ -43,7 +45,10 @@ class SchedulerOutputProcessorMixin:
We put them into a separate file to make the `scheduler.py` shorter.
"""
def process_batch_result_prebuilt(self: Scheduler, batch: ScheduleBatch):
@staticmethod
def process_batch_result_prebuilt(
self: "SchedulerBatchResultProcessor", batch: ScheduleBatch
):
assert self.disaggregation_mode == DisaggregationMode.DECODE
use_free_group = self.server_args.disaggregation_decode_enable_radix_cache
if use_free_group:
@@ -53,7 +58,7 @@ class SchedulerOutputProcessorMixin:
req.check_finished()
if req.finished():
req.time_stats.set_quick_finish_time()
if self.enable_hisparse:
if self.server_args.enable_hisparse:
self.hisparse_coordinator.request_finished(req)
release_kv_cache(req, self.tree_cache)
@@ -62,7 +67,8 @@ class SchedulerOutputProcessorMixin:
if use_free_group:
self.token_to_kv_pool_allocator.free_group_end()
def maybe_collect_routed_experts(self: Scheduler, req: Req):
@staticmethod
def _maybe_collect_routed_experts(self: "SchedulerBatchResultProcessor", req: Req):
"""Collect routed experts for a finished request.
Returns immediately if `return_routed_experts` was not set on the
@@ -105,7 +111,8 @@ class SchedulerOutputProcessorMixin:
req.routed_experts_start_len,
)
def maybe_collect_indexer_topk(self: Scheduler, req: Req):
@staticmethod
def _maybe_collect_indexer_topk(self: "SchedulerBatchResultProcessor", req: Req):
capturer = get_global_indexer_capturer()
if capturer is None:
return
@@ -115,8 +122,12 @@ class SchedulerOutputProcessorMixin:
req_to_token_pool=self.req_to_token_pool,
)
def maybe_collect_customized_info(
self: Scheduler, i: int, req: Req, logits_output: LogitsProcessorOutput
@staticmethod
def _maybe_collect_customized_info(
self: "SchedulerBatchResultProcessor",
i: int,
req: Req,
logits_output: LogitsProcessorOutput,
):
if logits_output is not None and logits_output.customized_info is not None:
if req.customized_info is None:
@@ -133,8 +144,9 @@ class SchedulerOutputProcessorMixin:
elem = elem.copy()
req.customized_info[k].append(elem)
@staticmethod
def process_batch_result_prefill(
self: Scheduler,
self: "SchedulerBatchResultProcessor",
batch: ScheduleBatch,
result: Union[GenerationBatchResult, EmbeddingBatchResult],
):
@@ -202,20 +214,28 @@ class SchedulerOutputProcessorMixin:
# req output_ids are set here
req.output_ids.append(next_token_id)
self._maybe_update_reasoning_tokens(req, next_token_id)
SchedulerOutputProcessorMixin._maybe_update_reasoning_tokens(
self, req, next_token_id
)
req.check_finished()
if req.finished():
self.maybe_collect_routed_experts(req)
self.maybe_collect_indexer_topk(req)
SchedulerOutputProcessorMixin._maybe_collect_routed_experts(
self, req
)
SchedulerOutputProcessorMixin._maybe_collect_indexer_topk(
self, req
)
release_kv_cache(req, self.tree_cache)
req.time_stats.set_completion_time()
elif not batch.decoding_reqs or req not in batch.decoding_reqs:
maybe_cache_unfinished_req(req, self.tree_cache)
if self.enable_hisparse:
if self.server_args.enable_hisparse:
self.hisparse_coordinator.admit_request_into_staging(req)
self.maybe_collect_customized_info(i, req, logits_output)
SchedulerOutputProcessorMixin._maybe_collect_customized_info(
self, i, req, logits_output
)
if batch.return_logprob:
assert extend_logprob_start_len_per_req is not None
@@ -369,8 +389,11 @@ class SchedulerOutputProcessorMixin:
dp_cooperation_info=batch.dp_cooperation_info,
)
@staticmethod
def _resolve_spec_overlap_tokens(
self: Scheduler, result: GenerationBatchResult, batch: ScheduleBatch
self: "SchedulerBatchResultProcessor",
result: GenerationBatchResult,
batch: ScheduleBatch,
) -> List[List[int]]:
"""Resolve the padding next token ids for speculative decoding with overlap."""
assert result.next_token_ids.is_cpu
@@ -416,8 +439,9 @@ class SchedulerOutputProcessorMixin:
return predict_tokens
@staticmethod
def process_batch_result_idle(
self: Scheduler,
self: "SchedulerBatchResultProcessor",
batch: ScheduleBatch,
result: GenerationBatchResult,
):
@@ -428,8 +452,9 @@ class SchedulerOutputProcessorMixin:
batch.reqs, batch.return_logprob, is_idle_batch=True
)
@staticmethod
def process_batch_result_decode(
self: Scheduler,
self: "SchedulerBatchResultProcessor",
batch: ScheduleBatch,
result: GenerationBatchResult,
):
@@ -450,7 +475,11 @@ class SchedulerOutputProcessorMixin:
if batch.spec_algorithm.is_none() or batch.is_spec_v2:
if batch.is_spec_v2:
next_token_ids = self._resolve_spec_overlap_tokens(result, batch)
next_token_ids = (
SchedulerOutputProcessorMixin._resolve_spec_overlap_tokens(
self, result, batch
)
)
elif isinstance(next_token_ids, list):
pass # MLX path: already a list[int], skip torch round-trip
else:
@@ -479,7 +508,7 @@ class SchedulerOutputProcessorMixin:
self.metrics_reporter.update_spec_metrics(
batch.batch_size(), result.num_correct_drafts
)
if self.metrics_reporter.enable_metrics:
if self.server_args.enable_metrics:
self.metrics_collector.increment_decode_cuda_graph_pass(
value=can_run_cuda_graph
)
@@ -501,9 +530,13 @@ class SchedulerOutputProcessorMixin:
continue
if is_spec_v1:
self._mamba_prefix_cache_update(req, batch, result, i)
SchedulerOutputProcessorMixin._mamba_prefix_cache_update(
self, req, batch, result, i
)
req.time_stats.set_last_decode_finish_time()
self._handle_finished_req(req, i, logits_output)
SchedulerOutputProcessorMixin._handle_finished_req(
self, req, i, logits_output
)
if req.return_hidden_states and logits_output.hidden_states is not None:
req.hidden_states.append(
logits_output.hidden_states[i].cpu().clone().tolist()
@@ -521,14 +554,20 @@ class SchedulerOutputProcessorMixin:
req.output_ids.extend(next_token_id)
new_accepted_len = len(next_token_id)
self._maybe_update_reasoning_tokens(req, next_token_id)
SchedulerOutputProcessorMixin._maybe_update_reasoning_tokens(
self, req, next_token_id
)
# Update Mamba last track seqlen
self._mamba_prefix_cache_update(req, batch, result, i)
SchedulerOutputProcessorMixin._mamba_prefix_cache_update(
self, req, batch, result, i
)
req.time_stats.set_last_decode_finish_time()
req.check_finished(new_accepted_len)
self._handle_finished_req(req, i, logits_output)
SchedulerOutputProcessorMixin._handle_finished_req(
self, req, i, logits_output
)
if req.return_logprob:
# Spec v1 handles logprobs inside its own worker.
@@ -598,8 +637,12 @@ class SchedulerOutputProcessorMixin:
num_correct_drafts=result.num_correct_drafts,
)
@staticmethod
def _handle_finished_req(
self: Scheduler, req: Req, i: int, logits_output: LogitsProcessorOutput
self: "SchedulerBatchResultProcessor",
req: Req,
i: int,
logits_output: LogitsProcessorOutput,
):
if (
self.server_args.disaggregation_decode_enable_offload_kvcache
@@ -611,31 +654,37 @@ class SchedulerOutputProcessorMixin:
# delete feature to save memory
if req.multimodal_inputs is not None and req.session is None:
req.multimodal_inputs.release_features()
self.maybe_collect_routed_experts(req)
self.maybe_collect_indexer_topk(req)
SchedulerOutputProcessorMixin._maybe_collect_routed_experts(self, req)
SchedulerOutputProcessorMixin._maybe_collect_indexer_topk(self, req)
if self.server_args.disaggregation_decode_enable_offload_kvcache:
# Asynchronously offload KV cache; release_kv_cache will be called after Device->Host transfer completes
if not self.decode_offload_manager.offload_kv_cache(req):
self.decode_offload_manager.finalize_release_on_finish(req)
else:
if self.enable_hisparse:
if self.server_args.enable_hisparse:
self.hisparse_coordinator.request_finished(req)
release_kv_cache(req, self.tree_cache)
req.time_stats.set_completion_time()
self.maybe_collect_customized_info(i, req, logits_output)
SchedulerOutputProcessorMixin._maybe_collect_customized_info(
self, i, req, logits_output
)
@staticmethod
def _maybe_update_reasoning_tokens(
self: Scheduler, req: Req, next_token_id: Union[int, List[int]]
self: "SchedulerBatchResultProcessor",
req: Req,
next_token_id: Union[int, List[int]],
):
think_end_id = self.model_config.think_end_id
if req.require_reasoning and think_end_id is not None:
req.update_reasoning_tokens(next_token_id, think_end_id)
@staticmethod
def _mamba_prefix_cache_update(
self: Scheduler,
self: "SchedulerBatchResultProcessor",
req: Req,
batch: ScheduleBatch,
result: GenerationBatchResult,