From 0c204fbd57a0deb12e0736b9ba783ac1d24ece85 Mon Sep 17 00:00:00 2001 From: huangtingwei <141888744+huangtingwei9988@users.noreply.github.com> Date: Wed, 8 Apr 2026 01:34:58 +0800 Subject: [PATCH] [HiSparse] Optimize the scheduling of decode backup. (#21932) Co-authored-by: hzh0425 Co-authored-by: Zhiqiang Xie --- .../srt/managers/hisparse_coordinator.py | 45 +++++++++++++++---- .../sglang/srt/model_executor/model_runner.py | 6 +++ 2 files changed, 42 insertions(+), 9 deletions(-) diff --git a/python/sglang/srt/managers/hisparse_coordinator.py b/python/sglang/srt/managers/hisparse_coordinator.py index 89740f736..933657197 100644 --- a/python/sglang/srt/managers/hisparse_coordinator.py +++ b/python/sglang/srt/managers/hisparse_coordinator.py @@ -78,8 +78,11 @@ class HiSparseCoordinator: ) self.write_staging_stream = device_module.Stream() + self.decode_backup_stream = device_module.Stream() self.ack_staging_queue: List[HiSparseAct] = [] self.decode_producer_stream = None + self._backup_done_event = device_module.Event() + self._has_pending_backup = False self.tp_group = tp_group self.tp_world_size = torch.distributed.get_world_size(group=self.tp_group) @@ -391,9 +394,6 @@ class HiSparseCoordinator: The only exception is the first decode step right after staging: all prefill tokens were already backed up during staging, so there is nothing new to save yet. """ - if self.decode_producer_stream is not None: - device_module.current_stream().wait_stream(self.decode_producer_stream) - # Build the list of batch positions that need a host backup. # Skip the first decode step after staging (prefill already backed up). backup_indices = [] @@ -431,12 +431,36 @@ class HiSparseCoordinator: host_locs = host_locs.to(device=self.device) self.req_to_host_pool[backup_req_indices, actual_token_pos] = host_locs - self.mem_pool_host.backup_from_device_all_layer( - self.mem_pool_device, - host_locs, - device_locs.contiguous(), - io_backend="kernel", - ) + if self._has_pending_backup: + self._backup_done_event.wait(device_module.current_stream()) + self._has_pending_backup = False + schedule_stream = device_module.current_stream() + with device_module.stream(self.decode_backup_stream): + self.decode_backup_stream.wait_stream(schedule_stream) + if self.decode_producer_stream is not None: + self.decode_backup_stream.wait_stream(self.decode_producer_stream) + self.mem_pool_host.backup_from_device_all_layer( + self.mem_pool_device, + host_locs, + device_locs, + io_backend="kernel", + ) + self._backup_done_event.record() + if host_locs.is_cuda: + host_locs.record_stream(self.decode_backup_stream) + if backup_req_indices.is_cuda: + backup_req_indices.record_stream(self.decode_backup_stream) + if actual_token_pos.is_cuda: + actual_token_pos.record_stream(self.decode_backup_stream) + if device_locs.is_cuda: + device_locs.record_stream(self.decode_backup_stream) + self._has_pending_backup = True + + def wait_for_pending_backup(self) -> None: + if not self._has_pending_backup: + return + self._backup_done_event.wait(device_module.current_stream()) + self._has_pending_backup = False def get_front_topk_tokens( self, @@ -569,6 +593,9 @@ class HiSparseCoordinator: # release resources only after the execution of a potential overlapped batch if self.decode_producer_stream is not None: device_module.current_stream().wait_stream(self.decode_producer_stream) + if self._has_pending_backup: + self._backup_done_event.wait(device_module.current_stream()) + self._has_pending_backup = False # release memory — only free actually-allocated buffer indices current_cap = int(self.req_device_buffer_size[req.req_pool_idx]) diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index 669cab133..26d8bd82a 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -2817,6 +2817,12 @@ class ModelRunner(ModelRunnerKVCacheMixin): and self.graph_runner.can_run(forward_batch) ) + if ( + self.hisparse_coordinator is not None + and forward_batch.forward_mode.is_decode() + ): + self.hisparse_coordinator.wait_for_pending_backup() + if can_run_graph: ret = self.graph_runner.replay( forward_batch,