[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
|
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
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
Reference in New Issue
Block a user