[RadixTree][2/N Refactor]: swa cache init tiny refactor (#17397)

This commit is contained in:
Yi Zhang
2026-01-21 15:48:30 +08:00
committed by GitHub
parent 0d49b13fdd
commit 236772c0e1
5 changed files with 12 additions and 12 deletions
+4 -6
View File
@@ -606,6 +606,7 @@ class Scheduler(
or self.tp_worker.model_runner.mamba2_config is not None or self.tp_worker.model_runner.mamba2_config is not None
) )
self.sliding_window_size = None
if self.is_hybrid_swa: if self.is_hybrid_swa:
self.sliding_window_size = self.tp_worker.sliding_window_size self.sliding_window_size = self.tp_worker.sliding_window_size
self.full_tokens_per_layer, self.swa_tokens_per_layer = ( self.full_tokens_per_layer, self.swa_tokens_per_layer = (
@@ -635,6 +636,7 @@ class Scheduler(
pp_rank=self.pp_rank, pp_rank=self.pp_rank,
pp_size=self.pp_size, pp_size=self.pp_size,
chunked_prefill_size=server_args.chunked_prefill_size, chunked_prefill_size=server_args.chunked_prefill_size,
sliding_window_size=self.sliding_window_size,
) )
if ( if (
@@ -648,9 +650,7 @@ class Scheduler(
else: else:
from sglang.srt.mem_cache.chunk_cache import SWAChunkCache from sglang.srt.mem_cache.chunk_cache import SWAChunkCache
self.tree_cache = SWAChunkCache( self.tree_cache = SWAChunkCache(params)
params, sliding_window_size=self.sliding_window_size
)
else: else:
if envs.SGLANG_EXPERIMENTAL_CPP_RADIX_TREE.get(): if envs.SGLANG_EXPERIMENTAL_CPP_RADIX_TREE.get():
@@ -669,9 +669,7 @@ class Scheduler(
elif self.is_hybrid_swa: elif self.is_hybrid_swa:
from sglang.srt.mem_cache.swa_radix_cache import SWARadixCache from sglang.srt.mem_cache.swa_radix_cache import SWARadixCache
self.tree_cache = SWARadixCache( self.tree_cache = SWARadixCache(params=params)
params=params, sliding_window_size=self.sliding_window_size
)
elif self.is_hybrid_ssm: elif self.is_hybrid_ssm:
from sglang.srt.mem_cache.mamba_radix_cache import MambaRadixCache from sglang.srt.mem_cache.mamba_radix_cache import MambaRadixCache
@@ -31,3 +31,5 @@ class CacheInitParams:
pp_size: int = 1 pp_size: int = 1
chunked_prefill_size: Optional[int] = None chunked_prefill_size: Optional[int] = None
sliding_window_size: Optional[int] = None
+2 -2
View File
@@ -90,11 +90,11 @@ class ChunkCache(BasePrefixCache):
class SWAChunkCache(ChunkCache): class SWAChunkCache(ChunkCache):
"""ChunkCache with support for sliding window attention.""" """ChunkCache with support for sliding window attention."""
def __init__(self, params: CacheInitParams, sliding_window_size: int): def __init__(self, params: CacheInitParams):
assert isinstance(params.token_to_kv_pool_allocator, SWATokenToKVPoolAllocator) assert isinstance(params.token_to_kv_pool_allocator, SWATokenToKVPoolAllocator)
super().__init__(params) super().__init__(params)
self.sliding_window_size = sliding_window_size self.sliding_window_size = params.sliding_window_size
self.chunked_prefill_size = params.chunked_prefill_size self.chunked_prefill_size = params.chunked_prefill_size
def supports_swa(self) -> bool: def supports_swa(self) -> bool:
@@ -333,7 +333,7 @@ class LRUList:
class SWARadixCache(BasePrefixCache): class SWARadixCache(BasePrefixCache):
def __init__(self, params: CacheInitParams, sliding_window_size: int): def __init__(self, params: CacheInitParams):
assert isinstance(params.token_to_kv_pool_allocator, SWATokenToKVPoolAllocator) assert isinstance(params.token_to_kv_pool_allocator, SWATokenToKVPoolAllocator)
self.req_to_token_pool = params.req_to_token_pool self.req_to_token_pool = params.req_to_token_pool
self.token_to_kv_pool_allocator = params.token_to_kv_pool_allocator self.token_to_kv_pool_allocator = params.token_to_kv_pool_allocator
@@ -361,7 +361,7 @@ class SWARadixCache(BasePrefixCache):
if params.enable_metrics: if params.enable_metrics:
self.init_metrics_collector() self.init_metrics_collector()
self.sliding_window_size = sliding_window_size self.sliding_window_size = params.sliding_window_size
self.reset() self.reset()
##### Public API ##### ##### Public API #####
@@ -127,8 +127,8 @@ class TestSWA(unittest.TestCase):
token_to_kv_pool_allocator=allocator, token_to_kv_pool_allocator=allocator,
disable=False, disable=False,
page_size=page_size, page_size=page_size,
),
sliding_window_size=sliding_window_size, sliding_window_size=sliding_window_size,
),
) )
# test # test
@@ -264,8 +264,8 @@ class TestSWA(unittest.TestCase):
page_size=page_size, page_size=page_size,
disable=False, disable=False,
is_eagle=True, is_eagle=True,
),
sliding_window_size=sliding_window_size, sliding_window_size=sliding_window_size,
),
) )
# test # test