Using UnifiedRadixTree by default for SWA, Mamba, and DSA models (#30468)
Co-authored-by: ispobock <ispobaoke@gmail.com>
This commit is contained in:
@@ -26,6 +26,7 @@ def _make_ctx(
|
||||
enable_lmcache=False,
|
||||
is_hybrid_swa=False,
|
||||
is_hybrid_ssm=False,
|
||||
is_dsa=False,
|
||||
enable_hierarchical_cache=False,
|
||||
disable_radix_cache=False,
|
||||
effective_chunked_prefill_size=None,
|
||||
@@ -41,6 +42,7 @@ def _make_ctx(
|
||||
params=MagicMock(),
|
||||
is_hybrid_swa=is_hybrid_swa,
|
||||
is_hybrid_ssm=is_hybrid_ssm,
|
||||
is_dsa=is_dsa,
|
||||
enable_hierarchical_cache=enable_hierarchical_cache,
|
||||
disable_radix_cache=disable_radix_cache,
|
||||
effective_chunked_prefill_size=effective_chunked_prefill_size,
|
||||
@@ -291,13 +293,42 @@ class TestDefaultRadixCacheFactory(CustomTestCase):
|
||||
ctx.tp_worker.register_hicache_layer_transfer_counter.assert_called_once()
|
||||
self.assertIs(result, fake_radix.UnifiedRadixCache.return_value)
|
||||
|
||||
def test_unified_radix_cache_when_hierarchical_and_dsa(self):
|
||||
ctx = _make_ctx(enable_hierarchical_cache=True, is_dsa=True)
|
||||
# DSA models (e.g. DeepSeek V3.2 / GLM-5.1) with hierarchical cache
|
||||
# use UnifiedRadixCache.
|
||||
fake_components = MagicMock()
|
||||
fake_radix = MagicMock()
|
||||
with patch.dict(
|
||||
"sys.modules",
|
||||
{
|
||||
"sglang.srt.mem_cache.unified_cache_components": fake_components,
|
||||
"sglang.srt.mem_cache.unified_radix_cache": fake_radix,
|
||||
},
|
||||
):
|
||||
result = default_radix_cache_factory(ctx)
|
||||
fake_radix.UnifiedRadixCache.assert_called_once_with(ctx.params)
|
||||
fake_radix.UnifiedRadixCache.return_value.init_hicache.assert_called_once_with(
|
||||
ctx.server_args, ctx.params
|
||||
)
|
||||
ctx.tp_worker.register_hicache_layer_transfer_counter.assert_called_once()
|
||||
self.assertIs(result, fake_radix.UnifiedRadixCache.return_value)
|
||||
|
||||
def test_swa_radix_cache_when_hybrid_swa(self):
|
||||
ctx = _make_ctx(is_hybrid_swa=True)
|
||||
with patch("sglang.srt.mem_cache.swa_radix_cache.SWARadixCache") as SWA:
|
||||
SWA.return_value = MagicMock()
|
||||
# SWA hybrid models now default to the unified radix tree.
|
||||
fake_components = MagicMock()
|
||||
fake_radix = MagicMock()
|
||||
with patch.dict(
|
||||
"sys.modules",
|
||||
{
|
||||
"sglang.srt.mem_cache.unified_cache_components": fake_components,
|
||||
"sglang.srt.mem_cache.unified_radix_cache": fake_radix,
|
||||
},
|
||||
):
|
||||
result = default_radix_cache_factory(ctx)
|
||||
SWA.assert_called_once_with(params=ctx.params)
|
||||
self.assertIs(result, SWA.return_value)
|
||||
fake_radix.UnifiedRadixCache.assert_called_once_with(ctx.params)
|
||||
self.assertIs(result, fake_radix.UnifiedRadixCache.return_value)
|
||||
|
||||
def test_pure_swa_radix_cache_when_all_swa(self):
|
||||
ctx = _make_ctx(is_hybrid_swa=True, full_tokens_per_layer=0)
|
||||
@@ -311,11 +342,19 @@ class TestDefaultRadixCacheFactory(CustomTestCase):
|
||||
|
||||
def test_mamba_radix_cache_when_hybrid_ssm(self):
|
||||
ctx = _make_ctx(is_hybrid_ssm=True)
|
||||
with patch("sglang.srt.mem_cache.mamba_radix_cache.MambaRadixCache") as Mamba:
|
||||
Mamba.return_value = MagicMock()
|
||||
# Mamba hybrid models now default to the unified radix tree.
|
||||
fake_components = MagicMock()
|
||||
fake_radix = MagicMock()
|
||||
with patch.dict(
|
||||
"sys.modules",
|
||||
{
|
||||
"sglang.srt.mem_cache.unified_cache_components": fake_components,
|
||||
"sglang.srt.mem_cache.unified_radix_cache": fake_radix,
|
||||
},
|
||||
):
|
||||
result = default_radix_cache_factory(ctx)
|
||||
Mamba.assert_called_once_with(ctx.params)
|
||||
self.assertIs(result, Mamba.return_value)
|
||||
fake_radix.UnifiedRadixCache.assert_called_once_with(ctx.params)
|
||||
self.assertIs(result, fake_radix.UnifiedRadixCache.return_value)
|
||||
|
||||
def test_lmc_radix_cache_when_enable_lmcache(self):
|
||||
ctx = _make_ctx(enable_lmcache=True)
|
||||
|
||||
Reference in New Issue
Block a user