[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
@@ -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: