[HiSparse] Optimize the scheduling of decode backup. (#21932)

Co-authored-by: hzh0425 <hzh0425@apache.org>
Co-authored-by: Zhiqiang Xie <xiezhq@stanford.edu>
This commit is contained in:
huangtingwei
2026-04-07 10:34:58 -07:00
committed by GitHub
co-authored by hzh0425 Zhiqiang Xie
parent 6131fb5882
commit 0c204fbd57
2 changed files with 42 additions and 9 deletions
@@ -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])
@@ -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,