perf(disaggregation): fast-path common decode result processing

This commit is contained in:
wxiwnd
2026-04-09 00:25:35 +08:00
parent 14c40f4e7d
commit 91facd5bb2
2 changed files with 226 additions and 0 deletions

View File

@@ -369,11 +369,27 @@ class SchedulerOutputProcessorMixin:
result.can_run_cuda_graph,
)
decode_fast_path = False
if batch.spec_algorithm.is_none() or batch.is_spec_v2:
if batch.is_spec_v2:
next_token_ids = self._resolve_spec_overlap_token_ids(result, batch)
else:
next_token_ids = next_token_ids.tolist()
if not batch.return_logprob:
decode_fast_path = (
logits_output.hidden_states is None
and logits_output.customized_info is None
and not self.server_args.disaggregation_decode_enable_offload_kvcache
and not self.enable_hisparse
and all(
not req.return_hidden_states
and req.grammar is None
and req.token_ids_logprob is None
and req.top_logprobs_num <= 0
and req.multimodal_inputs is None
for req in batch.reqs
)
)
if batch.return_logprob:
next_token_logprobs = logits_output.next_token_logprobs.tolist()
@@ -401,6 +417,38 @@ class SchedulerOutputProcessorMixin:
self.token_to_kv_pool_allocator.free_group_begin()
if decode_fast_path:
for i, (req, next_token_id) in enumerate(zip(batch.reqs, next_token_ids)):
req: Req
if self.enable_overlap and (req.finished() or req.is_retracted):
continue
req.output_ids.append(next_token_id)
# Update Mamba last track seqlen
self._mamba_prefix_cache_update(req, batch, result, i)
req.time_stats.set_last_decode_finish_time()
req.check_finished(1)
if req.finished():
self.maybe_collect_routed_experts(req)
release_kv_cache(req, self.tree_cache)
req.time_stats.set_completion_time()
self.stream_output(batch.reqs, batch.return_logprob)
self.token_to_kv_pool_allocator.free_group_end()
self.forward_ct_decode = (self.forward_ct_decode + 1) % (1 << 30)
self.report_decode_stats(
can_run_cuda_graph,
running_batch=batch,
num_accepted_tokens=result.num_accepted_tokens,
)
return
# NOTE: in any case, we should check finish here
# if finished, also clean up committed kv cache and over-allocated kv cache here

View File

