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