From 1d8e3c248b0dc7b312a5292fc572b0a116c498e0 Mon Sep 17 00:00:00 2001 From: Junjie Cao Date: Fri, 10 Jul 2026 15:27:02 +0800 Subject: [PATCH] Fix DSV4 HiSparse SWA tail allocation forwarding (#30408) Co-authored-by: Zhangheng --- python/sglang/srt/disaggregation/decode.py | 53 ++++-- .../srt/mem_cache/allocator/hisparse.py | 20 +++ .../unit/managers/test_hisparse_unit.py | 1 + .../unit/mem_cache/test_hisparse_allocator.py | 155 ++++++++++++++++++ 4 files changed, 212 insertions(+), 17 deletions(-) create mode 100644 test/registered/unit/mem_cache/test_hisparse_allocator.py diff --git a/python/sglang/srt/disaggregation/decode.py b/python/sglang/srt/disaggregation/decode.py index 1b51c66ed..acd094a18 100644 --- a/python/sglang/srt/disaggregation/decode.py +++ b/python/sglang/srt/disaggregation/decode.py @@ -68,10 +68,7 @@ from sglang.srt.managers.schedule_batch import ( from sglang.srt.managers.schedule_policy import match_prefix_for_req from sglang.srt.managers.utils import GenerationBatchResult from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator -from sglang.srt.mem_cache.base_prefix_cache import ( - BasePrefixCache, - EvictParams, -) +from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache, EvictParams from sglang.srt.mem_cache.common import ( kv_to_page_indices, page_align_floor, @@ -1285,11 +1282,15 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin): ) if self.scheduler.enable_hisparse: - # HiSparse pre-alloc only allocates logical indices (alloc_logical_only), - # so the logical pool is the binding constraint for admission control. - available_size = ( - self.token_to_kv_pool_allocator.logical_attn_allocator.available_size() - ) + logical_allocator = self.token_to_kv_pool_allocator.logical_attn_allocator + if self._uses_swa_tail_prealloc() and hasattr( + logical_allocator, "full_available_size" + ): + available_size = logical_allocator.full_available_size() + else: + # HiSparse pre-alloc only allocates logical indices, so the + # logical pool is the binding constraint for admission control. + available_size = logical_allocator.available_size() elif self._uses_swa_tail_prealloc(): available_size = self.token_to_kv_pool_allocator.full_available_size() if self.scheduler.server_args.disaggregation_decode_enable_radix_cache: @@ -1471,14 +1472,32 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin): # device indices) and allocate host indices for RDMA destination. coordinator = self.scheduler.hisparse_coordinator device = self.token_to_kv_pool_allocator.device - kv_loc = self.token_to_kv_pool_allocator.alloc_logical_only( - prefix_lens=torch.tensor([0], dtype=torch.int64, device=device), - prefix_lens_cpu=torch.tensor([0], dtype=torch.int64), - seq_lens=torch.tensor([fill_len], dtype=torch.int64, device=device), - seq_lens_cpu=torch.tensor([fill_len], dtype=torch.int64), - last_loc=torch.tensor([-1], dtype=torch.int64, device=device), - extend_num_tokens=fill_len, - ) + prefix_lens = torch.tensor([0], dtype=torch.int64, device=device) + prefix_lens_cpu = torch.tensor([0], dtype=torch.int64) + seq_lens = torch.tensor([fill_len], dtype=torch.int64, device=device) + seq_lens_cpu = torch.tensor([fill_len], dtype=torch.int64) + last_loc = torch.tensor([-1], dtype=torch.int64, device=device) + if self._uses_swa_tail_prealloc(): + swa_tail_len = self._swa_tail_len(fill_len) + kv_loc = self.token_to_kv_pool_allocator.alloc_extend_swa_tail( + prefix_lens=prefix_lens, + prefix_lens_cpu=prefix_lens_cpu, + seq_lens=seq_lens, + seq_lens_cpu=seq_lens_cpu, + last_loc=last_loc, + extend_num_tokens=fill_len, + swa_tail_len=swa_tail_len, + ) + req.swa_evicted_seqlen = fill_len - swa_tail_len + else: + kv_loc = self.token_to_kv_pool_allocator.alloc_logical_only( + prefix_lens=prefix_lens, + prefix_lens_cpu=prefix_lens_cpu, + seq_lens=seq_lens, + seq_lens_cpu=seq_lens_cpu, + last_loc=last_loc, + extend_num_tokens=fill_len, + ) # Allocate host indices for the RDMA transfer target. host_indices = coordinator.mem_pool_host.alloc_paged_token_slots( diff --git a/python/sglang/srt/mem_cache/allocator/hisparse.py b/python/sglang/srt/mem_cache/allocator/hisparse.py index cf30a7de2..79245ac29 100644 --- a/python/sglang/srt/mem_cache/allocator/hisparse.py +++ b/python/sglang/srt/mem_cache/allocator/hisparse.py @@ -402,6 +402,26 @@ class DeepSeekV4HiSparseTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator): extend_num_tokens, ) + def alloc_extend_swa_tail( + self, + prefix_lens: torch.Tensor, + prefix_lens_cpu: torch.Tensor, + seq_lens: torch.Tensor, + seq_lens_cpu: torch.Tensor, + last_loc: torch.Tensor, + extend_num_tokens: int, + swa_tail_len: int, + ): + return self.logical_attn_allocator.alloc_extend_swa_tail( + prefix_lens=prefix_lens, + prefix_lens_cpu=prefix_lens_cpu, + seq_lens=seq_lens, + seq_lens_cpu=seq_lens_cpu, + last_loc=last_loc, + extend_num_tokens=extend_num_tokens, + swa_tail_len=swa_tail_len, + ) + def alloc_device_buffer(self, allocated_indices, need_size: int): assert need_size % self.hisparse_page_size == 0 hisparse_indices = self.full_to_hisparse_device_index_mapping[allocated_indices] diff --git a/test/registered/unit/managers/test_hisparse_unit.py b/test/registered/unit/managers/test_hisparse_unit.py index 33d91cf9a..779f8b365 100644 --- a/test/registered/unit/managers/test_hisparse_unit.py +++ b/test/registered/unit/managers/test_hisparse_unit.py @@ -749,6 +749,7 @@ class TestHiSparseUnit(unittest.TestCase): queue = DecodePreallocQueue.__new__(DecodePreallocQueue) queue.req_to_token_pool = self.req_to_token_pool queue.token_to_kv_pool_allocator = self.allocator + queue.token_to_kv_pool = self.allocator.get_kvcache() queue.tree_cache = SimpleNamespace( evictable_size=lambda: 0, protected_size=lambda: 0, diff --git a/test/registered/unit/mem_cache/test_hisparse_allocator.py b/test/registered/unit/mem_cache/test_hisparse_allocator.py new file mode 100644 index 000000000..120ec6a28 --- /dev/null +++ b/test/registered/unit/mem_cache/test_hisparse_allocator.py @@ -0,0 +1,155 @@ +import unittest +from types import SimpleNamespace +from unittest.mock import MagicMock + +import torch + +from sglang.srt.mem_cache.allocator.hisparse import ( + DeepSeekV4HiSparseTokenToKVPoolAllocator, +) +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=1, suite="base-a-test-cpu") + + +class TestDeepSeekV4HiSparseAllocator(CustomTestCase): + def test_forwards_swa_tail_allocation_to_logical_allocator(self): + allocator = object.__new__(DeepSeekV4HiSparseTokenToKVPoolAllocator) + logical_allocator = MagicMock(spec=["alloc_extend_swa_tail"]) + allocator.logical_attn_allocator = logical_allocator + + expected = torch.tensor([8, 9, 10], dtype=torch.int64) + logical_allocator.alloc_extend_swa_tail.return_value = expected + + prefix_lens = torch.tensor([0], dtype=torch.int64) + prefix_lens_cpu = torch.tensor([0], dtype=torch.int64) + seq_lens = torch.tensor([512], dtype=torch.int64) + seq_lens_cpu = torch.tensor([512], dtype=torch.int64) + last_loc = torch.tensor([-1], dtype=torch.int64) + + result = allocator.alloc_extend_swa_tail( + prefix_lens=prefix_lens, + prefix_lens_cpu=prefix_lens_cpu, + seq_lens=seq_lens, + seq_lens_cpu=seq_lens_cpu, + last_loc=last_loc, + extend_num_tokens=512, + swa_tail_len=128, + ) + + self.assertIs(result, expected) + logical_allocator.alloc_extend_swa_tail.assert_called_once() + _, kwargs = logical_allocator.alloc_extend_swa_tail.call_args + self.assertIs(kwargs["prefix_lens"], prefix_lens) + self.assertIs(kwargs["prefix_lens_cpu"], prefix_lens_cpu) + self.assertIs(kwargs["seq_lens"], seq_lens) + self.assertIs(kwargs["seq_lens_cpu"], seq_lens_cpu) + self.assertIs(kwargs["last_loc"], last_loc) + self.assertEqual(kwargs["extend_num_tokens"], 512) + self.assertEqual(kwargs["swa_tail_len"], 128) + + def test_hisparse_budget_uses_full_logical_capacity_for_swa_tail(self): + from sglang.srt.disaggregation.decode import DecodePreallocQueue + + queue = DecodePreallocQueue.__new__(DecodePreallocQueue) + logical_allocator = SimpleNamespace( + available_size=MagicMock(return_value=32), + full_available_size=MagicMock(return_value=512), + ) + queue.token_to_kv_pool_allocator = SimpleNamespace( + logical_attn_allocator=logical_allocator + ) + queue.scheduler = SimpleNamespace(enable_hisparse=True, last_batch=None) + queue.retracted_queue = [] + queue.num_reserved_decode_tokens = 0 + queue._uses_swa_tail_prealloc = MagicMock(return_value=True) + queue._need_space_for_single_req = MagicMock(return_value=0) + queue._active_reserved_tokens = MagicMock(return_value=0) + + budget = queue._allocatable_token_budgets() + + self.assertEqual(budget, 512) + logical_allocator.full_available_size.assert_called_once_with() + logical_allocator.available_size.assert_not_called() + + def test_hisparse_prealloc_uses_swa_tail_for_direct_host_path(self): + from sglang.srt.disaggregation.decode import DecodePreallocQueue + + fill_len = 512 + swa_tail_len = 128 + kv_loc = torch.arange(fill_len, dtype=torch.int64) + host_indices = torch.arange(1000, 1000 + fill_len, dtype=torch.int64) + + req = SimpleNamespace( + rid="req-0", + origin_input_ids=list(range(fill_len)), + output_ids=[], + ) + + def set_extend_range(start, end): + req.extend_range = SimpleNamespace(start=start, end=end, length=end - start) + + req.set_extend_range = set_extend_range + + class ReqToTokenPool: + def __init__(self): + self.writes = [] + + def alloc(self, reqs): + for item in reqs: + item.req_pool_idx = 0 + return torch.tensor([0], dtype=torch.int64) + + def write(self, indices, values): + self.writes.append((indices, values)) + + req_to_token_pool = ReqToTokenPool() + allocator = SimpleNamespace( + device=torch.device("cpu"), + page_size=64, + available_size=MagicMock(return_value=fill_len), + alloc_extend_swa_tail=MagicMock(return_value=kv_loc), + alloc_logical_only=MagicMock(return_value=kv_loc), + ) + coordinator = SimpleNamespace( + mem_pool_host=SimpleNamespace( + alloc_paged_token_slots=MagicMock(return_value=host_indices) + ), + req_to_host_pool=object(), + req_to_host_pool_allocated_len=object(), + host_token_len=MagicMock(side_effect=lambda length: length), + ) + queue = DecodePreallocQueue.__new__(DecodePreallocQueue) + queue.req_to_token_pool = req_to_token_pool + queue.token_to_kv_pool_allocator = allocator + queue.tree_cache = SimpleNamespace( + evictable_size=MagicMock(return_value=0), + protected_size=MagicMock(return_value=0), + ) + queue.scheduler = SimpleNamespace( + enable_hisparse=True, + hisparse_coordinator=coordinator, + server_args=SimpleNamespace(disaggregation_decode_enable_radix_cache=False), + ) + queue._uses_swa_tail_prealloc = MagicMock(return_value=True) + queue._swa_tail_len = MagicMock(return_value=swa_tail_len) + + result = queue._pre_alloc(req) + + self.assertIs(result, host_indices) + allocator.alloc_extend_swa_tail.assert_called_once() + allocator.alloc_logical_only.assert_not_called() + _, kwargs = allocator.alloc_extend_swa_tail.call_args + self.assertEqual(kwargs["extend_num_tokens"], fill_len) + self.assertEqual(kwargs["swa_tail_len"], swa_tail_len) + self.assertEqual(req.swa_evicted_seqlen, fill_len - swa_tail_len) + self.assertEqual(req.kv_allocated_len, fill_len) + self.assertEqual(req.kv_committed_len, fill_len) + self.assertEqual(req.extend_range.length, fill_len) + self.assertEqual(len(req_to_token_pool.writes), 1) + coordinator.mem_pool_host.alloc_paged_token_slots.assert_called_once() + + +if __name__ == "__main__": + unittest.main()