@@ -0,0 +1,178 @@
import unittest
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import torch
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
from sglang.srt.managers.scheduler_output_processor_mixin import (
SchedulerOutputProcessorMixin,
)
from sglang.srt.managers.utils import GenerationBatchResult
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=8, suite="stage-a-test-cpu")
class FakeFinishReason:
def __init__(self):
self.value = None
class FakeReq:
def __init__(self, rid: str, *, finished_after_append: bool = False):
self.rid = rid
self.output_ids = []
self.return_logprob = False
self.return_hidden_states = False
self.grammar = None
self.token_ids_logprob = None
self.top_logprobs_num = 0
self.is_retracted = False
self.multimodal_inputs = None
self.finished_reason = None
self.time_stats = SimpleNamespace(
set_last_decode_finish_time=MagicMock(),
set_completion_time=MagicMock(),
)
self._finished = False
self._finished_after_append = finished_after_append
self.check_finished_calls = []
def finished(self):
return self._finished
def check_finished(self, accepted_len=1):
self.check_finished_calls.append(accepted_len)
if self._finished_after_append:
self._finished = True
class TestSchedulerOutputProcessorMixin(CustomTestCase):
def _make_scheduler(self):
scheduler = SimpleNamespace()
scheduler.num_generated_tokens = 0
scheduler.enable_metrics = True
scheduler.metrics_collector = SimpleNamespace(
increment_decode_cuda_graph_pass=MagicMock()
)
scheduler.token_to_kv_pool_allocator = SimpleNamespace(
free_group_begin=MagicMock(),
free_group_end=MagicMock(),
)
scheduler.enable_overlap = False
scheduler.server_args = SimpleNamespace(
disaggregation_decode_enable_offload_kvcache=False
)
scheduler.enable_hisparse = False
scheduler.tree_cache = object()
scheduler.forward_ct_decode = 0
scheduler.stream_output = MagicMock()
scheduler.report_decode_stats = MagicMock()
scheduler.maybe_collect_customized_info = MagicMock()
scheduler.maybe_collect_routed_experts = MagicMock()
scheduler._mamba_prefix_cache_update = MagicMock()
scheduler.update_spec_metrics = MagicMock()
scheduler.decode_offload_manager = MagicMock()
scheduler.hisparse_coordinator = MagicMock()
scheduler.abort_request = MagicMock()
return scheduler
def _make_batch(self, reqs, *, return_logprob=False):
return SimpleNamespace(
reqs=reqs,
return_logprob=return_logprob,
spec_algorithm=SimpleNamespace(is_none=lambda: True),
is_spec_v2=False,
batch_size=lambda: len(reqs),
)
def _make_result(self, *, next_token_ids, hidden_states=None, customized_info=None):
logits_output = LogitsProcessorOutput(
next_token_logits=None,
hidden_states=hidden_states,
customized_info=customized_info,
)
return GenerationBatchResult(
logits_output=logits_output,
next_token_ids=torch.tensor(next_token_ids),
can_run_cuda_graph=True,
)
def test_process_batch_result_decode_fast_path_common_case(self):
scheduler = self._make_scheduler()
finished_req = FakeReq("finished", finished_after_append=True)
ongoing_req = FakeReq("ongoing")
batch = self._make_batch([finished_req, ongoing_req])
result = self._make_result(next_token_ids=[11, 22])
released = []
with patch(
"sglang.srt.managers.scheduler_output_processor_mixin.release_kv_cache",
lambda req, tree_cache: released.append(req.rid),
):
SchedulerOutputProcessorMixin.process_batch_result_decode(
scheduler, batch, result
)
self.assertEqual(finished_req.output_ids, [11])
self.assertEqual(ongoing_req.output_ids, [22])
self.assertEqual(finished_req.check_finished_calls, [1])
self.assertEqual(ongoing_req.check_finished_calls, [1])
self.assertEqual(released, ["finished"])
finished_req.time_stats.set_last_decode_finish_time.assert_called_once()
ongoing_req.time_stats.set_last_decode_finish_time.assert_called_once()
finished_req.time_stats.set_completion_time.assert_called_once()
ongoing_req.time_stats.set_completion_time.assert_not_called()
scheduler._mamba_prefix_cache_update.assert_any_call(
finished_req, batch, result, 0
)
scheduler._mamba_prefix_cache_update.assert_any_call(
ongoing_req, batch, result, 1
)
self.assertEqual(scheduler.maybe_collect_customized_info.call_count, 0)
scheduler.maybe_collect_routed_experts.assert_called_once_with(finished_req)
scheduler.stream_output.assert_called_once_with(
batch.reqs, batch.return_logprob
)
scheduler.token_to_kv_pool_allocator.free_group_begin.assert_called_once()
scheduler.token_to_kv_pool_allocator.free_group_end.assert_called_once()
scheduler.metrics_collector.increment_decode_cuda_graph_pass.assert_called_once_with(
value=True
)
scheduler.report_decode_stats.assert_called_once_with(
True, running_batch=batch, num_accepted_tokens=result.num_accepted_tokens
)
self.assertEqual(scheduler.forward_ct_decode, 1)
self.assertEqual(scheduler.num_generated_tokens, 2)
def test_process_batch_result_decode_customized_info_stays_slow_path(self):
scheduler = self._make_scheduler()
req = FakeReq("req")
batch = self._make_batch([req])
result = self._make_result(
next_token_ids=[7], customized_info={"foo": [torch.tensor([1])]}
)
with patch(
"sglang.srt.managers.scheduler_output_processor_mixin.release_kv_cache",
lambda req, tree_cache: None,
):
SchedulerOutputProcessorMixin.process_batch_result_decode(
scheduler, batch, result
)
scheduler.maybe_collect_customized_info.assert_called_once_with(
0, req, result.logits_output
)
self.assertEqual(req.output_ids, [7])
scheduler.stream_output.assert_called_once_with(
batch.reqs, batch.return_logprob
)
if __name__ == "__main__":
unittest.main()