[Qwen3-next] support mamba radix cache for overlap scheduler (#14792)
This commit is contained in:
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user