diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index fd1390a53..6c324b297 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -606,6 +606,7 @@ class Scheduler( or self.tp_worker.model_runner.mamba2_config is not None ) + self.sliding_window_size = None if self.is_hybrid_swa: self.sliding_window_size = self.tp_worker.sliding_window_size self.full_tokens_per_layer, self.swa_tokens_per_layer = ( @@ -635,6 +636,7 @@ class Scheduler( pp_rank=self.pp_rank, pp_size=self.pp_size, chunked_prefill_size=server_args.chunked_prefill_size, + sliding_window_size=self.sliding_window_size, ) if ( @@ -648,9 +650,7 @@ class Scheduler( else: from sglang.srt.mem_cache.chunk_cache import SWAChunkCache - self.tree_cache = SWAChunkCache( - params, sliding_window_size=self.sliding_window_size - ) + self.tree_cache = SWAChunkCache(params) else: if envs.SGLANG_EXPERIMENTAL_CPP_RADIX_TREE.get(): @@ -669,9 +669,7 @@ class Scheduler( elif self.is_hybrid_swa: from sglang.srt.mem_cache.swa_radix_cache import SWARadixCache - self.tree_cache = SWARadixCache( - params=params, sliding_window_size=self.sliding_window_size - ) + self.tree_cache = SWARadixCache(params=params) elif self.is_hybrid_ssm: from sglang.srt.mem_cache.mamba_radix_cache import MambaRadixCache diff --git a/python/sglang/srt/mem_cache/cache_init_params.py b/python/sglang/srt/mem_cache/cache_init_params.py index 7f6111f80..7b1b4a7d6 100644 --- a/python/sglang/srt/mem_cache/cache_init_params.py +++ b/python/sglang/srt/mem_cache/cache_init_params.py @@ -31,3 +31,5 @@ class CacheInitParams: pp_size: int = 1 chunked_prefill_size: Optional[int] = None + + sliding_window_size: Optional[int] = None diff --git a/python/sglang/srt/mem_cache/chunk_cache.py b/python/sglang/srt/mem_cache/chunk_cache.py index 75f545166..87f88c7c6 100644 --- a/python/sglang/srt/mem_cache/chunk_cache.py +++ b/python/sglang/srt/mem_cache/chunk_cache.py @@ -90,11 +90,11 @@ class ChunkCache(BasePrefixCache): class SWAChunkCache(ChunkCache): """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) 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 def supports_swa(self) -> bool: diff --git a/python/sglang/srt/mem_cache/swa_radix_cache.py b/python/sglang/srt/mem_cache/swa_radix_cache.py index fceb83a3a..8148e8c40 100644 --- a/python/sglang/srt/mem_cache/swa_radix_cache.py +++ b/python/sglang/srt/mem_cache/swa_radix_cache.py @@ -333,7 +333,7 @@ class LRUList: 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) self.req_to_token_pool = params.req_to_token_pool self.token_to_kv_pool_allocator = params.token_to_kv_pool_allocator @@ -361,7 +361,7 @@ class SWARadixCache(BasePrefixCache): if params.enable_metrics: self.init_metrics_collector() - self.sliding_window_size = sliding_window_size + self.sliding_window_size = params.sliding_window_size self.reset() ##### Public API ##### diff --git a/test/registered/radix_cache/test_swa_unittest.py b/test/registered/radix_cache/test_swa_unittest.py index 63548401b..deb0bd2c0 100644 --- a/test/registered/radix_cache/test_swa_unittest.py +++ b/test/registered/radix_cache/test_swa_unittest.py @@ -127,8 +127,8 @@ class TestSWA(unittest.TestCase): token_to_kv_pool_allocator=allocator, disable=False, page_size=page_size, + sliding_window_size=sliding_window_size, ), - sliding_window_size=sliding_window_size, ) # test @@ -264,8 +264,8 @@ class TestSWA(unittest.TestCase): page_size=page_size, disable=False, is_eagle=True, + sliding_window_size=sliding_window_size, ), - sliding_window_size=sliding_window_size, ) # test