[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

View File

@@ -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

View File

@@ -31,3 +31,5 @@ class CacheInitParams:
pp_size: int = 1
chunked_prefill_size: Optional[int] = None
sliding_window_size: Optional[int] = None

View File

@@ -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:

View File

@@ -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 #####