[Mamba] extra buffer lazy support (#27118)

Co-authored-by: YAMY <74099316+YAMY1234@users.noreply.github.com>
This commit is contained in:
Hanming Lu
2026-06-03 12:42:11 -07:00
committed by GitHub
co-authored by YAMY
parent ac99794e64
commit e0b692600f
14 changed files with 429 additions and 94 deletions
+4
View File
@@ -263,6 +263,10 @@ class Envs:
SGLANG_TEST_RETRACT = EnvBool(False)
SGLANG_TEST_RETRACT_INTERVAL = EnvInt(3)
SGLANG_TEST_RETRACT_NO_PREFILL_BS = EnvInt(2 ** 31)
# Scheduler: force lazy extra_buffer prealloc to fail at decode boundaries
SGLANG_TEST_MAMBA_LAZY_ALLOC_FAIL = EnvBool(False)
# KL tests: skip the cache-hit count assertion (e.g. when alloc failure reduces hits)
SGLANG_TEST_SKIP_CACHE_HIT_ASSERT = EnvBool(False)
SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY = EnvInt(0)
SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_IDLE = EnvBool(True)
+49 -5
View File
@@ -75,6 +75,7 @@ from sglang.srt.managers.scheduler_components.new_token_ratio_tracker import (
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
from sglang.srt.mem_cache.base_prefix_cache import (
BasePrefixCache,
EvictParams,
MatchPrefixParams,
zero_match_result,
)
@@ -767,6 +768,9 @@ class Req(ReqDllmMixin):
self.mamba_cow_src_index: Optional[torch.Tensor] = None
# Deferred clear: newly allocated mamba slot needs zeroing on forward stream
self.mamba_needs_clear: bool = False
# Lazy extra buffer: skip radix cache insert when prealloc failed at
# boundary — the forward overwrites the only slot, corrupting the state.
self.mamba_lazy_is_insert: bool = True
# Check finish
self.tokenizer = None
@@ -2123,11 +2127,15 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
# so we need to add 1 to the seqlen to retrieve the correct mamba state from h.
mamba_track_seqlen = _force_track_h(mamba_track_seqlen_aligned)
req.mamba_next_track_idx = (
self.req_to_token_pool.get_mamba_ping_pong_other_idx(
req.mamba_next_track_idx
# In lazy mode, skip the swap — the second ping-pong slot is not
# allocated yet; it will be allocated on demand at the track boundary
# in mamba_lazy_prealloc_at_boundary during prepare_for_decode.
if not get_global_server_args().enable_mamba_extra_buffer_lazy():
req.mamba_next_track_idx = (
self.req_to_token_pool.get_mamba_ping_pong_other_idx(
req.mamba_next_track_idx
)
)
)
if req.mamba_branching_seqlen is not None:
# track branching point in this forward if the branching point
# is within the current extend batch.
@@ -2380,6 +2388,38 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
assert not ret or self.spec_algorithm.supports_spec_v2()
return ret
def mamba_lazy_prealloc_at_boundary(self, mamba_track_interval: int):
"""Allocate a temporary second ping-pong slot for reqs at a track boundary.
In lazy mode each request normally holds only 1 ping-pong slot.
When seq_len hits a track interval boundary, we allocate the
second slot so the forward pass can write the new tracked state
there. The old slot is freed after the forward in
mamba_lazy_post_decode_at_boundary.
"""
pool = self.req_to_token_pool
for i, req in enumerate(self.reqs):
buf = req.mamba_ping_pong_track_buffer
assert buf is not None
# Skip reqs not at a track boundary
if self.seq_lens_cpu[i].item() % mamba_track_interval != 0:
continue
other_idx = 1 - req.mamba_next_track_idx
if buf[other_idx].item() != -1:
# With overlap the previous forward's post-processing
# (which frees this slot) hasn't run yet. Skip.
continue
if envs.SGLANG_TEST_MAMBA_LAZY_ALLOC_FAIL.get():
new_slot = None
else:
new_slot = pool.mamba_pool.alloc(1)
if new_slot is None:
self.tree_cache.evict(EvictParams(num_tokens=0, mamba_num=1))
new_slot = pool.mamba_pool.alloc(1)
if new_slot is not None:
pool.set_mamba_ping_pong_slot(req, other_idx, new_slot[0])
req.mamba_next_track_idx = other_idx
def prepare_for_decode(self):
self.forward_mode = ForwardMode.DECODE
bs = len(self.reqs)
@@ -2460,16 +2500,20 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
)
if get_global_server_args().enable_mamba_extra_buffer():
mamba_track_interval = get_global_server_args().mamba_track_interval
if len(self.reqs) == 0:
self.mamba_track_indices = torch.empty(
(0,), dtype=torch.int64, device=self.device
)
else:
if get_global_server_args().enable_mamba_extra_buffer_lazy():
self.mamba_lazy_prealloc_at_boundary(mamba_track_interval)
set_mamba_track_indices_from_reqs(self)
# async H2D
self.mamba_track_mask = (
(self.seq_lens_cpu % get_global_server_args().mamba_track_interval == 0)
(self.seq_lens_cpu % mamba_track_interval == 0)
.pin_memory()
.to(device=self.device, non_blocking=True)
)
@@ -639,9 +639,10 @@ class SchedulerBatchResultProcessor:
continue
if is_spec_v1:
self._mamba_prefix_cache_update(req, batch, result, i)
req.time_stats.set_last_decode_finish_time()
self._handle_finished_req(req, i, logits_output)
self._handle_finish_state_updated_req(
req, batch, result, i, logits_output
)
if req.return_hidden_states and logits_output.hidden_states is not None:
req.hidden_states.append(
logits_output.hidden_states[i].cpu().clone().tolist()
@@ -661,12 +662,10 @@ class SchedulerBatchResultProcessor:
self._maybe_update_reasoning_tokens(req, next_token_id)
# Update Mamba last track seqlen
self._mamba_prefix_cache_update(req, batch, result, i)
req.time_stats.set_last_decode_finish_time()
req.update_finish_state(new_accepted_len)
self._handle_finished_req(req, i, logits_output)
self._handle_finish_state_updated_req(req, batch, result, i, logits_output)
if req.return_logprob:
self._apply_decode_logprobs(
@@ -802,12 +801,18 @@ class SchedulerBatchResultProcessor:
self.abort_request(AbortReq(rid=req.rid))
req.grammar.finished = req.finished()
def _handle_finished_req(
def _handle_finish_state_updated_req(
self,
req: Req,
batch: ScheduleBatch,
result: GenerationBatchResult,
i: int,
logits_output: LogitsProcessorOutput,
):
# Called here (after update_finish_state) so req.finished() is valid
# for mamba_lazy_post_decode_at_boundary inside.
self._mamba_prefix_cache_update(req, batch, result, i)
if (
self.server_args.disaggregation_decode_enable_offload_kvcache
and not req.finished()
@@ -833,7 +838,12 @@ class SchedulerBatchResultProcessor:
)
if callable(prepare_release):
prepare_release(req)
release_kv_cache(req, self.tree_cache)
is_insert = (
req.mamba_lazy_is_insert
if get_global_server_args().enable_mamba_extra_buffer_lazy()
else True
)
release_kv_cache(req, self.tree_cache, is_insert=is_insert)
req.time_stats.set_completion_time()
@@ -855,33 +865,80 @@ class SchedulerBatchResultProcessor:
result: GenerationBatchResult,
i: int,
) -> None:
seq_len = len(req.origin_input_ids) + len(req.output_ids) - 1
if req.mamba_ping_pong_track_buffer is not None:
mamba_track_interval = get_global_server_args().mamba_track_interval
if batch.spec_algorithm.is_none() and seq_len % mamba_track_interval == 0:
# for non-spec decode, we update mamba_last_track_seqlen at the end of each track interval
req.mamba_next_track_idx = (
batch.req_to_token_pool.get_mamba_ping_pong_other_idx(
req.mamba_next_track_idx
)
"""Update mamba track state at ping-pong boundaries.
Non-lazy: swap the ping-pong index so the next forward writes to
the alternate slot.
Lazy: keep the same index (prealloc handles the swap) and run
post-decode cleanup to free the temporary second slot.
"""
if req.mamba_ping_pong_track_buffer is None:
return
lazy = get_global_server_args().enable_mamba_extra_buffer_lazy()
at_boundary, track_seqlen = self._mamba_check_track_boundary(
req, batch, result, i
)
if not at_boundary:
return
req.mamba_last_track_seqlen = track_seqlen
if lazy:
self.mamba_lazy_post_decode_at_boundary(req, batch)
else:
req.mamba_next_track_idx = (
batch.req_to_token_pool.get_mamba_ping_pong_other_idx(
req.mamba_next_track_idx
)
req.mamba_last_track_seqlen = seq_len
elif (
not batch.spec_algorithm.is_none()
and result.num_correct_drafts_per_req_cpu is not None
):
# for spec decode, update mamba_last_track_seqlen if this iteration crosses a track interval
actual_seq_len = req.seqlen - 1
if (
actual_seq_len // mamba_track_interval
!= (actual_seq_len - result.num_correct_drafts_per_req_cpu[i] - 1)
// mamba_track_interval
):
req.mamba_next_track_idx = (
batch.req_to_token_pool.get_mamba_ping_pong_other_idx(
req.mamba_next_track_idx
)
)
req.mamba_last_track_seqlen = (
actual_seq_len // mamba_track_interval * mamba_track_interval
)
)
def _mamba_check_track_boundary(self, req, batch, result, i):
"""Check if this decode step crosses a mamba track interval boundary.
Returns (at_boundary, track_seqlen). The boundary condition
matches what the forward's tracking mask used:
``prepare_for_decode`` increments both ``seq_lens_cpu`` and
``kv_committed_len`` by 1, then checks
``seq_lens_cpu % interval == 0``. Using ``kv_committed_len``
here reproduces that check exactly, and the value is always a
multiple of ``interval`` (hence page-aligned).
For spec decode, the boundary is detected by comparing the
accepted seq_len range against interval boundaries.
"""
interval = get_global_server_args().mamba_track_interval
if batch.spec_algorithm.is_none():
if req.kv_committed_len % interval == 0:
return True, req.kv_committed_len
elif result.num_correct_drafts_per_req_cpu is not None:
cur = req.seqlen - 1
prev = cur - result.num_correct_drafts_per_req_cpu[i] - 1
if cur // interval != prev // interval:
return True, cur // interval * interval
return False, 0
def mamba_lazy_post_decode_at_boundary(self, req: Req, batch: ScheduleBatch):
"""Post-decode cleanup at a lazy-mode track boundary.
Finished reqs: if prealloc failed (other slot is -1), the forward
overwrote the only slot with corrupted state, so mark
is_insert=False to skip the cache insert. If the other slot is
occupied (stale prealloc from an overlap extra forward), free it
so the prealloc assert in the next prepare_for_decode holds.
Running reqs: free the old ping-pong slot so we go back to
holding only 1 slot until the next boundary.
"""
other_idx = 1 - req.mamba_next_track_idx
other_val = req.mamba_ping_pong_track_buffer[other_idx].item()
if other_val != -1:
pool = batch.req_to_token_pool
pool.mamba_pool.free(
req.mamba_ping_pong_track_buffer[other_idx].unsqueeze(0)
)
pool.set_mamba_ping_pong_slot(req, other_idx, -1)
elif req.finished():
req.mamba_lazy_is_insert = False
@@ -32,6 +32,7 @@ class CacheInitParams:
enable_kv_cache_events: bool = False
enable_mamba_extra_buffer: bool = False
enable_mamba_extra_buffer_lazy: bool = False
pp_rank: int = 0
pp_size: int = 1
+10 -5
View File
@@ -31,6 +31,8 @@ if TYPE_CHECKING:
# Needs 2 + 1 slots for mamba request with prefix cache. 2 for ping pong cache, 1 for running mamba state.
MAMBA_STATE_PER_REQ_PREFIX_CACHE = 3
# Lazy mode: 1 + 1 slots (1 ping-pong + 1 running), second ping-pong allocated on demand at boundary.
MAMBA_STATE_PER_REQ_PREFIX_CACHE_LAZY = 2
MAMBA_STATE_PER_REQ_NO_CACHE = 1
logger = logging.getLogger(__name__)
@@ -255,11 +257,14 @@ def alloc_req_slots(
num_reqs = len(reqs)
if isinstance(req_to_token_pool, HybridReqToTokenPool):
mamba_available_size = req_to_token_pool.mamba_pool.available_size()
factor = (
MAMBA_STATE_PER_REQ_PREFIX_CACHE
if tree_cache.supports_mamba()
else MAMBA_STATE_PER_REQ_NO_CACHE
)
if tree_cache.supports_mamba():
factor = (
MAMBA_STATE_PER_REQ_PREFIX_CACHE_LAZY
if req_to_token_pool.enable_mamba_extra_buffer_lazy
else MAMBA_STATE_PER_REQ_PREFIX_CACHE
)
else:
factor = MAMBA_STATE_PER_REQ_NO_CACHE
mamba_state_needed = num_reqs * factor
if mamba_available_size < mamba_state_needed:
if tree_cache is not None and tree_cache.supports_mamba():
@@ -218,6 +218,7 @@ def build_kv_cache(
enable_metrics=enable_metrics,
enable_kv_cache_events=enable_kv_cache_events,
enable_mamba_extra_buffer=server_args.enable_mamba_extra_buffer(),
enable_mamba_extra_buffer_lazy=server_args.enable_mamba_extra_buffer_lazy(),
pp_rank=ps.pp_rank,
pp_size=ps.pp_size,
chunked_prefill_size=effective_chunked_prefill_size,
@@ -431,6 +431,7 @@ class MambaRadixCache(KVCacheEventMixin, BasePrefixCache):
self.disable = params.disable
self.enable_kv_cache_events = params.enable_kv_cache_events
self.enable_mamba_extra_buffer = params.enable_mamba_extra_buffer
self.enable_mamba_extra_buffer_lazy = params.enable_mamba_extra_buffer_lazy
self.kv_event_queue = []
if not self.enable_mamba_extra_buffer:
@@ -559,9 +560,7 @@ class MambaRadixCache(KVCacheEventMixin, BasePrefixCache):
# insert the token_ids and kv_indices into the radix tree
if self.enable_mamba_extra_buffer:
mamba_ping_pong_track_buffer_to_keep = (
self.req_to_token_pool.get_mamba_ping_pong_other_idx(
req.mamba_next_track_idx
)
self.req_to_token_pool.get_mamba_ping_pong_keep_idx(req)
)
mamba_value = (
req.mamba_ping_pong_track_buffer[
@@ -570,6 +569,13 @@ class MambaRadixCache(KVCacheEventMixin, BasePrefixCache):
.unsqueeze(-1)
.clone()
)
assert mamba_value.item() != -1, (
f"Cached mamba slot is -1: keep_idx={mamba_ping_pong_track_buffer_to_keep}, "
f"buf={req.mamba_ping_pong_track_buffer.tolist()}, "
f"next_track_idx={req.mamba_next_track_idx}, "
f"last_track_seqlen={req.mamba_last_track_seqlen}, "
f"rid={req.rid}"
)
else:
mamba_value = req.mamba_pool_idx.unsqueeze(-1).clone()
mamba_ping_pong_track_buffer_to_keep = None
@@ -644,23 +650,10 @@ class MambaRadixCache(KVCacheEventMixin, BasePrefixCache):
# Donate the mamba index to the radix cache instead of copying.
# This avoids a data copy that would race with the forward stream.
if self.enable_mamba_extra_buffer:
mamba_ping_pong_track_buffer_to_keep = (
self.req_to_token_pool.get_mamba_ping_pong_other_idx(
req.mamba_next_track_idx
)
)
mamba_value_donated = (
req.mamba_ping_pong_track_buffer[mamba_ping_pong_track_buffer_to_keep]
.unsqueeze(-1)
.clone()
)
new_slot = self._alloc_mamba_slot()
req.mamba_ping_pong_track_buffer[mamba_ping_pong_track_buffer_to_keep] = (
new_slot[0]
mamba_value_donated = self.req_to_token_pool.donate_mamba_ping_pong_slot(
req, new_slot
)
self.req_to_token_pool.req_index_to_mamba_ping_pong_track_buffer_mapping[
req.req_pool_idx
] = req.mamba_ping_pong_track_buffer
else:
mamba_value_donated = self._alloc_mamba_slot()
self.req_to_token_pool.mamba_pool.copy_from(
+82 -7
View File
@@ -138,6 +138,8 @@ def _set_kv_buffer_impl(
class ReqToTokenPool:
"""A memory pool that maps a request to its token locations."""
enable_mamba_extra_buffer_lazy: bool = False
def __init__(
self,
size: int,
@@ -511,6 +513,7 @@ class HybridReqToTokenPool(ReqToTokenPool):
cache_params: BaseLinearStateParams,
mamba_layer_ids: List[int],
enable_mamba_extra_buffer: bool,
enable_mamba_extra_buffer_lazy: bool = False,
speculative_num_draft_tokens: int = None,
enable_overlap_schedule: bool = True,
start_layer: Optional[int] = None,
@@ -524,6 +527,7 @@ class HybridReqToTokenPool(ReqToTokenPool):
self.mamba_ping_pong_track_buffer_size = 2 if enable_overlap_schedule else 1
self.enable_mamba_extra_buffer = enable_mamba_extra_buffer
self.enable_mamba_extra_buffer_lazy = enable_mamba_extra_buffer_lazy
self.enable_memory_saver = enable_memory_saver
self.start_layer = start_layer if start_layer is not None else 0
self.layer_transfer_counter = None
@@ -599,13 +603,7 @@ class HybridReqToTokenPool(ReqToTokenPool):
mamba_indices.append(req.mamba_pool_idx)
if self.enable_mamba_extra_buffer:
if req.mamba_ping_pong_track_buffer is None:
req.mamba_ping_pong_track_buffer = self.mamba_pool.alloc(
self.mamba_ping_pong_track_buffer_size
)
assert (
req.mamba_ping_pong_track_buffer is not None
), "Not enough space for mamba ping pong idx, try to increase --mamba-full-memory-ratio."
req.mamba_next_track_idx = 0
self._alloc_ping_pong_buffer(req)
mamba_ping_pong_track_buffers.append(req.mamba_ping_pong_track_buffer)
assert len(select_index) == len(
mamba_indices
@@ -647,6 +645,77 @@ class HybridReqToTokenPool(ReqToTokenPool):
else:
return mamba_next_track_idx
def get_mamba_ping_pong_keep_idx(self, req: "Req") -> int:
"""Return the ping-pong index holding the most recent tracked state.
In lazy mode the valid state stays at next_track_idx (no eager swap).
In normal mode it is at the "other" index (swapped after each track).
"""
if self.enable_mamba_extra_buffer_lazy:
return req.mamba_next_track_idx
return self.get_mamba_ping_pong_other_idx(req.mamba_next_track_idx)
def _alloc_ping_pong_buffer(self, req: "Req"):
"""Allocate the ping-pong track buffer for a new request.
Lazy mode allocates 1 slot with the second set to -1 (allocated
on demand at track boundaries). Normal mode allocates all slots upfront.
"""
n = (
1
if self.enable_mamba_extra_buffer_lazy
else self.mamba_ping_pong_track_buffer_size
)
slots = self.mamba_pool.alloc(n)
assert slots is not None, (
"Not enough space for mamba ping pong idx, "
"try to increase --mamba-full-memory-ratio."
)
buf = torch.full(
(self.mamba_ping_pong_track_buffer_size,),
-1,
dtype=slots.dtype,
device=slots.device,
)
buf[:n] = slots
req.mamba_ping_pong_track_buffer = buf
req.mamba_next_track_idx = 0
def set_mamba_ping_pong_slot(self, req: "Req", idx: int, value):
"""Update a ping-pong slot value and sync the device-side mapping.
The req holds the authoritative buffer; this keeps the
req_index_to_mamba_ping_pong_track_buffer_mapping in sync so that
set_mamba_track_indices_from_reqs reads correct slot indices.
"""
req.mamba_ping_pong_track_buffer[idx] = value
self.req_index_to_mamba_ping_pong_track_buffer_mapping[req.req_pool_idx] = (
req.mamba_ping_pong_track_buffer
)
def donate_mamba_ping_pong_slot(
self, req: "Req", new_slot: torch.Tensor
) -> torch.Tensor:
"""Donate the tracked-state ping-pong slot to the radix cache.
Returns the old slot index (shape [1]) for cache insertion and
replaces it with new_slot so the request can continue tracking.
In lazy mode the valid state is at next_track_idx; in normal mode
it is at the "other" index.
"""
donate_idx = self.get_mamba_ping_pong_keep_idx(req)
mamba_value_donated = (
req.mamba_ping_pong_track_buffer[donate_idx].unsqueeze(-1).clone()
)
assert mamba_value_donated.item() != -1, (
f"Donated mamba slot is -1: donate_idx={donate_idx}, "
f"buf={req.mamba_ping_pong_track_buffer.tolist()}, "
f"next_track_idx={req.mamba_next_track_idx}, "
f"rid={req.rid}"
)
self.set_mamba_ping_pong_slot(req, donate_idx, new_slot[0])
return mamba_value_donated
def free_mamba_cache(
self, req: "Req", mamba_ping_pong_track_buffer_to_keep: Optional[int] = None
):
@@ -686,6 +755,12 @@ class HybridReqToTokenPool(ReqToTokenPool):
mamba_ping_pong_track_buffer_to_free = (
mamba_ping_pong_track_buffer_to_free[0:0]
)
if self.enable_mamba_extra_buffer_lazy:
mamba_ping_pong_track_buffer_to_free = (
mamba_ping_pong_track_buffer_to_free[
mamba_ping_pong_track_buffer_to_free != -1
]
)
self.mamba_pool.free(mamba_ping_pong_track_buffer_to_free)
def clear(self):
@@ -52,6 +52,7 @@ class MambaComponent(TreeComponent):
), f"MambaComponent requires page_size=1 when mamba_extra_buffer is disabled, got {cache.page_size}"
super().__init__(cache, params)
self.enable_mamba_extra_buffer = params.enable_mamba_extra_buffer
self.enable_mamba_extra_buffer_lazy = params.enable_mamba_extra_buffer_lazy
# HiCache state
self._mamba_pool_host = None # set to host mamba pool when HiCache enabled
@@ -307,8 +308,8 @@ class MambaComponent(TreeComponent):
if cache_len is None:
cache_len = 0
if self.enable_mamba_extra_buffer:
keep_idx = self.cache.req_to_token_pool.get_mamba_ping_pong_other_idx(
req.mamba_next_track_idx
keep_idx = self.cache.req_to_token_pool.get_mamba_ping_pong_keep_idx(
req
)
mamba_value = (
req.mamba_ping_pong_track_buffer[keep_idx].unsqueeze(-1).clone()
@@ -322,16 +323,12 @@ class MambaComponent(TreeComponent):
return 0
# Donate the mamba index to the radix cache instead of copying.
if self.enable_mamba_extra_buffer:
keep_idx = self.cache.req_to_token_pool.get_mamba_ping_pong_other_idx(
req.mamba_next_track_idx
)
new_slot = self._alloc_mamba_slot()
mamba_value_donated = (
req.mamba_ping_pong_track_buffer[keep_idx].unsqueeze(-1).clone()
self.cache.req_to_token_pool.donate_mamba_ping_pong_slot(
req, new_slot
)
)
req.mamba_ping_pong_track_buffer[keep_idx] = self._alloc_mamba_slot()[0]
self.cache.req_to_token_pool.req_index_to_mamba_ping_pong_track_buffer_mapping[
req.req_pool_idx
] = req.mamba_ping_pong_track_buffer
else:
mamba_value_donated = self._alloc_mamba_slot()
self.cache.req_to_token_pool.mamba_pool.copy_from(
@@ -352,8 +349,8 @@ class MambaComponent(TreeComponent):
insert_result.mamba_exist if insert_result is not None else True
)
if self.enable_mamba_extra_buffer:
keep_idx = self.cache.req_to_token_pool.get_mamba_ping_pong_other_idx(
req.mamba_next_track_idx
keep_idx = self.cache.req_to_token_pool.get_mamba_ping_pong_keep_idx(
req
)
else:
keep_idx = None
@@ -50,6 +50,7 @@ if TYPE_CHECKING:
# the ratio of mamba cache pool size to max_running_requests
MAMBA_CACHE_SIZE_MAX_RUNNING_REQUESTS_RATIO = 3
MAMBA_CACHE_V2_ADDITIONAL_RATIO_OVERLAP = 2
MAMBA_CACHE_V2_ADDITIONAL_RATIO_OVERLAP_LAZY = 1
MAMBA_CACHE_V2_ADDITIONAL_RATIO_NO_OVERLAP = 1
logger = logging.getLogger(__name__)
@@ -233,9 +234,16 @@ class ModelRunnerKVCacheMixin:
additional_ratio = 0
if self.server_args.enable_mamba_extra_buffer():
# ping-pong buffer size is 2 when overlap schedule is on, 1 otherwise.
# Lazy mode saves 1 slot (2 → 1) for overlap; non-overlap already uses 1.
if not self.server_args.disable_overlap_schedule:
additional_ratio = MAMBA_CACHE_V2_ADDITIONAL_RATIO_OVERLAP
if self.server_args.enable_mamba_extra_buffer_lazy():
additional_ratio = MAMBA_CACHE_V2_ADDITIONAL_RATIO_OVERLAP_LAZY
else:
additional_ratio = MAMBA_CACHE_V2_ADDITIONAL_RATIO_OVERLAP
else:
assert (
not self.server_args.enable_mamba_extra_buffer_lazy()
), "Lazy extra buffer requires overlap schedule (--disable-overlap-schedule is incompatible)"
additional_ratio = MAMBA_CACHE_V2_ADDITIONAL_RATIO_NO_OVERLAP
return MAMBA_CACHE_SIZE_MAX_RUNNING_REQUESTS_RATIO + additional_ratio
@@ -352,6 +360,7 @@ class ModelRunnerKVCacheMixin:
]
),
enable_mamba_extra_buffer=self.server_args.enable_mamba_extra_buffer(),
enable_mamba_extra_buffer_lazy=self.server_args.enable_mamba_extra_buffer_lazy(),
speculative_num_draft_tokens=max_spec_draft_tokens,
enable_overlap_schedule=not self.server_args.disable_overlap_schedule,
start_layer=self.start_layer,
+14 -2
View File
@@ -286,7 +286,12 @@ NSA_CHOICES = DSA_CHOICES # deprecated alias
DSA_TOPK_BACKEND_CHOICES = ["sgl-kernel", "torch", "flashinfer"]
MAMBA_SCHEDULER_STRATEGY_CHOICES = ["auto", "no_buffer", "extra_buffer"]
MAMBA_SCHEDULER_STRATEGY_CHOICES = [
"auto",
"no_buffer",
"extra_buffer",
"extra_buffer_lazy",
]
MAMBA_BACKEND_CHOICES = ["triton", "flashinfer"]
@@ -2663,6 +2668,10 @@ class ServerArgs:
is_cuda() or is_musa() or is_npu()
), "Mamba extra_buffer is only supported on CUDA and MUSA and NPU devices with FLA backend"
if self.speculative_num_draft_tokens is not None:
assert not self.enable_mamba_extra_buffer_lazy(), (
"extra_buffer_lazy is not yet supported with speculative decoding. "
"Use --mamba-scheduler-strategy extra_buffer instead."
)
assert (
self.mamba_track_interval >= self.speculative_num_draft_tokens
), f"mamba_track_interval {self.mamba_track_interval} must be greater than or equal to speculative_num_draft_tokens {self.speculative_num_draft_tokens}"
@@ -7240,7 +7249,10 @@ class ServerArgs:
)
def enable_mamba_extra_buffer(self) -> bool:
return self.mamba_scheduler_strategy == "extra_buffer"
return self.mamba_scheduler_strategy in ("extra_buffer", "extra_buffer_lazy")
def enable_mamba_extra_buffer_lazy(self) -> bool:
return self.mamba_scheduler_strategy == "extra_buffer_lazy"
@cached_property
def max_speculative_num_draft_tokens(self) -> Optional[int]:
+8 -6
View File
@@ -253,9 +253,10 @@ def test_input_output_logprobs_match_prefill_cache_hit_helper(
new_input_ids.append(input_ids[i] + result["output_ids"])
output_logprobs.append(_extract_output_logprobs(result))
assert len(new_input_ids) > 0.5 * len(
input_ids
), f"Too few prefill cache hits: {len(new_input_ids)}/{len(input_ids)}"
if not os.environ.get("SGLANG_TEST_SKIP_CACHE_HIT_ASSERT"):
assert len(new_input_ids) > 0.5 * len(
input_ids
), f"Too few prefill cache hits: {len(new_input_ids)}/{len(input_ids)}"
print("Flush Cache and run prefill to get input logprobs ...")
input_logprobs = _get_input_logprobs(base_url, new_input_ids, output_logprobs)
@@ -321,9 +322,10 @@ def test_input_output_logprobs_match_decode_cache_hit_helper(
new_input_ids.append(second_turn_input_ids[i] + result["output_ids"])
output_logprobs.append(_extract_output_logprobs(result))
assert len(new_input_ids) > 0.5 * len(
second_turn_input_ids
), f"Too few decode cache hits: {len(new_input_ids)}/{len(second_turn_input_ids)}"
if not os.environ.get("SGLANG_TEST_SKIP_CACHE_HIT_ASSERT"):
assert len(new_input_ids) > 0.5 * len(
second_turn_input_ids
), f"Too few decode cache hits: {len(new_input_ids)}/{len(second_turn_input_ids)}"
print("Flush Cache and run prefill to get input logprobs ...")
input_logprobs = _get_input_logprobs(base_url, new_input_ids, output_logprobs)