[Qwen3-next] support mamba radix cache for overlap scheduler (#14792)

This commit is contained in:
Hanming Lu
2025-12-14 18:54:16 -08:00
committed by GitHub
parent 36e7c8c59f
commit e61dabf5e4
30 changed files with 1414 additions and 204 deletions
@@ -57,6 +57,7 @@ from sglang.srt.disaggregation.decode_schedule_batch_mixin import (
from sglang.srt.disaggregation.utils import DisaggregationMode
from sglang.srt.distributed.parallel_state import get_tensor_model_parallel_rank
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 (
BaseTokenToKVPoolAllocator,
SWATokenToKVPoolAllocator,
@@ -543,6 +544,14 @@ class Req:
# Memory pool info
self.req_pool_idx: Optional[int] = None
self.mamba_pool_idx: Optional[torch.Tensor] = None # shape (1)
self.mamba_ping_pong_track_buffer: Optional[torch.Tensor] = None # shape (2)
self.mamba_next_track_idx: Optional[int] = None # 0 or 1
self.mamba_last_track_seqlen: Optional[int] = (
None # seq len of the last cached mamba state
)
# the branching point seqlen to track mamba state. If set, given by prefix match,
# it will be the tracked seqlen in the ping pong buffer for the right prefill pass.
self.mamba_branching_seqlen: Optional[int] = None
# Check finish
self.tokenizer = None
@@ -824,11 +833,13 @@ class Req:
self.last_node,
self.last_host_node,
self.host_hit_length,
self.mamba_branching_seqlen,
) = (
match_result.device_indices,
match_result.last_device_node,
match_result.last_host_node,
match_result.host_hit_length,
match_result.mamba_branching_seqlen,
)
self.cache_protected_len = len(self.prefix_indices)
@@ -1027,6 +1038,10 @@ class Req:
self.extend_logprob_start_len = 0
self.is_chunked = 0
self.mamba_pool_idx = None
self.mamba_ping_pong_track_buffer = None
self.mamba_next_track_idx = None
self.mamba_last_track_seqlen = None
self.mamba_branching_seqlen = None
self.already_computed = 0
self.kv_allocated_len = 0
self.kv_committed_len = 0
@@ -1115,6 +1130,11 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
out_cache_loc: torch.Tensor = None # shape: [b], int64
output_ids: torch.Tensor = None # shape: [b], int64
# For hybrid GDN prefix cache
mamba_track_indices: torch.Tensor = None # shape: [b], int64
mamba_track_mask: torch.Tensor = None # shape: [b], bool
mamba_track_seqlens: torch.Tensor = None # shape: [b], int64
# For multimodal inputs
multimodal_inputs: Optional[List] = None
@@ -1380,6 +1400,9 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
input_embeds = []
extend_input_logprob_token_ids = []
multimodal_inputs = []
mamba_track_mask_cpu = []
mamba_track_indices_cpu = []
mamba_track_seqlens_cpu = []
for i, (req, seq_len, pre_len) in enumerate(zip(reqs, seq_lens, prefix_lens)):
req.req_pool_idx = req_pool_indices[i]
@@ -1403,6 +1426,14 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
req.already_computed = seq_len
req.is_retracted = False
if get_global_server_args().enable_mamba_extra_buffer():
self._mamba_radix_cache_v2_req_prepare_for_extend(
req,
mamba_track_mask_cpu,
mamba_track_indices_cpu,
mamba_track_seqlens_cpu,
)
# Compute the relative logprob_start_len in an extend batch
#
# Key variables:
@@ -1512,6 +1543,23 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
self.extend_logprob_start_lens = [r.extend_logprob_start_len for r in reqs]
self.extend_input_logprob_token_ids = extend_input_logprob_token_ids
if get_global_server_args().enable_mamba_extra_buffer():
self.mamba_track_indices = torch.tensor(
mamba_track_indices_cpu,
dtype=torch.int64,
device=self.device,
)
self.mamba_track_mask = torch.tensor(
mamba_track_mask_cpu,
dtype=torch.bool,
device=self.device,
)
self.mamba_track_seqlens = torch.tensor(
mamba_track_seqlens_cpu,
dtype=torch.int64,
device=self.device,
)
if self.model_config.is_encoder_decoder:
self.prepare_encoder_info_extend(input_ids, seq_lens)
@@ -1521,6 +1569,60 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
self.model_config.vocab_size,
)
def _mamba_radix_cache_v2_req_prepare_for_extend(
self,
req: Req,
mamba_track_mask_cpu: List[bool],
mamba_track_indices_cpu: List[int],
mamba_track_seqlens_cpu: List[int],
):
mask = (req.extend_input_len // FLA_CHUNK_SIZE) * FLA_CHUNK_SIZE > 0
mamba_track_mask_cpu.append(mask)
mamba_track_indices_cpu.append(
req.mamba_ping_pong_track_buffer[req.mamba_next_track_idx].item()
)
mamba_track_seqlen = -1
if mask:
# mamba_track_seqlen is used to calculate the indices to track in
# hybrid_linear_attn_backend's _init_track_ssm_indices. Due to the
# fact that the ssm state between aligned and non-aligned are retrieved differently,
# if 1) last pos and 2) is aligned, then retrieved from the last_recurrent_state,
# otherwise retrieved from h (i.e. unaligned).
# We need to pass the non-aligned seqlen to the calculation. Even though
# we pass in mamba_track_seqlen, the actual tracked seqlen is mamba_last_track_seqlen.
mamba_track_seqlen = len(req.prefix_indices) + req.extend_input_len
# mamba_last_track_seqlen is actual tracked seqlen. Used to pass to
# mamba radix cache to track which seqlen this mamba state should store at.
mamba_track_seqlen_aligned = (
len(req.prefix_indices)
+ (req.extend_input_len // FLA_CHUNK_SIZE) * FLA_CHUNK_SIZE
)
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.
branching_seqlen_aligned_mask = (
req.mamba_branching_seqlen - len(req.prefix_indices)
) % FLA_CHUNK_SIZE == 0
if (
req.mamba_branching_seqlen > len(req.prefix_indices)
and req.mamba_branching_seqlen < mamba_track_seqlen
and branching_seqlen_aligned_mask
):
# NOTE: See the comment above for mamba_track_seqlen, the +1 is necessary
# because the branching point is not the last aligned position, so we need
# to retrieve its state from h. Adding 1 will give us the correct index in h,
# otherwise the calculation will retrieve the state from the last_recurrent_state,
# which is not correct.
mamba_track_seqlen = req.mamba_branching_seqlen + 1
mamba_track_seqlen_aligned = req.mamba_branching_seqlen
req.mamba_last_track_seqlen = mamba_track_seqlen_aligned
mamba_track_seqlens_cpu.append(mamba_track_seqlen)
def prepare_for_split_prefill(self):
self.prepare_for_extend()
# For split prefill, we need to set the forward mode to SPLIT_PREFILL
@@ -1786,6 +1888,24 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
self.orig_seq_lens.add_(1)
self.seq_lens_sum += bs
if get_global_server_args().enable_mamba_extra_buffer():
self.mamba_track_indices = torch.tensor(
[
req.mamba_ping_pong_track_buffer[req.mamba_next_track_idx]
for req in self.reqs
],
dtype=torch.int64,
device=self.device,
)
self.mamba_track_mask = torch.tensor(
[
sl % get_global_server_args().mamba_track_interval == 0
for sl in self.seq_lens_cpu
],
dtype=torch.bool,
device=self.device,
)
def maybe_wait_verify_done(self):
if self.is_v2_eagle:
draft_input: EagleDraftInput = self.spec_info
@@ -1842,6 +1962,9 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
self.out_cache_loc = None
self.seq_lens_sum = self.seq_lens.sum().item()
self.output_ids = self.output_ids[keep_indices_device]
self.mamba_track_indices = None
self.mamba_track_mask = None
self.mamba_track_seqlens = None
self.return_logprob = any(req.return_logprob for req in self.reqs)
if self.return_logprob:
self.top_logprobs_nums = [self.top_logprobs_nums[i] for i in keep_indices]
@@ -1889,6 +2012,9 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
self.seq_lens_sum += other.seq_lens_sum
if self.output_ids is not None:
self.output_ids = torch.cat([self.output_ids, other.output_ids])
self.mamba_track_indices = None
self.mamba_track_mask = None
self.mamba_track_seqlens = None
if self.return_logprob and other.return_logprob:
self.top_logprobs_nums.extend(other.top_logprobs_nums)
self.token_ids_logprobs.extend(other.token_ids_logprobs)
@@ -1982,6 +2108,9 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
dllm_config=self.dllm_config,
reqs=self.reqs,
has_grammar=self.has_grammar,
mamba_track_indices=self.mamba_track_indices,
mamba_track_mask=self.mamba_track_mask,
mamba_track_seqlens=self.mamba_track_seqlens,
)
def copy(self):
@@ -2003,6 +2132,9 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
is_prefill_only=self.is_prefill_only,
seq_lens_cpu=self.seq_lens_cpu,
enable_overlap=self.enable_overlap,
mamba_track_indices=self.mamba_track_indices,
mamba_track_mask=self.mamba_track_mask,
mamba_track_seqlens=self.mamba_track_seqlens,
)
def _is_available_size_sufficient(self, num_tokens: int) -> bool:
@@ -2104,3 +2236,8 @@ class ModelWorkerBatch:
# FIXME(lsyin): remove this after fully overlap grammar
reqs: Optional[List[Req]] = None
has_grammar: bool = False
# For mamba state tracking
mamba_track_indices: Optional[torch.Tensor] = None # shape: [b], int64
mamba_track_mask: Optional[torch.Tensor] = None # shape: [b], bool
mamba_track_seqlens: Optional[torch.Tensor] = None # shape: [b], int64
@@ -363,7 +363,7 @@ class PrefillAdder:
self.is_hybrid_swa = isinstance(
self.token_to_kv_pool_allocator, SWATokenToKVPoolAllocator
)
self.is_ssm_radix_cache = isinstance(self.tree_cache, MambaRadixCache)
self.is_hybrid_ssm_cache = isinstance(self.tree_cache, MambaRadixCache)
self.priority_scheduling_preemption_threshold = (
priority_scheduling_preemption_threshold
@@ -389,7 +389,7 @@ class PrefillAdder:
self.token_to_kv_pool_allocator.swa_available_size()
+ self.tree_cache.swa_evictable_size(),
)
elif self.is_ssm_radix_cache:
elif self.is_hybrid_ssm_cache:
available_and_evictable = (
self.token_to_kv_pool_allocator.available_size()
+ self.tree_cache.full_evictable_size()
@@ -411,7 +411,7 @@ class PrefillAdder:
self.token_to_kv_pool_allocator.swa_available_size()
+ self.tree_cache.swa_evictable_size(),
)
elif self.is_ssm_radix_cache:
elif self.is_hybrid_ssm_cache:
available_and_evictable = (
self.token_to_kv_pool_allocator.available_size()
+ self.tree_cache.full_evictable_size()
+3 -2
View File
@@ -405,7 +405,7 @@ class Scheduler(
# Hybrid memory pool
self.is_hybrid_swa = self.tp_worker.is_hybrid_swa
self.is_ssm_model = (
self.is_hybrid_ssm = (
self.tp_worker.model_runner.hybrid_gdn_config is not None
or self.tp_worker.model_runner.mamba2_config is not None
)
@@ -772,6 +772,7 @@ class Scheduler(
eviction_policy=server_args.radix_eviction_policy,
enable_metrics=self.enable_metrics,
enable_kv_cache_events=self.enable_kv_cache_events,
enable_mamba_extra_buffer=server_args.enable_mamba_extra_buffer(),
)
if (
@@ -808,7 +809,7 @@ class Scheduler(
self.tree_cache = SWARadixCache(
params=params, sliding_window_size=self.sliding_window_size
)
elif self.is_ssm_model:
elif self.is_hybrid_ssm:
from sglang.srt.mem_cache.mamba_radix_cache import MambaRadixCache
self.tree_cache = MambaRadixCache(params)
@@ -112,7 +112,7 @@ class SchedulerMetricsMixin:
f"full token usage: {full_token_usage:.2f}, "
f"swa token usage: {swa_token_usage:.2f}, "
)
elif self.is_ssm_model:
elif self.is_hybrid_ssm:
(
full_num_used,
_,
@@ -166,7 +166,7 @@ class SchedulerMetricsMixin:
self.stats.token_usage = token_usage
if self.is_hybrid_swa:
self.stats.swa_token_usage = swa_token_usage
if self.is_ssm_model:
if self.is_hybrid_ssm:
self.stats.mamba_usage = mamba_usage
self.stats.num_queue_reqs = len(self.waiting_queue)
self.stats.num_grammar_queue_reqs = len(self.grammar_queue)
@@ -238,7 +238,7 @@ class SchedulerMetricsMixin:
f"#swa token: {swa_num_used}, "
f"swa token usage: {swa_token_usage:.2f}, "
)
elif self.is_ssm_model:
elif self.is_hybrid_ssm:
(
full_num_used,
mamba_used,
@@ -315,7 +315,7 @@ class SchedulerMetricsMixin:
self.stats.token_usage = token_usage
if self.is_hybrid_swa:
self.stats.swa_token_usage = swa_token_usage
if self.is_ssm_model:
if self.is_hybrid_ssm:
self.stats.mamba_usage = mamba_usage
self.stats.decode_sum_seq_lens = batch.seq_lens_cpu.sum().item()
self.stats.gen_throughput = self.last_gen_throughput
@@ -402,7 +402,7 @@ class SchedulerMetricsMixin:
if self.is_hybrid_swa:
full_num_used, swa_num_used, *_ = self._get_swa_token_info()
num_tokens = max(full_num_used, swa_num_used)
elif self.is_ssm_model:
elif self.is_hybrid_ssm:
num_tokens = self._get_mamba_token_info()[0]
else:
num_tokens = self._get_token_info()[0]
@@ -21,6 +21,7 @@ from sglang.srt.managers.schedule_batch import (
ScheduleBatch,
)
from sglang.srt.mem_cache.common import release_kv_cache
from sglang.srt.server_args import get_global_server_args
from sglang.srt.tracing.trace import trace_slice, trace_slice_batch, trace_slice_end
if TYPE_CHECKING:
@@ -267,6 +268,7 @@ class SchedulerOutputProcessorMixin:
next_token_ids = result.next_token_ids.tolist()
accept_lens = result.accept_lens.tolist()
result.num_accepted_tokens = sum(accept_lens) - len(batch.reqs)
result.accept_length_per_req_cpu = [x - 1 for x in accept_lens]
predict_tokens = []
stride = self.draft_worker.speculative_num_draft_tokens
@@ -359,6 +361,9 @@ class SchedulerOutputProcessorMixin:
req.output_ids.extend(next_token_id)
new_accepted_len = len(next_token_id)
# Update Mamba last track seqlen
self._mamba_prefix_cache_update(req, batch, result, i)
req.check_finished(new_accepted_len)
if req.finished():
@@ -424,6 +429,31 @@ class SchedulerOutputProcessorMixin:
):
self.log_decode_stats(can_run_cuda_graph, running_batch=batch)
def _mamba_prefix_cache_update(
self, req: Req, batch: ScheduleBatch, 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 = 1 - req.mamba_next_track_idx
req.mamba_last_track_seqlen = seq_len
elif (
not batch.spec_algorithm.is_none()
and result.accept_length_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.accept_length_per_req_cpu[i])
// mamba_track_interval
):
req.mamba_last_track_seqlen = (
actual_seq_len // mamba_track_interval * mamba_track_interval
)
def _process_input_token_logprobs(
self, req: Req, input_token_logprobs: List
) -> None:
@@ -121,9 +121,22 @@ class SchedulerRuntimeCheckerMixin:
full_num_used != self.tree_cache.full_protected_size()
or mamba_num_used != self.tree_cache.mamba_protected_size()
)
free_full_pages = set(
self.token_to_kv_pool_allocator.free_pages.tolist()
+ self.token_to_kv_pool_allocator.release_pages.tolist()
)
cached_full_pages = set(self.tree_cache.all_values_flatten().tolist())
expected_full_pages = set(range(1, self.token_to_kv_pool_allocator.size + 1))
leaked_full_pages = expected_full_pages - free_full_pages - cached_full_pages
free_mamba_pages = set(self.req_to_token_pool.mamba_pool.free_slots.tolist())
cached_mamba_pages = set(self.tree_cache.all_mamba_values_flatten().tolist())
expected_mamba_pages = set(range(self.req_to_token_pool.mamba_pool.size))
leaked_mamba_pages = (
expected_mamba_pages - free_mamba_pages - cached_mamba_pages
)
token_msg = (
f"{full_available_size=}, {full_evictable_size=}, {self.token_to_kv_pool_allocator.size=}, {self.tree_cache.full_protected_size()=}\n"
f"{mamba_available_size=}, {mamba_evictable_size=}, {self.req_to_token_pool.mamba_pool.size=}, {self.tree_cache.mamba_protected_size()=}\n"
f"{mamba_available_size=}, {mamba_evictable_size=}, {self.req_to_token_pool.mamba_pool.size=}, {self.tree_cache.mamba_protected_size()=}, leaked_full_pages={leaked_full_pages if len(leaked_full_pages) > 0 else None}, leaked_mamba_pages={leaked_mamba_pages if len(leaked_mamba_pages) > 0 else None}\n"
)
return memory_leak, token_msg
@@ -207,7 +220,7 @@ class SchedulerRuntimeCheckerMixin:
def check_memory(self: Scheduler):
if self.is_hybrid_swa:
memory_leak, token_msg = self._check_hybrid_memory()
elif self.is_ssm_model and isinstance(self.tree_cache, MambaRadixCache):
elif self.is_hybrid_ssm and isinstance(self.tree_cache, MambaRadixCache):
memory_leak, token_msg = self._check_mamba_memory()
else:
memory_leak, token_msg = self._check_radix_cache_memory()
@@ -242,7 +255,7 @@ class SchedulerRuntimeCheckerMixin:
) = self._get_swa_token_info()
num_used = max(full_num_used, swa_num_used)
token_usage = max(full_token_usage, swa_token_usage)
elif self.is_ssm_model:
elif self.is_hybrid_ssm:
(
num_used,
_,
@@ -281,7 +294,7 @@ class SchedulerRuntimeCheckerMixin:
def check_tree_cache(self: Scheduler):
if (self.is_hybrid_swa and isinstance(self.tree_cache, SWARadixCache)) or (
self.is_ssm_model and isinstance(self.tree_cache, MambaRadixCache)
self.is_hybrid_ssm and isinstance(self.tree_cache, MambaRadixCache)
):
self.tree_cache.sanity_check()
@@ -344,7 +357,7 @@ class SchedulerWatchdog:
# Print batch size and memory pool info to check whether there are de-sync issues.
if self.scheduler.is_hybrid_swa:
_, info_msg = self.scheduler._check_hybrid_memory()
elif self.scheduler.is_ssm_model and isinstance(
elif self.scheduler.is_hybrid_ssm and isinstance(
self.scheduler.tree_cache, MambaRadixCache
):
_, info_msg = self.scheduler._check_mamba_memory()
+2 -1
View File
@@ -24,7 +24,8 @@ class GenerationBatchResult:
logits_output: Optional[LogitsProcessorOutput] = None
pp_hidden_states_proxy_tensors: Optional[PPProxyTensors] = None
next_token_ids: Optional[torch.Tensor] = None
num_accepted_tokens: Optional[int] = None
num_accepted_tokens: int = 0
accept_length_per_req_cpu: Optional[List[int]] = None
can_run_cuda_graph: bool = False
# For output processing