[feat] Support extra_buffer in Mamba2-based models (#15829)
Signed-off-by: Roi Koren <roik@nvidia.com>
This commit is contained in:
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user