From 5a1275a5195104b824e6bd0df71678cfe25a2aa1 Mon Sep 17 00:00:00 2001 From: Cheng Wan <54331508+ch-wan@users.noreply.github.com> Date: Wed, 2 Sep 2026 16:55:46 -0700 Subject: [PATCH] Converge the two SWA predicates, and stop conditioning the capture sink on the pool (#37550) Co-authored-by: Claude Opus 5 --- .../srt/layers/attention/aiter_backend.py | 11 +++++--- .../attention/flashattention_backend.py | 28 ++++++++----------- .../layers/attention/flashinfer_backend.py | 5 +--- .../srt/layers/attention/triton_backend.py | 16 +++++++---- .../layers/attention/trtllm_mha_backend.py | 20 +++++++++---- .../srt/mem_cache/kv_index_translator.py | 6 ++-- 6 files changed, 48 insertions(+), 38 deletions(-) diff --git a/python/sglang/srt/layers/attention/aiter_backend.py b/python/sglang/srt/layers/attention/aiter_backend.py index fb7981fbb..578fe3152 100755 --- a/python/sglang/srt/layers/attention/aiter_backend.py +++ b/python/sglang/srt/layers/attention/aiter_backend.py @@ -1192,11 +1192,14 @@ class AiterAttnBackend(AttentionBackend): if self.use_sliding_window_kv_pool and forward_batch.out_cache_loc is not None: n = forward_batch.out_cache_loc.shape[0] self.cuda_graph_swa_out_cache_loc[n:].zero_() - self.cuda_graph_swa_out_cache_loc[:n].copy_( - self.token_to_kv_pool.translate_loc_from_full_to_swa( - forward_batch.out_cache_loc + if in_capture: + self.cuda_graph_swa_out_cache_loc[:n].zero_() + else: + self.cuda_graph_swa_out_cache_loc[:n].copy_( + self.token_to_kv_pool.translate_loc_from_full_to_swa( + forward_batch.out_cache_loc + ) ) - ) self.forward_metadata.swa_out_cache_loc = self.cuda_graph_swa_out_cache_loc[ :n ] diff --git a/python/sglang/srt/layers/attention/flashattention_backend.py b/python/sglang/srt/layers/attention/flashattention_backend.py index 6dab7c141..cb962e6ed 100644 --- a/python/sglang/srt/layers/attention/flashattention_backend.py +++ b/python/sglang/srt/layers/attention/flashattention_backend.py @@ -494,6 +494,18 @@ class FlashAttentionBackend(AttentionBackend): spec_info = forward_batch.spec_info out_cache_loc = getattr(forward_batch, "out_cache_loc", None) + # Refill the SWA write-target buffer (bound as a metadata view in + # _bind_metadata_buffers) from the live out_cache_loc before replay. + if self.use_sliding_window_kv_pool and out_cache_loc is not None: + n = out_cache_loc.shape[0] + self.cuda_graph_swa_out_cache_loc[n:].zero_() + if in_capture: + self.cuda_graph_swa_out_cache_loc[:n].zero_() + else: + self.cuda_graph_swa_out_cache_loc[:n].copy_( + self.kv_index_translator.sliding_window_write_loc_for(out_cache_loc) + ) + if in_capture: num_tokens = forward_batch.positions.numel() seq_lens_cpu = seq_lens.cpu() @@ -537,7 +549,6 @@ class FlashAttentionBackend(AttentionBackend): spec_info=spec_info, seq_lens_cpu=seq_lens_cpu, out_cache_loc=out_cache_loc, - in_capture=True, ) if forward_mode.is_decode_or_idle() and spec_info is None: @@ -2754,7 +2765,6 @@ class FlashAttentionBackend(AttentionBackend): spec_info: Optional[SpecInput], seq_lens_cpu: Optional[torch.Tensor], out_cache_loc: Optional[torch.Tensor] = None, - in_capture: bool = False, ): """Shared capture+replay body for the cuda-graph init path. @@ -2772,20 +2782,6 @@ class FlashAttentionBackend(AttentionBackend): metadata = None metadata_expand = None - # Refill the SWA write-target buffer (bound as a metadata view in - # _bind_metadata_buffers) from the live out_cache_loc before replay. - if self.use_sliding_window_kv_pool and out_cache_loc is not None: - n = out_cache_loc.shape[0] - self.cuda_graph_swa_out_cache_loc[n:].zero_() - if in_capture and self.kv_index_translator.is_translating: - # A capture batch never went through `init_new`, so there is no - # rebound write loc; zeros are the page-0 sink. - self.cuda_graph_swa_out_cache_loc[:n].zero_() - else: - self.cuda_graph_swa_out_cache_loc[:n].copy_( - self.kv_index_translator.sliding_window_write_loc_for(out_cache_loc) - ) - if forward_mode.is_decode_or_idle(): if spec_info is not None: # Draft Decode diff --git a/python/sglang/srt/layers/attention/flashinfer_backend.py b/python/sglang/srt/layers/attention/flashinfer_backend.py index 8935fe3c2..504456fd9 100644 --- a/python/sglang/srt/layers/attention/flashinfer_backend.py +++ b/python/sglang/srt/layers/attention/flashinfer_backend.py @@ -842,10 +842,7 @@ class FlashInferAttnBackend(AttentionBackend): if self.use_sliding_window_kv_pool and forward_batch.out_cache_loc is not None: n = forward_batch.out_cache_loc.shape[0] self.cuda_graph_swa_out_cache_loc[n:].zero_() - if in_capture and self.kv_index_translator.is_translating: - # A runner-built capture batch never went through `init_new`, - # so there is no prepared write loc to resolve -- and zeros are the - # page-0 sink in every id space. Replay refills below. + if in_capture: self.cuda_graph_swa_out_cache_loc[:n].zero_() else: self.cuda_graph_swa_out_cache_loc[:n].copy_( diff --git a/python/sglang/srt/layers/attention/triton_backend.py b/python/sglang/srt/layers/attention/triton_backend.py index 45e396111..0d736d7e7 100644 --- a/python/sglang/srt/layers/attention/triton_backend.py +++ b/python/sglang/srt/layers/attention/triton_backend.py @@ -675,7 +675,9 @@ class TritonAttnBackend(AttentionBackend): out_cache_loc_full_physical = self._fill_cuda_graph_write_locs( forward_batch, bs ) - swa_out_cache_loc = self._fill_cuda_graph_swa_out_cache_loc(forward_batch) + swa_out_cache_loc = self._fill_cuda_graph_swa_out_cache_loc( + forward_batch, in_capture=True + ) self.forward_metadata = self._build_cuda_graph_forward_metadata( bs, forward_mode, @@ -696,7 +698,7 @@ class TritonAttnBackend(AttentionBackend): self._fill_cuda_graph_swa_out_cache_loc(forward_batch) def _fill_cuda_graph_swa_out_cache_loc( - self, forward_batch: ForwardBatch + self, forward_batch: ForwardBatch, in_capture: bool = False ) -> Optional[torch.Tensor]: """Refill the SWA write-target buffer from the batch's derived sliding-window write loc, returning the [:n] view (None for non-SWA / @@ -710,12 +712,14 @@ class TritonAttnBackend(AttentionBackend): or out_cache_loc.shape[0] > self.cuda_graph_swa_out_cache_loc.shape[0] ): return None - swa_write_loc = self.kv_index_translator.sliding_window_write_loc_for( - out_cache_loc - ) n = out_cache_loc.shape[0] self.cuda_graph_swa_out_cache_loc[n:].zero_() - self.cuda_graph_swa_out_cache_loc[:n].copy_(swa_write_loc) + if in_capture: + self.cuda_graph_swa_out_cache_loc[:n].zero_() + else: + self.cuda_graph_swa_out_cache_loc[:n].copy_( + self.kv_index_translator.sliding_window_write_loc_for(out_cache_loc) + ) return self.cuda_graph_swa_out_cache_loc[:n] def _fill_cuda_graph_write_locs( diff --git a/python/sglang/srt/layers/attention/trtllm_mha_backend.py b/python/sglang/srt/layers/attention/trtllm_mha_backend.py index c46397638..e5f8359ae 100644 --- a/python/sglang/srt/layers/attention/trtllm_mha_backend.py +++ b/python/sglang/srt/layers/attention/trtllm_mha_backend.py @@ -885,7 +885,11 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend): ragged_layout = resolve_ragged_verify_layout(forward_batch) if ragged_layout is not None: self._write_ragged_verify_graph_metadata( - self.forward_metadata, forward_batch, ragged_layout, bs + self.forward_metadata, + forward_batch, + ragged_layout, + bs, + in_capture=in_capture, ) elif forward_mode.is_draft_extend_v2(): self.forward_metadata = self.draft_extend_metadata[bs] @@ -918,7 +922,7 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend): ): n = forward_batch.out_cache_loc.shape[0] self.cuda_graph_swa_out_cache_loc[n:].zero_() - if in_capture and self.kv_index_translator.is_translating: + if in_capture: self.cuda_graph_swa_out_cache_loc[:n].zero_() else: self.cuda_graph_swa_out_cache_loc[:n].copy_( @@ -941,6 +945,7 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend): forward_batch: ForwardBatch, ragged_layout: RaggedVerifyLayout, bs: int, + in_capture: bool = False, ) -> None: """Eagerly rebuild the target-verify graph metadata for ragged verify. @@ -968,11 +973,14 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend): if self.use_sliding_window_kv_pool and forward_batch.out_cache_loc is not None: n = forward_batch.out_cache_loc.shape[0] self.cuda_graph_swa_out_cache_loc[n:].zero_() - self.cuda_graph_swa_out_cache_loc[:n].copy_( - self.token_to_kv_pool.translate_loc_from_full_to_swa( - forward_batch.out_cache_loc + if in_capture: + self.cuda_graph_swa_out_cache_loc[:n].zero_() + else: + self.cuda_graph_swa_out_cache_loc[:n].copy_( + self.token_to_kv_pool.translate_loc_from_full_to_swa( + forward_batch.out_cache_loc + ) ) - ) def init_forward_metadata_in_graph(self, forward_batch: ForwardBatch): self._apply_cuda_graph_metadata( diff --git a/python/sglang/srt/mem_cache/kv_index_translator.py b/python/sglang/srt/mem_cache/kv_index_translator.py index 89c636d64..7701c95dd 100644 --- a/python/sglang/srt/mem_cache/kv_index_translator.py +++ b/python/sglang/srt/mem_cache/kv_index_translator.py @@ -64,11 +64,11 @@ from sglang.kernels.ops.kvcache.kv_read_table import ( build_kv_read_table, build_kv_read_table_packed, ) +from sglang.srt.mem_cache.base_swa_memory_pool import BaseSWAKVPool from sglang.srt.mem_cache.multi_ended_allocator import ( UnifiedMambaTokenToKVPoolAllocator, UnifiedSWATokenToKVPoolAllocator, ) -from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool from sglang.srt.runtime_context import get_parallel @@ -154,9 +154,11 @@ class KVIndexTranslator: self.defer_read_translate = False self._swa_v2p_table = None self._swa_page_multiplier = 1 + # `translate_loc_from_full_to_swa` is abstract on `BaseSWAKVPool`, + # which is also what the backends' `_resolve_swa_kv_pool` keys on. self._swa_write_loc_from_full = ( token_to_kv_pool.translate_loc_from_full_to_swa - if isinstance(token_to_kv_pool, SWAKVPool) + if isinstance(token_to_kv_pool, BaseSWAKVPool) else None )