[RadixTree][2/N Refactor]: swa cache init tiny refactor (#17397)
This commit is contained in:
@@ -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
|
||||
|
||||
|
||||
@@ -31,3 +31,5 @@ class CacheInitParams:
|
||||
pp_size: int = 1
|
||||
|
||||
chunked_prefill_size: Optional[int] = None
|
||||
|
||||
sliding_window_size: Optional[int] = None
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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 #####
|
||||
|
||||
Reference in New Issue
Block a user