From 17119a697de72910d77d3bfffc34d097cf6cad09 Mon Sep 17 00:00:00 2001 From: sky Date: Wed, 4 Mar 2026 16:32:42 +0800 Subject: [PATCH] Optimization: Reduce the number of D2H operations (#19424) Signed-off-by: wangfakang --- python/sglang/srt/batch_overlap/two_batch_overlap.py | 6 ++++-- python/sglang/srt/managers/scheduler_dp_attn_mixin.py | 6 ++++-- 2 files changed, 8 insertions(+), 4 deletions(-) diff --git a/python/sglang/srt/batch_overlap/two_batch_overlap.py b/python/sglang/srt/batch_overlap/two_batch_overlap.py index c05e10b5b..c0d1a4923 100644 --- a/python/sglang/srt/batch_overlap/two_batch_overlap.py +++ b/python/sglang/srt/batch_overlap/two_batch_overlap.py @@ -416,8 +416,10 @@ class TboDPAttentionPreparer: return local_can_run_tbo, local_forward_mode def compute_output(self, partial_global_info): - local_can_run_tbo_aggregated = min(partial_global_info[:, 0].tolist()) - forward_modes = partial_global_info[:, 1].tolist() + # Perform only one Device-to-Host (D2H) memory copy + cpu_data = partial_global_info[:, :2].cpu() + local_can_run_tbo_aggregated = min(cpu_data[:, 0].tolist()) + forward_modes = cpu_data[:, 1].tolist() global_forward_mode, forward_mode_agree = self._compute_global_forward_mode( forward_modes diff --git a/python/sglang/srt/managers/scheduler_dp_attn_mixin.py b/python/sglang/srt/managers/scheduler_dp_attn_mixin.py index 58122c8e4..5331fc033 100644 --- a/python/sglang/srt/managers/scheduler_dp_attn_mixin.py +++ b/python/sglang/srt/managers/scheduler_dp_attn_mixin.py @@ -94,8 +94,10 @@ class MLPSyncBatchInfo: tp0_info = global_info_tensor[:, 0, :] self.tp0_info = tp0_info - self.global_num_tokens = tp0_info[:, 0].tolist() - self.global_num_tokens_for_logprob = tp0_info[:, 1].tolist() + # Perform only one Device-to-Host (D2H) memory copy + cpu_data = tp0_info[:, :2].cpu() + self.global_num_tokens = cpu_data[:, 0].tolist() + self.global_num_tokens_for_logprob = cpu_data[:, 1].tolist() self.can_cuda_graph = bool(tp0_info[:, 2].min().item()) self.is_extend_in_batch = bool(tp0_info[:, 3].max().item()) if _ENABLE_METRICS_DP_ATTENTION: