From 9a4d6402440a0f0e69cb6664bb015357b4e598fc Mon Sep 17 00:00:00 2001 From: Liangsheng Yin Date: Thu, 16 Jul 2026 15:00:41 -0700 Subject: [PATCH] [Perf] Cache uniform ragged-verify layout for DSpark verify-all compact (#31434) --- .../dspark_components/dspark_planner.py | 16 ++++++++++++++++ 1 file changed, 16 insertions(+) diff --git a/python/sglang/srt/speculative/dspark_components/dspark_planner.py b/python/sglang/srt/speculative/dspark_components/dspark_planner.py index d71f8b224..e76206d5c 100644 --- a/python/sglang/srt/speculative/dspark_components/dspark_planner.py +++ b/python/sglang/srt/speculative/dspark_components/dspark_planner.py @@ -124,6 +124,7 @@ class DSparkVerifyPlanner: self._dynamic_graph_tier = False self._dp_tier_gather_enabled = False self._is_verify_all = True + self._uniform_layout_cache: dict = {} if self._ragged_verify_mode is not RaggedVerifyMode.STATIC: if self._confidence_head is None: raise ValueError( @@ -417,6 +418,21 @@ class DSparkVerifyPlanner: ) -> Optional[RaggedVerifyLayout]: if self._ragged_verify_mode is RaggedVerifyMode.STATIC: return None + if self._is_verify_all and self._ragged_verify_mode is RaggedVerifyMode.COMPACT: + # Verify-all: the uniform layout (or None, past the captured grid) + # is constant per (bs, tier); serve it from cache instead of paying + # the per-step schedule and its host<->device round-trips. + key = (int(req_pool_indices.shape[0]), global_num_reqs) + if key not in self._uniform_layout_cache: + self._uniform_layout_cache[key] = uniform_ragged_layout( + bs=key[0], + device=device, + verify_num_draft_tokens=self.verify_num_draft_tokens, + ragged_verify_mode=self._ragged_verify_mode, + model_runner=self.model_runner, + tier_num_reqs=global_num_reqs, + ) + return self._uniform_layout_cache[key] verify_lens = self._schedule_verify_lens( req_pool_indices=req_pool_indices, prefix_lens=prefix_lens,