[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
@@ -16,7 +16,7 @@ register_cuda_ci(est_time=310, stage="extra-b", runner_config="8-gpu-h200")
@unittest.skipIf(is_in_ci(), "Temporarily disable the flaky test.")
class TestDisaggregationHybridAttentionMamba(PDDisaggregationServerBase):
class TestDisaggregationHybridAttentionGDN(PDDisaggregationServerBase):
@classmethod
def setUpClass(cls):
super().setUpClass()
@@ -88,7 +88,7 @@ class TestDisaggregationHybridAttentionMamba(PDDisaggregationServerBase):
self.assertGreater(metrics["score"], 0.93)
class TestDisaggregationHybridAttentionMambaExtraBuffer(PDDisaggregationServerBase):
class TestDisaggregationHybridAttentionGDNExtraBuffer(PDDisaggregationServerBase):
@classmethod
def setUpClass(cls):
super().setUpClass()
@@ -165,7 +165,7 @@ class TestDisaggregationHybridAttentionMambaExtraBuffer(PDDisaggregationServerBa
self.assertGreater(metrics["score"], 0.90)
class TestDisaggregationHybridAttentionMambaDPDecode(PDDisaggregationServerBase):
class TestDisaggregationHybridAttentionGDNDPDecode(PDDisaggregationServerBase):
"""Test with prefill tp=2 and decode tp=2/dp=2 with dp-attention enabled."""
@classmethod
@@ -244,5 +244,153 @@ class TestDisaggregationHybridAttentionMambaDPDecode(PDDisaggregationServerBase)
self.assertGreater(metrics["score"], 0.90)
class TestDisaggregationHybridAttentionMamba(PDDisaggregationServerBase):
@classmethod
def setUpClass(cls):
super().setUpClass()
cls.model = "nvidia/NVIDIA-Nemotron-Nano-9B-v2"
# Non blocking start servers
cls.start_prefill()
cls.start_decode()
# Block until both
cls.wait_server_ready(cls.prefill_url + "/health", process=cls.process_prefill)
cls.wait_server_ready(cls.decode_url + "/health", process=cls.process_decode)
cls.launch_lb()
@classmethod
def start_prefill(cls):
prefill_args = [
"--trust-remote-code",
"--disaggregation-mode",
"prefill",
"--disaggregation-bootstrap-port",
cls.bootstrap_port,
"--tp",
"4",
]
prefill_args += cls.transfer_backend + cls.rdma_devices
cls.process_prefill = popen_launch_pd_server(
cls.model,
cls.prefill_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
other_args=prefill_args,
)
@classmethod
def start_decode(cls):
decode_args = [
"--trust-remote-code",
"--disaggregation-mode",
"decode",
"--disaggregation-bootstrap-port",
cls.bootstrap_port,
"--tp",
"4",
"--base-gpu-id",
"4",
]
decode_args += cls.transfer_backend + cls.rdma_devices
cls.process_decode = popen_launch_pd_server(
cls.model,
cls.decode_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
other_args=decode_args,
)
def test_gsm8k(self):
args = SimpleNamespace(
base_url=self.base_url,
model=self.model,
eval_name="gsm8k",
api="completion",
max_tokens=512,
num_examples=200,
num_threads=128,
)
metrics = run_eval(args)
print(f"Evaluation metrics: {metrics}")
self.assertGreater(metrics["score"], 0.87)
class TestDisaggregationHybridAttentionMambaExtraBuffer(PDDisaggregationServerBase):
@classmethod
def setUpClass(cls):
super().setUpClass()
cls.model = "nvidia/NVIDIA-Nemotron-Nano-9B-v2"
# Non blocking start servers
cls.start_prefill()
cls.start_decode()
# Block until both
cls.wait_server_ready(cls.prefill_url + "/health", process=cls.process_prefill)
cls.wait_server_ready(cls.decode_url + "/health", process=cls.process_decode)
cls.launch_lb()
@classmethod
def start_prefill(cls):
prefill_args = [
"--trust-remote-code",
"--disaggregation-mode",
"prefill",
"--disaggregation-bootstrap-port",
cls.bootstrap_port,
"--tp",
"4",
"--mamba-scheduler-strategy",
"extra_buffer",
]
prefill_args += cls.transfer_backend + cls.rdma_devices
cls.process_prefill = popen_launch_pd_server(
cls.model,
cls.prefill_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
other_args=prefill_args,
)
@classmethod
def start_decode(cls):
decode_args = [
"--trust-remote-code",
"--disaggregation-mode",
"decode",
"--disaggregation-bootstrap-port",
cls.bootstrap_port,
"--tp",
"4",
"--base-gpu-id",
"4",
"--mamba-scheduler-strategy",
"extra_buffer",
]
decode_args += cls.transfer_backend + cls.rdma_devices
cls.process_decode = popen_launch_pd_server(
cls.model,
cls.decode_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
other_args=decode_args,
)
def test_gsm8k(self):
args = SimpleNamespace(
base_url=self.base_url,
model=self.model,
eval_name="gsm8k",
api="completion",
max_tokens=512,
num_examples=200,
num_threads=128,
)
metrics = run_eval(args)
print(f"Evaluation metrics: {metrics}")
self.assertGreater(metrics["score"], 0.87)
if __name__ == "__main__":
unittest.main()
@@ -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