perf(disaggregation): fast-path common decode result processing
This commit is contained in:
@@ -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
|
||||
|
||||
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user