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() new_prebuilt_batch = self.get_new_prebuilt_batch()
if new_prebuilt_batch: if new_prebuilt_batch:
assert self.chunked_req is None 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() new_prebuilt_batch.filter_batch()
if not new_prebuilt_batch.is_empty(): if not new_prebuilt_batch.is_empty():
if self.running_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_logprob_start_len = extend_logprob_start_len_per_req[i]
extend_input_len = extend_input_len_per_req[i] extend_input_len = extend_input_len_per_req[i]
num_input_logprobs = extend_input_len - extend_logprob_start_len 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, i,
req, req,
logprob_pt, logprob_pt,
@@ -572,7 +572,7 @@ class SchedulerDisaggregationPrefillMixin:
if extend_logprob_start_len < extend_input_len: if extend_logprob_start_len < extend_input_len:
# Update input logprobs. # Update input logprobs.
num_input_logprobs = extend_input_len - extend_logprob_start_len 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, i,
req, req,
logits_output, logits_output,
+32 -9
View File
@@ -164,6 +164,9 @@ from sglang.srt.managers.schedule_policy import (
PrefillAdder, PrefillAdder,
SchedulePolicy, SchedulePolicy,
) )
from sglang.srt.managers.scheduler_components.batch_result_processor import (
SchedulerBatchResultProcessor,
)
from sglang.srt.managers.scheduler_components.dp_attn import ( from sglang.srt.managers.scheduler_components.dp_attn import (
SchedulerDPAttnAdapter, SchedulerDPAttnAdapter,
) )
@@ -740,11 +743,6 @@ class Scheduler(
get_spec_total_num_forward_ct=lambda: self.metrics_reporter.spec_total_num_forward_ct, 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( self.output_streamer = SchedulerOutputStreamer(
send_to_detokenizer=self.send_to_detokenizer, send_to_detokenizer=self.send_to_detokenizer,
tree_cache=self.tree_cache, tree_cache=self.tree_cache,
@@ -757,6 +755,29 @@ class Scheduler(
load_inquirer_get_loads=lambda req: self.load_inquirer.get_loads(req), 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 self.is_initializing = False
def init_zbal_on_npu(self): def init_zbal_on_npu(self):
@@ -3069,18 +3090,20 @@ class Scheduler(
result: Union[GenerationBatchResult, EmbeddingBatchResult], result: Union[GenerationBatchResult, EmbeddingBatchResult],
): ):
if batch.forward_mode.is_decode(): 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(): elif batch.forward_mode.is_extend():
if batch.is_dllm(): if batch.is_dllm():
self.process_batch_result_dllm(batch, result) self.process_batch_result_dllm(batch, result)
elif self.disaggregation_mode == DisaggregationMode.PREFILL: elif self.disaggregation_mode == DisaggregationMode.PREFILL:
self.process_batch_result_disagg_prefill(batch, result) self.process_batch_result_disagg_prefill(batch, result)
else: 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(): 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(): 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) 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, EmbeddingBatchResult,
GenerationBatchResult, GenerationBatchResult,
ScheduleBatch, ScheduleBatch,
Scheduler, )
from sglang.srt.managers.scheduler_components.batch_result_processor import (
SchedulerBatchResultProcessor,
) )
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -43,7 +45,10 @@ 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.
""" """
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 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
if use_free_group: if use_free_group:
@@ -53,7 +58,7 @@ class SchedulerOutputProcessorMixin:
req.check_finished() req.check_finished()
if req.finished(): if req.finished():
req.time_stats.set_quick_finish_time() req.time_stats.set_quick_finish_time()
if self.enable_hisparse: if self.server_args.enable_hisparse:
self.hisparse_coordinator.request_finished(req) self.hisparse_coordinator.request_finished(req)
release_kv_cache(req, self.tree_cache) release_kv_cache(req, self.tree_cache)
@@ -62,7 +67,8 @@ class SchedulerOutputProcessorMixin:
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()
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. """Collect routed experts for a finished request.
Returns immediately if `return_routed_experts` was not set on the Returns immediately if `return_routed_experts` was not set on the
@@ -105,7 +111,8 @@ class SchedulerOutputProcessorMixin:
req.routed_experts_start_len, 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() capturer = get_global_indexer_capturer()
if capturer is None: if capturer is None:
return return
@@ -115,8 +122,12 @@ class SchedulerOutputProcessorMixin:
req_to_token_pool=self.req_to_token_pool, req_to_token_pool=self.req_to_token_pool,
) )
def maybe_collect_customized_info( @staticmethod
self: Scheduler, i: int, req: Req, logits_output: LogitsProcessorOutput 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 logits_output is not None and logits_output.customized_info is not None:
if req.customized_info is None: if req.customized_info is None:
@@ -133,8 +144,9 @@ class SchedulerOutputProcessorMixin:
elem = elem.copy() elem = elem.copy()
req.customized_info[k].append(elem) req.customized_info[k].append(elem)
@staticmethod
def process_batch_result_prefill( def process_batch_result_prefill(
self: Scheduler, self: "SchedulerBatchResultProcessor",
batch: ScheduleBatch, batch: ScheduleBatch,
result: Union[GenerationBatchResult, EmbeddingBatchResult], result: Union[GenerationBatchResult, EmbeddingBatchResult],
): ):
@@ -202,20 +214,28 @@ class SchedulerOutputProcessorMixin:
# req output_ids are set here # req output_ids are set here
req.output_ids.append(next_token_id) 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() req.check_finished()
if req.finished(): if req.finished():
self.maybe_collect_routed_experts(req) SchedulerOutputProcessorMixin._maybe_collect_routed_experts(
self.maybe_collect_indexer_topk(req) self, req
)
SchedulerOutputProcessorMixin._maybe_collect_indexer_topk(
self, req
)
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()
elif not batch.decoding_reqs or req not in batch.decoding_reqs: elif not batch.decoding_reqs or req not in batch.decoding_reqs:
maybe_cache_unfinished_req(req, self.tree_cache) 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.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: if batch.return_logprob:
assert extend_logprob_start_len_per_req is not None assert extend_logprob_start_len_per_req is not None
@@ -369,8 +389,11 @@ class SchedulerOutputProcessorMixin:
dp_cooperation_info=batch.dp_cooperation_info, dp_cooperation_info=batch.dp_cooperation_info,
) )
@staticmethod
def _resolve_spec_overlap_tokens( def _resolve_spec_overlap_tokens(
self: Scheduler, result: GenerationBatchResult, batch: ScheduleBatch self: "SchedulerBatchResultProcessor",
result: GenerationBatchResult,
batch: ScheduleBatch,
) -> List[List[int]]: ) -> List[List[int]]:
"""Resolve the padding next token ids for speculative decoding with overlap.""" """Resolve the padding next token ids for speculative decoding with overlap."""
assert result.next_token_ids.is_cpu assert result.next_token_ids.is_cpu
@@ -416,8 +439,9 @@ class SchedulerOutputProcessorMixin:
return predict_tokens return predict_tokens
@staticmethod
def process_batch_result_idle( def process_batch_result_idle(
self: Scheduler, self: "SchedulerBatchResultProcessor",
batch: ScheduleBatch, batch: ScheduleBatch,
result: GenerationBatchResult, result: GenerationBatchResult,
): ):
@@ -428,8 +452,9 @@ class SchedulerOutputProcessorMixin:
batch.reqs, batch.return_logprob, is_idle_batch=True batch.reqs, batch.return_logprob, is_idle_batch=True
) )
@staticmethod
def process_batch_result_decode( def process_batch_result_decode(
self: Scheduler, self: "SchedulerBatchResultProcessor",
batch: ScheduleBatch, batch: ScheduleBatch,
result: GenerationBatchResult, result: GenerationBatchResult,
): ):
@@ -450,7 +475,11 @@ class SchedulerOutputProcessorMixin:
if batch.spec_algorithm.is_none() or batch.is_spec_v2: if batch.spec_algorithm.is_none() or batch.is_spec_v2:
if 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): elif isinstance(next_token_ids, list):
pass # MLX path: already a list[int], skip torch round-trip pass # MLX path: already a list[int], skip torch round-trip
else: else:
@@ -479,7 +508,7 @@ class SchedulerOutputProcessorMixin:
self.metrics_reporter.update_spec_metrics( self.metrics_reporter.update_spec_metrics(
batch.batch_size(), result.num_correct_drafts 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( self.metrics_collector.increment_decode_cuda_graph_pass(
value=can_run_cuda_graph value=can_run_cuda_graph
) )
@@ -501,9 +530,13 @@ class SchedulerOutputProcessorMixin:
continue continue
if is_spec_v1: 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() 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: if req.return_hidden_states and logits_output.hidden_states is not None:
req.hidden_states.append( req.hidden_states.append(
logits_output.hidden_states[i].cpu().clone().tolist() logits_output.hidden_states[i].cpu().clone().tolist()
@@ -521,14 +554,20 @@ class SchedulerOutputProcessorMixin:
req.output_ids.extend(next_token_id) req.output_ids.extend(next_token_id)
new_accepted_len = len(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 # 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.time_stats.set_last_decode_finish_time()
req.check_finished(new_accepted_len) 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: if req.return_logprob:
# Spec v1 handles logprobs inside its own worker. # Spec v1 handles logprobs inside its own worker.
@@ -598,8 +637,12 @@ class SchedulerOutputProcessorMixin:
num_correct_drafts=result.num_correct_drafts, num_correct_drafts=result.num_correct_drafts,
) )
@staticmethod
def _handle_finished_req( def _handle_finished_req(
self: Scheduler, req: Req, i: int, logits_output: LogitsProcessorOutput self: "SchedulerBatchResultProcessor",
req: Req,
i: int,
logits_output: LogitsProcessorOutput,
): ):
if ( if (
self.server_args.disaggregation_decode_enable_offload_kvcache self.server_args.disaggregation_decode_enable_offload_kvcache
@@ -611,31 +654,37 @@ class SchedulerOutputProcessorMixin:
# delete feature to save memory # delete feature to save memory
if req.multimodal_inputs is not None and req.session is None: if req.multimodal_inputs is not None and req.session is None:
req.multimodal_inputs.release_features() req.multimodal_inputs.release_features()
self.maybe_collect_routed_experts(req) SchedulerOutputProcessorMixin._maybe_collect_routed_experts(self, req)
self.maybe_collect_indexer_topk(req) SchedulerOutputProcessorMixin._maybe_collect_indexer_topk(self, req)
if self.server_args.disaggregation_decode_enable_offload_kvcache: if self.server_args.disaggregation_decode_enable_offload_kvcache:
# Asynchronously offload KV cache; release_kv_cache will be called after Device->Host transfer completes # 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): if not self.decode_offload_manager.offload_kv_cache(req):
self.decode_offload_manager.finalize_release_on_finish(req) self.decode_offload_manager.finalize_release_on_finish(req)
else: else:
if self.enable_hisparse: if self.server_args.enable_hisparse:
self.hisparse_coordinator.request_finished(req) self.hisparse_coordinator.request_finished(req)
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.maybe_collect_customized_info(i, req, logits_output) SchedulerOutputProcessorMixin._maybe_collect_customized_info(
self, i, req, logits_output
)
@staticmethod
def _maybe_update_reasoning_tokens( 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 think_end_id = self.model_config.think_end_id
if req.require_reasoning and think_end_id is not None: if req.require_reasoning and think_end_id is not None:
req.update_reasoning_tokens(next_token_id, think_end_id) req.update_reasoning_tokens(next_token_id, think_end_id)
@staticmethod
def _mamba_prefix_cache_update( def _mamba_prefix_cache_update(
self: Scheduler, self: "SchedulerBatchResultProcessor",
req: Req, req: Req,
batch: ScheduleBatch, batch: ScheduleBatch,
result: GenerationBatchResult, result: GenerationBatchResult,