[feat] Support extra_buffer in Mamba2-based models (#15829)

Signed-off-by: Roi Koren <roik@nvidia.com>
This commit is contained in:
roikoren755
2026-05-26 16:03:29 +08:00
committed by GitHub
parent 7e6e5efe51
commit e958f4561f
17 changed files with 405 additions and 130 deletions
@@ -6,6 +6,7 @@ import torch
from sglang.srt.configs.mamba_utils import Mamba2CacheParams, Mamba2StateShape
from sglang.srt.disaggregation.kv_events import BlockRemoved, BlockStored
from sglang.srt.environ import envs
from sglang.srt.layers.attention.fla.chunk_delta_h import CHUNK_SIZE as FLA_CHUNK_SIZE
from sglang.srt.managers.schedule_batch import Req
from sglang.srt.mem_cache.allocator import TokenToKVPoolAllocator
from sglang.srt.mem_cache.base_prefix_cache import (
@@ -410,9 +411,12 @@ class TestMamba(unittest.TestCase):
def _setup_tree_and_allocator(self, enable_kv_cache_events=False):
"""Helper to create a MambaRadixCache with allocator for testing."""
set_global_server_args_for_scheduler(
ServerArgs(model_path="dummy", page_size=1)
)
server_args = ServerArgs(model_path="dummy", page_size=1)
# MambaRadixCache reads mamba_cache_chunk_size, whose property otherwise
# loads the HF config for self.model_path — impossible for the dummy model.
# Mirror the property's default for a dummy HF config: FLA_CHUNK_SIZE.
server_args._mamba_cache_chunk_size = FLA_CHUNK_SIZE
set_global_server_args_for_scheduler(server_args)
size = 128
dtype = torch.bfloat16
head_num = 2
@@ -21,6 +21,7 @@ import torch
from sglang.srt.configs.mamba_utils import Mamba2CacheParams, Mamba2StateShape
from sglang.srt.environ import envs
from sglang.srt.layers.attention.fla.chunk_delta_h import CHUNK_SIZE as FLA_CHUNK_SIZE
from sglang.srt.mem_cache.allocator import TokenToKVPoolAllocator
from sglang.srt.mem_cache.base_prefix_cache import (
DecLockRefParams,
@@ -691,9 +692,11 @@ def run_all_benchmarks(
if benchmarks is None or "all" in benchmarks:
benchmarks = list(ALL_BENCHMARKS.keys())
set_global_server_args_for_scheduler(
ServerArgs(model_path="dummy", page_size=page_size)
)
server_args = ServerArgs(model_path="dummy", page_size=page_size)
# MambaRadixCache reads mamba_cache_chunk_size, whose property otherwise
# loads the HF config for self.model_path — impossible for the dummy model.
server_args._mamba_cache_chunk_size = max(FLA_CHUNK_SIZE, page_size)
set_global_server_args_for_scheduler(server_args)
impl_name = (tree_cls or UnifiedRadixCache).__name__
results = []
@@ -780,9 +783,11 @@ class _BenchSuite:
@classmethod
def setUpClass(cls):
set_global_server_args_for_scheduler(
ServerArgs(model_path="dummy", page_size=cls.bench_cfg["page_size"])
)
page_size = cls.bench_cfg["page_size"]
server_args = ServerArgs(model_path="dummy", page_size=page_size)
# See run_all_benchmarks for why _mamba_cache_chunk_size is preset.
server_args._mamba_cache_chunk_size = max(FLA_CHUNK_SIZE, page_size)
set_global_server_args_for_scheduler(server_args)
def _run(self, bench_fn):
cfg = self.bench_cfg
@@ -10,6 +10,7 @@ import torch
from sglang.srt.configs.mamba_utils import Mamba2CacheParams, Mamba2StateShape
from sglang.srt.environ import envs
from sglang.srt.layers.attention.fla.chunk_delta_h import CHUNK_SIZE as FLA_CHUNK_SIZE
from sglang.srt.managers.schedule_batch import Req
from sglang.srt.mem_cache.allocator import TokenToKVPoolAllocator
from sglang.srt.mem_cache.base_prefix_cache import (
@@ -114,9 +115,12 @@ class CacheConfig:
def build_fixture(cfg: CacheConfig):
"""Create (tree, allocator, req_to_token_pool) from a CacheConfig."""
set_global_server_args_for_scheduler(
ServerArgs(model_path="dummy", page_size=cfg.page_size)
)
server_args = ServerArgs(model_path="dummy", page_size=cfg.page_size)
# MambaRadixCache reads mamba_cache_chunk_size, whose property otherwise
# loads the HF config for self.model_path — impossible for the dummy model.
# Mirror the property's default for a dummy HF config: FLA_CHUNK_SIZE.
server_args._mamba_cache_chunk_size = max(FLA_CHUNK_SIZE, cfg.page_size)
set_global_server_args_for_scheduler(server_args)
device = get_device()
mamba2_cache_params = None
@@ -1336,6 +1340,8 @@ class UnifiedRadixCacheSuite:
hicache_io_backend="direct",
hicache_write_policy=write_policy,
)
# See build_fixture for why _mamba_cache_chunk_size is preset.
server_args._mamba_cache_chunk_size = max(FLA_CHUNK_SIZE, self.cfg.page_size)
set_global_server_args_for_scheduler(server_args)
tree.init_hicache(server_args, tree.cache_init_params)
tree.write_through_threshold = 1 << 30