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