From 33719cfb317b1dd7fb86289e77d04f2ea63f0926 Mon Sep 17 00:00:00 2001 From: cctry Date: Mon, 15 Jun 2026 11:01:55 -0700 Subject: [PATCH] [PD] Optimize SWA allocation (#28085) Co-authored-by: cctry Co-authored-by: Lianmin Zheng --- python/sglang/srt/disaggregation/decode.py | 1 + python/sglang/srt/mem_cache/allocator/swa.py | 38 +++++++++---------- .../test_swa_alloc_extend_page_estimation.py | 17 ++++++++- 3 files changed, 34 insertions(+), 22 deletions(-) diff --git a/python/sglang/srt/disaggregation/decode.py b/python/sglang/srt/disaggregation/decode.py index 6c4045bd2..6fb6aa2f6 100644 --- a/python/sglang/srt/disaggregation/decode.py +++ b/python/sglang/srt/disaggregation/decode.py @@ -1379,6 +1379,7 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin): extend_num_tokens=fill_len, swa_tail_len=self._swa_tail_len(fill_len), ) + req.swa_evicted_seqlen = fill_len - self._swa_tail_len(fill_len) else: kv_loc = self.token_to_kv_pool_allocator.alloc_extend( prefix_lens=torch.tensor( diff --git a/python/sglang/srt/mem_cache/allocator/swa.py b/python/sglang/srt/mem_cache/allocator/swa.py index 66e49aadb..31745d645 100644 --- a/python/sglang/srt/mem_cache/allocator/swa.py +++ b/python/sglang/srt/mem_cache/allocator/swa.py @@ -151,14 +151,17 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator): assert alloc_full_indices is not None assert alloc_swa_indices is not None - if _is_npu: - self.full_to_swa_index_mapping[alloc_full_indices.to(torch.int64)] = ( - alloc_swa_indices.to(torch.int64) - ) - else: - self.full_to_swa_index_mapping[alloc_full_indices] = alloc_swa_indices + self.set_full_to_swa_mapping(alloc_full_indices, alloc_swa_indices) return alloc_full_indices + def new_pages_available(self, num_full_pages: int, num_swa_pages: int) -> bool: + return ( + num_full_pages + <= self.full_attn_allocator.available_size() // self.page_size + and num_swa_pages + <= self.swa_attn_allocator.available_size() // self.page_size + ) + def alloc_extend( self, prefix_lens: torch.Tensor, @@ -173,9 +176,7 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator): num_new_pages = get_num_new_pages( seq_lens=seq_lens_cpu, page_size=self.page_size, prefix_lens=prefix_lens_cpu ) - if num_new_pages > self.full_attn_allocator.available_size() // self.page_size: - return None - if num_new_pages > self.swa_attn_allocator.available_size() // self.page_size: + if not self.new_pages_available(num_new_pages, num_new_pages): return None swa_last_loc = self.translate_loc_from_full_to_swa(last_loc) @@ -201,12 +202,7 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator): assert alloc_full_indices is not None assert alloc_swa_indices is not None - if _is_npu: - self.full_to_swa_index_mapping[alloc_full_indices.to(torch.int64)] = ( - alloc_swa_indices.to(torch.int64) - ) - else: - self.full_to_swa_index_mapping[alloc_full_indices] = alloc_swa_indices + self.set_full_to_swa_mapping(alloc_full_indices, alloc_swa_indices) return alloc_full_indices @@ -235,9 +231,7 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator): seq_lens=seq_lens_cpu, page_size=self.page_size, prefix_lens=prefix_lens_cpu ) num_swa_pages = (swa_tail_len + self.page_size - 1) // self.page_size - if num_full_pages > self.full_attn_allocator.available_size() // self.page_size: - return None - if num_swa_pages > self.swa_attn_allocator.available_size() // self.page_size: + if not self.new_pages_available(num_full_pages, num_swa_pages): return None alloc_full_indices = self.full_attn_allocator.alloc_extend( @@ -247,6 +241,7 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator): seq_lens_cpu, last_loc, extend_num_tokens, + num_new_pages=num_full_pages, ) assert alloc_full_indices is not None @@ -267,12 +262,13 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator): swa_seq_lens_cpu, swa_last_loc, swa_tail_len, + num_new_pages=num_swa_pages, ) assert alloc_swa_indices is not None - self.full_to_swa_index_mapping[ - alloc_full_indices[-swa_tail_len:].to(torch.int64) - ] = alloc_swa_indices.to(torch.int64) + self.set_full_to_swa_mapping( + alloc_full_indices[-swa_tail_len:], alloc_swa_indices + ) if swa_tail_len < extend_num_tokens: self.full_to_swa_index_mapping[ alloc_full_indices[:-swa_tail_len].to(torch.int64) diff --git a/test/registered/unit/mem_cache/test_swa_alloc_extend_page_estimation.py b/test/registered/unit/mem_cache/test_swa_alloc_extend_page_estimation.py index 3700ca73b..7cef341de 100644 --- a/test/registered/unit/mem_cache/test_swa_alloc_extend_page_estimation.py +++ b/test/registered/unit/mem_cache/test_swa_alloc_extend_page_estimation.py @@ -21,6 +21,19 @@ register_cpu_ci(est_time=2, suite="base-a-test-cpu") def _make_self(*, page_size: int, full_available: int, swa_available: int): full_indices = torch.tensor([10, 11], dtype=torch.int64) swa_indices = torch.tensor([20, 21], dtype=torch.int64) + full_to_swa_index_mapping = torch.zeros(64, dtype=torch.int64) + + def new_pages_available(num_full_pages: int, num_swa_pages: int) -> bool: + return ( + num_full_pages <= full_available // page_size + and num_swa_pages <= swa_available // page_size + ) + + def set_full_to_swa_mapping( + full_indices: torch.Tensor, swa_indices: torch.Tensor + ) -> None: + full_to_swa_index_mapping[full_indices] = swa_indices + return SimpleNamespace( page_size=page_size, full_attn_allocator=SimpleNamespace( @@ -32,7 +45,9 @@ def _make_self(*, page_size: int, full_available: int, swa_available: int): alloc_extend=MagicMock(return_value=swa_indices), ), translate_loc_from_full_to_swa=lambda last_loc: last_loc, - full_to_swa_index_mapping=torch.zeros(64, dtype=torch.int64), + new_pages_available=new_pages_available, + set_full_to_swa_mapping=set_full_to_swa_mapping, + full_to_swa_index_mapping=full_to_swa_index_mapping, )