Enable optimistic prefill for Mamba radix-cache models (#40184)
This commit is contained in:
@@ -13,7 +13,6 @@ from typing import Any
|
|||||||
from sglang.srt.arg_groups.overrides import (
|
from sglang.srt.arg_groups.overrides import (
|
||||||
declare_resolution,
|
declare_resolution,
|
||||||
model_config_of,
|
model_config_of,
|
||||||
resolved_view,
|
|
||||||
resolving_view,
|
resolving_view,
|
||||||
)
|
)
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
@@ -486,15 +485,6 @@ def handle_other_validations(server_args: Any):
|
|||||||
"_handle_other_validations",
|
"_handle_other_validations",
|
||||||
optimistic_prefill_attempts=0,
|
optimistic_prefill_attempts=0,
|
||||||
)
|
)
|
||||||
elif resolved_view(server_args).uses_mamba_radix_cache:
|
|
||||||
logger.warning(
|
|
||||||
"Optimistic prefill does not support models that use mamba radix cache."
|
|
||||||
)
|
|
||||||
declare_resolution(
|
|
||||||
server_args,
|
|
||||||
"_handle_other_validations",
|
|
||||||
optimistic_prefill_attempts=0,
|
|
||||||
)
|
|
||||||
|
|
||||||
# Handle model inference tensor dump.
|
# Handle model inference tensor dump.
|
||||||
if cfg.debug_tensor_dump_output_folder is not None:
|
if cfg.debug_tensor_dump_output_folder is not None:
|
||||||
|
|||||||
@@ -1510,7 +1510,10 @@ class SchedulerDisaggregationPrefillMixin:
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
self._release_aborted_request(req)
|
self._release_aborted_request(req)
|
||||||
release_kv_cache(req, self.tree_cache)
|
# Mamba insertion donates the checkpoint and clears its sequence marker.
|
||||||
|
release_kv_cache(
|
||||||
|
req, self.tree_cache, is_insert=not self.tree_cache.supports_mamba()
|
||||||
|
)
|
||||||
req.reset_for_retract()
|
req.reset_for_retract()
|
||||||
req.output_ids = array("q")
|
req.output_ids = array("q")
|
||||||
req.start_send_idx = 0
|
req.start_send_idx = 0
|
||||||
|
|||||||
@@ -4,20 +4,41 @@ import tempfile
|
|||||||
import time
|
import time
|
||||||
import unittest
|
import unittest
|
||||||
import uuid
|
import uuid
|
||||||
|
from array import array
|
||||||
from concurrent.futures import ThreadPoolExecutor, as_completed
|
from concurrent.futures import ThreadPoolExecutor, as_completed
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
|
from unittest.mock import patch
|
||||||
|
|
||||||
import requests
|
import requests
|
||||||
|
import torch
|
||||||
from prometheus_client.parser import text_string_to_metric_families
|
from prometheus_client.parser import text_string_to_metric_families
|
||||||
|
|
||||||
from sglang.srt.disaggregation.prefill import should_force_retry
|
from sglang.kernels.ops.attention.fla.chunk_delta_h import CHUNK_SIZE as FLA_CHUNK_SIZE
|
||||||
|
from sglang.srt.arg_groups.model_override_base import resolving_view
|
||||||
|
from sglang.srt.arg_groups.overrides import declare_resolution
|
||||||
|
from sglang.srt.arg_groups.serving_hook import handle_other_validations
|
||||||
|
from sglang.srt.configs.mamba_utils import Mamba2CacheParams, Mamba2StateShape
|
||||||
|
from sglang.srt.disaggregation.prefill import (
|
||||||
|
SchedulerDisaggregationPrefillMixin,
|
||||||
|
should_force_retry,
|
||||||
|
)
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
|
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 MatchPrefixParams
|
||||||
|
from sglang.srt.mem_cache.cache_init_params import CacheInitParams
|
||||||
|
from sglang.srt.mem_cache.mamba_radix_cache import MambaRadixCache
|
||||||
|
from sglang.srt.mem_cache.memory_pool import HybridLinearKVPool, HybridReqToTokenPool
|
||||||
|
from sglang.srt.mem_cache.radix_cache import RadixKey
|
||||||
|
from sglang.srt.sampling.sampling_params import SamplingParams
|
||||||
|
from sglang.srt.server_args import ServerArgs, set_global_server_args_for_scheduler
|
||||||
|
from sglang.srt.utils import get_device
|
||||||
from sglang.test.ci.ci_register import register_cuda_ci
|
from sglang.test.ci.ci_register import register_cuda_ci
|
||||||
from sglang.test.run_eval import run_eval
|
from sglang.test.run_eval import run_eval
|
||||||
from sglang.test.server_fixtures.disaggregation_fixture import (
|
from sglang.test.server_fixtures.disaggregation_fixture import (
|
||||||
PDDisaggregationServerBase,
|
PDDisaggregationServerBase,
|
||||||
)
|
)
|
||||||
from sglang.test.test_utils import DEFAULT_MODEL_NAME_FOR_TEST
|
from sglang.test.test_utils import DEFAULT_MODEL_NAME_FOR_TEST, CustomTestCase
|
||||||
|
|
||||||
register_cuda_ci(est_time=300, stage="base-b", runner_config="2-gpu-large")
|
register_cuda_ci(est_time=300, stage="base-b", runner_config="2-gpu-large")
|
||||||
|
|
||||||
@@ -291,5 +312,187 @@ class TestOptimisticPrefillL3BufferWriteThrough(
|
|||||||
time.sleep(1) # trigger memory check
|
time.sleep(1) # trigger memory check
|
||||||
|
|
||||||
|
|
||||||
|
class TestOptimisticPrefillMambaAdmission(CustomTestCase):
|
||||||
|
"""Mamba radix-cache models must keep optimistic prefill enabled.
|
||||||
|
|
||||||
|
Regression guard: the validation hook used to zero
|
||||||
|
``optimistic_prefill_attempts`` whenever resolution marked the config as
|
||||||
|
using the mamba radix cache, silently disabling optimistic prefill for
|
||||||
|
hybrid-mamba models in disaggregated prefill. The PP and HiCache
|
||||||
|
write-policy restrictions are independent of that gate and must remain.
|
||||||
|
"""
|
||||||
|
|
||||||
|
BASE = dict(disaggregation_mode="prefill", optimistic_prefill_attempts=3)
|
||||||
|
|
||||||
|
def _make_args(self, **fields) -> ServerArgs:
|
||||||
|
server_args = ServerArgs(model_path="dummy")
|
||||||
|
for name, value in fields.items():
|
||||||
|
setattr(server_args, name, value)
|
||||||
|
return server_args
|
||||||
|
|
||||||
|
def _resolved_attempts(self, server_args) -> int:
|
||||||
|
handle_other_validations(server_args)
|
||||||
|
return resolving_view(server_args).optimistic_prefill_attempts
|
||||||
|
|
||||||
|
def test_mamba_radix_cache_keeps_optimistic_prefill(self):
|
||||||
|
server_args = self._make_args(**self.BASE)
|
||||||
|
# Simulate resolution having identified a mamba radix-cache model.
|
||||||
|
declare_resolution(server_args, "test", uses_mamba_radix_cache=True)
|
||||||
|
self.assertEqual(self._resolved_attempts(server_args), 3)
|
||||||
|
|
||||||
|
def test_pp_restriction_retained(self):
|
||||||
|
server_args = self._make_args(pp_size=2, **self.BASE)
|
||||||
|
self.assertEqual(self._resolved_attempts(server_args), 0)
|
||||||
|
|
||||||
|
def test_hicache_write_policy_restriction_retained(self):
|
||||||
|
server_args = self._make_args(
|
||||||
|
enable_hierarchical_cache=True,
|
||||||
|
hicache_write_policy="write_through",
|
||||||
|
**self.BASE,
|
||||||
|
)
|
||||||
|
self.assertEqual(self._resolved_attempts(server_args), 0)
|
||||||
|
|
||||||
|
|
||||||
|
class TestOptimisticPrefillMambaRetryRelease(CustomTestCase):
|
||||||
|
"""Optimistic retry cleanup must not treat the donated Mamba checkpoint
|
||||||
|
as a second donation.
|
||||||
|
|
||||||
|
Bug mechanism: the retry path first inserts the unfinished prefix, which
|
||||||
|
donates the tracked checkpoint and clears ``mamba_last_track_seqlen``.
|
||||||
|
Releasing with ``is_insert=True`` afterwards re-enters the donation path
|
||||||
|
with the cleared marker, inserting a zero-length radix entry that pins a
|
||||||
|
clone of the request's live state. Releasing with ``is_insert=False``
|
||||||
|
retains the donated prefix/checkpoint and frees only the uncached tail
|
||||||
|
and the request-owned Mamba buffers.
|
||||||
|
"""
|
||||||
|
|
||||||
|
SIZE = 128
|
||||||
|
MAMBA_SIZE = 8
|
||||||
|
TRACK_SEQLEN = 4
|
||||||
|
PROMPT = [1, 2, 3, 4, 5, 6, 7, 8]
|
||||||
|
|
||||||
|
def _setup_mamba_tree(self):
|
||||||
|
server_args = ServerArgs(model_path="dummy", page_size=1)
|
||||||
|
# MambaRadixCache reads mamba_cache_chunk_size, whose property
|
||||||
|
# otherwise loads the HF config for the dummy model.
|
||||||
|
server_args._mamba_cache_chunk_size = FLA_CHUNK_SIZE
|
||||||
|
set_global_server_args_for_scheduler(server_args)
|
||||||
|
num_layers = 48
|
||||||
|
global_interval = 4
|
||||||
|
full_attention_layer_ids = [
|
||||||
|
i for i in range(global_interval - 1, num_layers, global_interval)
|
||||||
|
]
|
||||||
|
mamba_layers = [
|
||||||
|
i for i in range(num_layers) if i not in full_attention_layer_ids
|
||||||
|
]
|
||||||
|
device = get_device()
|
||||||
|
with envs.SGLANG_MAMBA_SSM_DTYPE.override("bfloat16"):
|
||||||
|
shape = Mamba2StateShape.create(
|
||||||
|
tp_world_size=1,
|
||||||
|
intermediate_size=4096,
|
||||||
|
n_groups=16,
|
||||||
|
num_heads=32,
|
||||||
|
head_dim=128,
|
||||||
|
state_size=128,
|
||||||
|
conv_kernel=4,
|
||||||
|
)
|
||||||
|
cache_params = Mamba2CacheParams(shape=shape, layers=mamba_layers)
|
||||||
|
req_to_token_pool = HybridReqToTokenPool(
|
||||||
|
size=4,
|
||||||
|
mamba_size=self.MAMBA_SIZE,
|
||||||
|
mamba_spec_state_size=4,
|
||||||
|
max_context_len=64,
|
||||||
|
device=device,
|
||||||
|
enable_memory_saver=False,
|
||||||
|
cache_params=cache_params,
|
||||||
|
mamba_layer_ids=mamba_layers,
|
||||||
|
enable_mamba_extra_buffer=True,
|
||||||
|
speculative_num_draft_tokens=3,
|
||||||
|
)
|
||||||
|
pool = HybridLinearKVPool(
|
||||||
|
size=self.SIZE,
|
||||||
|
dtype=torch.bfloat16,
|
||||||
|
page_size=1,
|
||||||
|
head_num=2,
|
||||||
|
head_dim=256,
|
||||||
|
full_attention_layer_ids=full_attention_layer_ids,
|
||||||
|
device=device,
|
||||||
|
enable_memory_saver=False,
|
||||||
|
mamba_pool=req_to_token_pool.mamba_pool,
|
||||||
|
)
|
||||||
|
allocator = TokenToKVPoolAllocator(
|
||||||
|
size=self.SIZE,
|
||||||
|
dtype=torch.bfloat16,
|
||||||
|
device=device,
|
||||||
|
kvcache=pool,
|
||||||
|
need_sort=False,
|
||||||
|
)
|
||||||
|
tree = MambaRadixCache(
|
||||||
|
params=CacheInitParams(
|
||||||
|
disable=False,
|
||||||
|
req_to_token_pool=req_to_token_pool,
|
||||||
|
token_to_kv_pool_allocator=allocator,
|
||||||
|
page_size=1,
|
||||||
|
enable_mamba_extra_buffer=True,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
return tree, allocator, req_to_token_pool
|
||||||
|
|
||||||
|
def test_retry_release_retains_donated_checkpoint(self):
|
||||||
|
tree, allocator, req_to_token_pool = self._setup_mamba_tree()
|
||||||
|
|
||||||
|
# A request that finished optimistic prefill of PROMPT and tracked one
|
||||||
|
# Mamba checkpoint at TRACK_SEQLEN.
|
||||||
|
req = Req(
|
||||||
|
rid="optimistic-mamba-retry",
|
||||||
|
origin_input_text="",
|
||||||
|
origin_input_ids=array("q", self.PROMPT),
|
||||||
|
sampling_params=SamplingParams(max_new_tokens=1),
|
||||||
|
)
|
||||||
|
req_to_token_pool.alloc([req])
|
||||||
|
kv_indices = allocator.alloc(len(self.PROMPT))
|
||||||
|
req_to_token_pool.write(
|
||||||
|
(req.kv.req_pool_idx, slice(0, len(self.PROMPT))), kv_indices
|
||||||
|
)
|
||||||
|
req.full_untruncated_fill_ids = array("q", self.PROMPT)
|
||||||
|
req.set_extend_range(0, len(self.PROMPT))
|
||||||
|
req.kv.kv_committed_len = len(self.PROMPT)
|
||||||
|
req.kv.kv_allocated_len = len(self.PROMPT)
|
||||||
|
req.kv.mamba_last_track_seqlen = self.TRACK_SEQLEN
|
||||||
|
req.last_node = tree.root_node
|
||||||
|
|
||||||
|
scheduler = SimpleNamespace(
|
||||||
|
tree_cache=tree,
|
||||||
|
waiting_queue=[],
|
||||||
|
disagg_prefill_bootstrap_queue=SimpleNamespace(queue=[]),
|
||||||
|
metrics_reporter=SimpleNamespace(enable_metrics=False),
|
||||||
|
processed_tokens_counter=0,
|
||||||
|
_release_aborted_request=lambda req: None,
|
||||||
|
clear_pending_chunk_send=lambda req: None,
|
||||||
|
)
|
||||||
|
with patch(
|
||||||
|
"sglang.srt.disaggregation.prefill.get_disagg",
|
||||||
|
return_value=SimpleNamespace(optimistic_prefill_attempts=3),
|
||||||
|
):
|
||||||
|
SchedulerDisaggregationPrefillMixin.optimistic_release_and_requeue(
|
||||||
|
scheduler, req
|
||||||
|
)
|
||||||
|
|
||||||
|
# The donated prefix and exactly one checkpoint stay in the tree; the
|
||||||
|
# old double-donation path either asserts or pins a second state.
|
||||||
|
self.assertEqual(tree.total_size(), (self.TRACK_SEQLEN, 1))
|
||||||
|
match = tree.match_prefix(
|
||||||
|
MatchPrefixParams(key=RadixKey(array("q", self.PROMPT)))
|
||||||
|
)
|
||||||
|
self.assertIsNotNone(match.last_device_node.mamba_value)
|
||||||
|
|
||||||
|
# Only the uncached tail and the request-owned Mamba buffers are
|
||||||
|
# freed; the tree keeps the donated checkpoint slot.
|
||||||
|
self.assertEqual(allocator.available_size(), self.SIZE - self.TRACK_SEQLEN)
|
||||||
|
self.assertEqual(
|
||||||
|
req_to_token_pool.mamba_allocator.available_size(), self.MAMBA_SIZE - 1
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
Reference in New Issue
Block a user