[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
+9 -1
View File
@@ -155,8 +155,16 @@ class HybridMambaDecodeReqToTokenPool(HybridReqToTokenPool):
pre_alloc_size=pre_alloc_size,
)
self.enable_memory_saver = enable_memory_saver
self.enable_mamba_extra_buffer = (
False # TODO: add PD support for mamba cache extra_buffer
)
self._init_mamba_pool(
size + pre_alloc_size, cache_params, device, speculative_num_draft_tokens
size=size + pre_alloc_size,
mamba_spec_state_size=size + pre_alloc_size,
cache_params=cache_params,
device=device,
enable_mamba_extra_buffer=self.enable_mamba_extra_buffer,
speculative_num_draft_tokens=speculative_num_draft_tokens,
)
def clear(self):
@@ -67,7 +67,7 @@ def chunk_gated_delta_rule_fwd(
cu_seqlens=cu_seqlens,
)
if SUPPRESS_LEVEL < 3:
return g, o, A, final_state, None, None, None
return g, o, A, final_state, None, h, None
elif SUPPRESS_LEVEL >= 3:
return g, o, A, final_state, w, h, v_new
@@ -108,7 +108,7 @@ class ChunkGatedDeltaRuleFunction(torch.autograd.Function):
output_final_state=output_final_state,
cu_seqlens=cu_seqlens,
)
return o.to(q.dtype), final_state
return o.to(q.dtype), final_state, h
@torch.compiler.disable
@@ -224,7 +224,7 @@ def chunk_gated_delta_rule(
)
if scale is None:
scale = k.shape[-1] ** -0.5
o, final_state = ChunkGatedDeltaRuleFunction.apply(
o, final_state, h = ChunkGatedDeltaRuleFunction.apply(
q,
k,
v,
@@ -238,4 +238,4 @@ def chunk_gated_delta_rule(
)
if head_first:
o = rearrange(o, "b t h ... -> b h t ...")
return o, final_state
return o, final_state, h
@@ -16,6 +16,7 @@ from sglang.srt.layers.attention.fla.op import exp, safe_exp
from sglang.srt.layers.attention.fla.utils import is_nvidia_hopper
NUM_WARPS = [2, 4] if is_nvidia_hopper else [2, 4, 8, 16]
CHUNK_SIZE = 64
# @triton.autotune(
@@ -274,16 +275,15 @@ def chunk_gated_delta_rule_fwd_h(
gk: Optional[torch.Tensor] = None,
initial_state: Optional[torch.Tensor] = None,
output_final_state: bool = False,
chunk_size: int = 64, # SY: remove this argument and force chunk size 64?
save_new_value: bool = True,
cu_seqlens: Optional[torch.LongTensor] = None,
) -> Tuple[torch.Tensor, torch.Tensor]:
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
B, T, Hg, K, V = *k.shape, u.shape[-1]
H = u.shape[-2]
BT = chunk_size
BT = CHUNK_SIZE
chunk_indices = (
prepare_chunk_indices(cu_seqlens, chunk_size)
prepare_chunk_indices(cu_seqlens, CHUNK_SIZE)
if cu_seqlens is not None
else None
)
@@ -353,6 +353,7 @@ def fused_recurrent_gated_delta_rule_update_fwd_kernel(
cu_seqlens,
scale,
intermediate_states_buffer,
intermediate_state_indices,
cache_steps,
retrieve_parent_token_ptr,
stride_retrieve_parent_token_seq: tl.constexpr,
@@ -431,7 +432,7 @@ def fused_recurrent_gated_delta_rule_update_fwd_kernel(
# Prepare intermediate state cache variables if enabled
cache_idx = -1
if CACHE_INTERMEDIATE_STATES:
cache_idx = tl.load(h0_indices + i_n)
cache_idx = tl.load(intermediate_state_indices + i_n)
step_idx = 0
for _ in range(0, T):
@@ -532,6 +533,7 @@ def fused_recurrent_gated_delta_rule_update_fwd(
disable_state_update: bool = False,
disable_output_calculation: bool = False,
intermediate_states_buffer: Optional[torch.Tensor] = None,
intermediate_state_indices: Optional[torch.Tensor] = None,
cache_steps: Optional[int] = None,
retrieve_parent_token: Optional[torch.Tensor] = None,
) -> torch.Tensor:
@@ -574,6 +576,7 @@ def fused_recurrent_gated_delta_rule_update_fwd(
cu_seqlens=cu_seqlens,
scale=scale,
intermediate_states_buffer=intermediate_states_buffer,
intermediate_state_indices=intermediate_state_indices,
cache_steps=0 if cache_steps is None else cache_steps,
retrieve_parent_token_ptr=retrieve_parent_token,
stride_retrieve_parent_token_seq=stride_retrieve_parent_token_seq,
@@ -588,13 +591,13 @@ def fused_recurrent_gated_delta_rule_update_fwd(
BK=BK,
BV=BV,
USE_INITIAL_STATE=initial_state_source is not None,
IS_BETA_HEADWISE=beta.ndim == v.ndim,
USE_QK_L2NORM_IN_KERNEL=use_qk_l2norm_in_kernel,
IS_VARLEN=cu_seqlens is not None,
DISABLE_STATE_UPDATE=disable_state_update,
DISABLE_OUTPUT_CALCULATION=disable_output_calculation,
CACHE_INTERMEDIATE_STATES=intermediate_states_buffer is not None,
HAS_EAGLE_TREE_CUSTOM_ATTN_MASK=retrieve_parent_token is not None,
IS_BETA_HEADWISE=beta.ndim == v.ndim,
USE_QK_L2NORM_IN_KERNEL=use_qk_l2norm_in_kernel,
DISABLE_STATE_UPDATE=disable_state_update,
DISABLE_OUTPUT_CALCULATION=disable_output_calculation,
num_warps=num_warps,
num_stages=num_stages,
)
@@ -621,6 +624,7 @@ class FusedRecurrentUpdateFunction(torch.autograd.Function):
disable_state_update: bool = False,
disable_output_calculation: bool = False,
intermediate_states_buffer: Optional[torch.Tensor] = None,
intermediate_state_indices: Optional[torch.Tensor] = None,
cache_steps: Optional[int] = None,
retrieve_parent_token: Optional[torch.Tensor] = None,
):
@@ -638,6 +642,7 @@ class FusedRecurrentUpdateFunction(torch.autograd.Function):
disable_state_update=disable_state_update,
disable_output_calculation=disable_output_calculation,
intermediate_states_buffer=intermediate_states_buffer,
intermediate_state_indices=intermediate_state_indices,
cache_steps=cache_steps,
retrieve_parent_token=retrieve_parent_token,
)
@@ -668,6 +673,7 @@ def fused_recurrent_gated_delta_rule_update(
disable_state_update: bool = False,
disable_output_calculation: bool = False,
intermediate_states_buffer: Optional[torch.Tensor] = None,
intermediate_state_indices: Optional[torch.Tensor] = None,
cache_steps: Optional[int] = None,
retrieve_parent_token: Optional[torch.Tensor] = None,
) -> torch.Tensor:
@@ -677,14 +683,17 @@ def fused_recurrent_gated_delta_rule_update(
f"The batch size is expected to be 1 rather than {q.shape[0]} when using `cu_seqlens`."
f"Please flatten variable-length inputs before processing."
)
if (
initial_state_source is not None
and initial_state_indices.shape[0] != len(cu_seqlens) - 1
):
raise ValueError(
f"The number of initial states is expected to be equal to the number of input sequences, "
f"i.e., {len(cu_seqlens) - 1} rather than {initial_state_indices.shape[0]}."
)
if initial_state_source is not None:
if initial_state_indices.shape[0] != len(cu_seqlens) - 1:
raise ValueError(
f"The number of initial states is expected to be equal to the number of input sequences, "
f"i.e., {len(cu_seqlens) - 1} rather than {initial_state_indices.shape[0]}."
)
if initial_state_indices.shape[0] != intermediate_state_indices.shape[0]:
raise ValueError(
f"The number of intermediate state indices is expected to be equal to the number of input sequences, "
f"i.e., {initial_state_indices.shape[0]} != {intermediate_state_indices.shape[0]}."
)
if scale is None:
scale = k.shape[-1] ** -0.5
else:
@@ -705,6 +714,7 @@ def fused_recurrent_gated_delta_rule_update(
disable_state_update,
disable_output_calculation,
intermediate_states_buffer,
intermediate_state_indices,
cache_steps,
retrieve_parent_token,
)
@@ -1,10 +1,13 @@
from typing import Optional, Union
import torch
import triton
import triton.language as tl
from einops import rearrange
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
from sglang.srt.layers.attention.fla.chunk import chunk_gated_delta_rule
from sglang.srt.layers.attention.fla.chunk_delta_h import CHUNK_SIZE as FLA_CHUNK_SIZE
from sglang.srt.layers.attention.fla.fused_gdn_gating import fused_gdn_gating
from sglang.srt.layers.attention.fla.fused_recurrent import (
fused_recurrent_gated_delta_rule_update,
@@ -57,6 +60,108 @@ elif is_npu():
causal_conv1d_update = causal_conv1d_update_npu
# Kernel to track mamba states if needed based on track mask
@triton.jit
def track_mamba_state_if_needed_kernel(
conv_states_ptr,
ssm_states_ptr,
cache_indices_ptr,
mamba_track_mask_ptr,
mamba_track_indices_ptr,
conv_state_stride_0, # stride for first dimension (batch/pool index)
ssm_state_stride_0, # stride for first dimension (batch/pool index)
conv_state_numel_per_row: tl.constexpr, # total elements per row
ssm_state_numel_per_row: tl.constexpr, # total elements per row
BLOCK_SIZE: tl.constexpr,
):
"""
Track conv_states and ssm_states rows based on track mask.
This kernel replaces a Python loop that copies state tensors for mamba attention.
For each batch element, if the track mask is True, it copies the entire row from
the source index (cache_indices[i]) to the destination index (mamba_track_indices[i]).
Grid: (batch_size,)
Each block handles one batch element, using multiple threads to copy data in parallel.
"""
batch_idx = tl.program_id(0)
# Load the copy mask for this batch element
track_mask = tl.load(mamba_track_mask_ptr + batch_idx)
# Early exit if we don't need to track
if not track_mask:
return
# Load source and destination indices
src_idx = tl.load(cache_indices_ptr + batch_idx)
dst_idx = tl.load(mamba_track_indices_ptr + batch_idx)
# Copy conv_states
# Each thread handles BLOCK_SIZE elements
for offset in range(0, conv_state_numel_per_row, BLOCK_SIZE):
element_indices = offset + tl.arange(0, BLOCK_SIZE)
mask = element_indices < conv_state_numel_per_row
src_ptr = conv_states_ptr + src_idx * conv_state_stride_0 + element_indices
dst_ptr = conv_states_ptr + dst_idx * conv_state_stride_0 + element_indices
data = tl.load(src_ptr, mask=mask, other=0.0)
tl.store(dst_ptr, data, mask=mask)
# Copy ssm_states
for offset in range(0, ssm_state_numel_per_row, BLOCK_SIZE):
element_indices = offset + tl.arange(0, BLOCK_SIZE)
mask = element_indices < ssm_state_numel_per_row
src_ptr = ssm_states_ptr + src_idx * ssm_state_stride_0 + element_indices
dst_ptr = ssm_states_ptr + dst_idx * ssm_state_stride_0 + element_indices
data = tl.load(src_ptr, mask=mask, other=0.0)
tl.store(dst_ptr, data, mask=mask)
def track_mamba_states_if_needed(
conv_states: torch.Tensor,
ssm_states: torch.Tensor,
cache_indices: torch.Tensor,
mamba_track_mask: torch.Tensor,
mamba_track_indices: torch.Tensor,
batch_size: int,
):
"""
Track mamba states using Triton kernel for better performance.
Args:
conv_states: Convolution states tensor [pool_size, ...]
ssm_states: SSM states tensor [pool_size, ...]
cache_indices: Source indices for each batch element [batch_size]
mamba_track_mask: Boolean mask indicating which elements to track [batch_size]
mamba_track_indices: Indices to track for each batch element [batch_size]
batch_size: Number of batch elements
"""
conv_state_numel_per_row = conv_states[0].numel()
ssm_state_numel_per_row = ssm_states[0].numel()
# Choose BLOCK_SIZE based on the size of the data
BLOCK_SIZE = 1024
# Launch kernel with batch_size blocks
grid = (batch_size,)
track_mamba_state_if_needed_kernel[grid](
conv_states,
ssm_states,
cache_indices,
mamba_track_mask,
mamba_track_indices,
conv_states.stride(0),
ssm_states.stride(0),
conv_state_numel_per_row,
ssm_state_numel_per_row,
BLOCK_SIZE,
)
class MambaAttnBackendBase(AttentionBackend):
def __init__(self, model_runner: ModelRunner):
super().__init__()
@@ -71,6 +176,7 @@ class MambaAttnBackendBase(AttentionBackend):
self.retrieve_parent_token_list = []
self.cached_cuda_graph_decode_query_start_loc: torch.Tensor = None
self.cached_cuda_graph_verify_query_start_loc: torch.Tensor = None
self.conv_states_shape: tuple[int, int] = None
def _forward_metadata(self, forward_batch: ForwardBatch):
bs = forward_batch.batch_size
@@ -78,6 +184,15 @@ class MambaAttnBackendBase(AttentionBackend):
retrieve_next_token = None
retrieve_next_sibling = None
retrieve_parent_token = None
track_conv_indices = None
track_ssm_h_src = None
track_ssm_h_dst = None
track_ssm_final_src = None
track_ssm_final_dst = None
mamba_cache_indices = self.req_to_token_pool.get_mamba_indices(
forward_batch.req_pool_indices
)
if forward_batch.forward_mode.is_decode_or_idle():
query_start_loc = torch.arange(
@@ -108,22 +223,160 @@ class MambaAttnBackendBase(AttentionBackend):
forward_batch.extend_start_loc[-1]
+ forward_batch.extend_seq_lens[-1]
)
if (
forward_batch.mamba_track_mask is not None
and forward_batch.mamba_track_mask.any()
):
track_conv_indices = self._init_track_conv_indices(
query_start_loc, forward_batch
)
(
track_ssm_h_src,
track_ssm_h_dst,
track_ssm_final_src,
track_ssm_final_dst,
) = self._init_track_ssm_indices(mamba_cache_indices, forward_batch)
else:
raise ValueError(f"Invalid forward mode: {forward_batch.forward_mode=}")
mamba_cache_indices = self.req_to_token_pool.get_mamba_indices(
forward_batch.req_pool_indices
)
return ForwardMetadata(
query_start_loc=query_start_loc,
mamba_cache_indices=mamba_cache_indices,
retrieve_next_token=retrieve_next_token,
retrieve_next_sibling=retrieve_next_sibling,
retrieve_parent_token=retrieve_parent_token,
track_conv_indices=track_conv_indices,
track_ssm_h_src=track_ssm_h_src,
track_ssm_h_dst=track_ssm_h_dst,
track_ssm_final_src=track_ssm_final_src,
track_ssm_final_dst=track_ssm_final_dst,
)
def init_forward_metadata(self, forward_batch: ForwardBatch):
self.forward_metadata = self._forward_metadata(forward_batch)
def _init_track_conv_indices(
self, query_start_loc: torch.Tensor, forward_batch: ForwardBatch
):
"""
Compute indices for extracting conv states from the input sequence during extend.
In Mamba models, the conv layer maintains a sliding window of recent inputs.
After processing a prefill chunk, we need to save the last `conv_state_len` tokens
of the processed region for prefix caching.
The key insight is that FLA (Flash Linear Attention) processes sequences in chunks
of FLA_CHUNK_SIZE. We only track the conv state up to the last complete chunk boundary
(aligned_len).
start_indices is the starting token index of the conv state to track in this extend batch.
indices include all pos to track in this extend batch, conv_state_len for each req that
needs to be tracked (i.e. mamba_track_mask is True)
Returns:
indices: Tensor of shape [num_tracked_requests, conv_state_len] containing
flattened positions into the packed input tensor.
"""
conv_state_len = self.conv_states_shape[-1]
# Calculate the end position of the last aligned chunk
lens_to_track = (
forward_batch.mamba_track_seqlens - forward_batch.extend_prefix_lens
)
aligned_len = (lens_to_track // FLA_CHUNK_SIZE) * FLA_CHUNK_SIZE
start_indices = query_start_loc[:-1] + aligned_len - conv_state_len
start_indices = start_indices[forward_batch.mamba_track_mask]
# Create indices: [batch_size, conv_state_len]
indices = start_indices.unsqueeze(-1) + torch.arange(
conv_state_len,
device=self.device,
dtype=start_indices.dtype,
)
return indices.clamp(0, query_start_loc[-1] - 1)
def _init_track_ssm_indices(
self, mamba_cache_indices: torch.Tensor, forward_batch: ForwardBatch
):
"""
Compute source and destination indices for tracking SSM states for prefix caching.
After processing a prefill, we need to save the SSM recurrent state for prefix caching.
The FLA kernel outputs intermediate hidden states `h` at each chunk boundary,
plus a `last_recurrent_state` at the end of the chunked prefill size.
The challenge is that sequences may or may not end on a chunk boundary:
- Aligned case (len % FLA_CHUNK_SIZE == 0): In this case, FLA will store the to-cache
state in the last_recurrent_state.
- Unaligned case (len % FLA_CHUNK_SIZE != 0): The last_recurrent_state includes the
unaligned position, but we only want state up to the last chunk boundary.
We must extract from the intermediate `h` tensor at the appropriate chunk index.
We compute the src and dst indices for all requests that need to be cached
(i.e. mamba_track_mask is True) based on the rule above.
For example:
1. If chunked prefill length is < 64, then only final state has value. In this case we
cache `final` state.
2. if chunked prefill length == 64, then only final state has value. In this case we
cache pos 64, from `final` state
3. if chunked prefill length >64 and < 128, then both h and final state have value.
We cache pos 64 from `h` state
4. if chunked prefill length ==128, then both h and final state have value. We cache
pos 128 from `final` state. Note `h` doesn't include the pos 128.
Returns:
track_ssm_h_src: Source indices into the packed `h` tensor (for unaligned seqs)
track_ssm_h_dst: Destination cache slot indices (for unaligned seqs)
track_ssm_final_src: Source indices into last_recurrent_state buffer (for aligned seqs)
track_ssm_final_dst: Destination cache slot indices (for aligned seqs)
"""
# Move to CPU to avoid kernel launches for masking operations
mamba_track_mask = forward_batch.mamba_track_mask.cpu()
extend_seq_lens = forward_batch.extend_seq_lens.cpu()
mamba_track_indices = forward_batch.mamba_track_indices.cpu()
mamba_cache_indices = mamba_cache_indices.cpu()
mamba_track_seqlens = forward_batch.mamba_track_seqlens.cpu()
prefix_lens = forward_batch.extend_prefix_lens.cpu()
# Calculate the number of hidden states per request
num_h_states = (extend_seq_lens - 1) // FLA_CHUNK_SIZE + 1
# Calculate the starting offset for each sequence in the packed batch
track_ssm_src_offset = torch.zeros_like(num_h_states)
track_ssm_src_offset[1:] = torch.cumsum(num_h_states[:-1], dim=0)
# Filter variables by track mask
lens_to_track = mamba_track_seqlens - prefix_lens
lens_masked = lens_to_track[mamba_track_mask]
offset_masked = track_ssm_src_offset[mamba_track_mask]
dst_masked = mamba_track_indices[mamba_track_mask]
# Determine if the sequence ends at a chunk boundary
is_aligned = (lens_masked % FLA_CHUNK_SIZE) == 0
# Case 1: Aligned. Use last_recurrent_state from ssm_states.
track_ssm_final_src = mamba_cache_indices[mamba_track_mask][is_aligned]
track_ssm_final_dst = dst_masked[is_aligned]
# Case 2: Unaligned. Use intermediate state from h.
# TODO: if support FLA_CHUNK_SIZE % page size != 0, then need to modify this
not_aligned = ~is_aligned
track_ssm_h_src = offset_masked[not_aligned] + (
lens_masked[not_aligned] // FLA_CHUNK_SIZE
)
track_ssm_h_dst = dst_masked[not_aligned]
# Move back to GPU
return (
track_ssm_h_src.to(self.device, non_blocking=True),
track_ssm_h_dst.to(self.device, non_blocking=True),
track_ssm_final_src.to(self.device, non_blocking=True),
track_ssm_final_dst.to(self.device, non_blocking=True),
)
def init_forward_metadata_capture_cuda_graph(
self,
bs: int,
@@ -165,7 +418,7 @@ class MambaAttnBackendBase(AttentionBackend):
)
)
self.query_start_loc_list.append(
torch.empty((i + 2,), dtype=torch.int32, device=self.device)
torch.zeros((i + 2,), dtype=torch.int32, device=self.device)
)
self.retrieve_next_token_list.append(
torch.zeros(
@@ -277,7 +530,6 @@ class MambaAttnBackendBase(AttentionBackend):
# If topk > 1, we need to use retrieve_next_token and retrieve_next_sibling to handle the eagle tree custom attention mask
if forward_mode.is_target_verify() and spec_info.topk > 1:
bs_without_pad = spec_info.retrive_next_token.shape[0]
# print(spec_info.retrive_next_token, spec_info.retrive_next_sibling)
self.retrieve_next_token_list[bs - 1][:bs_without_pad].copy_(
spec_info.retrive_next_token
)
@@ -300,6 +552,70 @@ class MambaAttnBackendBase(AttentionBackend):
def get_cuda_graph_seq_len_fill_value(self):
return 1 # Mamba attn does not use seq lens to index kv cache
def _track_mamba_state_decode(
self,
forward_batch: ForwardBatch,
conv_states: torch.Tensor,
ssm_states: torch.Tensor,
cache_indices: torch.Tensor,
):
"""
Track and copy Mamba conv/SSM states during decode for prefix caching.
During decode, each token update modifies conv_states and ssm_states in-place
at positions indexed by cache_indices (the working slots). For prefix caching,
we need to copy these updated states to persistent cache slots (mamba_track_indices)
so they can be prefix cached.
This delegates to `track_mamba_states_if_needed`, which performs:
conv_states[mamba_track_indices[i]] = conv_states[cache_indices[i]]
ssm_states[mamba_track_indices[i]] = ssm_states[cache_indices[i]]
for all requests where mamba_track_mask[i] is True.
"""
if forward_batch.mamba_track_mask is not None:
track_mamba_states_if_needed(
conv_states,
ssm_states,
cache_indices,
forward_batch.mamba_track_mask,
forward_batch.mamba_track_indices,
forward_batch.batch_size,
)
def _track_mamba_state_extend(
self,
forward_batch: ForwardBatch,
h: torch.Tensor,
ssm_states: torch.Tensor,
forward_metadata: ForwardMetadata,
):
"""
Track and copy SSM states during extend for prefix caching.
After the FLA chunked prefill kernel runs, we need to save the SSM recurrent
state at the last chunk boundary so it can be reused for prefix caching.
The source of the state depends on whether the sequence length is aligned
to FLA_CHUNK_SIZE. See `_init_track_ssm_indices` for more details on how
the source and destination indices are computed.
Note: Conv state tracking for extend is handled separately via gather operations
using indices computed by `_init_track_conv_indices`.
"""
if (
forward_batch.mamba_track_mask is not None
and forward_batch.mamba_track_mask.any()
):
h = h.squeeze(0)
if forward_metadata.track_ssm_h_src.numel() > 0:
ssm_states[forward_metadata.track_ssm_h_dst] = h[
forward_metadata.track_ssm_h_src
].to(ssm_states.dtype, copy=False)
if forward_metadata.track_ssm_final_src.numel() > 0:
ssm_states[forward_metadata.track_ssm_final_dst] = ssm_states[
forward_metadata.track_ssm_final_src
]
class KimiLinearAttnBackend(MambaAttnBackendBase):
"""Attention backend using Mamba kernel."""
@@ -521,6 +837,15 @@ class KimiLinearAttnBackend(MambaAttnBackendBase):
class GDNAttnBackend(MambaAttnBackendBase):
"""Attention backend using Mamba kernel."""
def __init__(self, model_runner: ModelRunner):
super().__init__(model_runner)
self.conv_states_shape = (
model_runner.req_to_token_pool.mamba_pool.mamba_cache.conv[0].shape
)
assert (
self.conv_states_shape[-1] < FLA_CHUNK_SIZE
), f"{self.conv_states_shape[-1]=} should be less than {FLA_CHUNK_SIZE}"
def forward_decode(
self,
q: torch.Tensor,
@@ -593,6 +918,10 @@ class GDNAttnBackend(MambaAttnBackendBase):
softplus_threshold=20.0,
)
self._track_mamba_state_decode(
forward_batch, conv_states, ssm_states, cache_indices
)
return core_attn_out
def forward_extend(
@@ -622,12 +951,13 @@ class GDNAttnBackend(MambaAttnBackendBase):
seq_len = kwargs["seq_len"]
is_target_verify = forward_batch.forward_mode.is_target_verify()
forward_metadata = self.forward_metadata
query_start_loc = self.forward_metadata.query_start_loc
cache_indices = self.forward_metadata.mamba_cache_indices
retrieve_next_token = self.forward_metadata.retrieve_next_token
retrieve_next_sibling = self.forward_metadata.retrieve_next_sibling
retrieve_parent_token = self.forward_metadata.retrieve_parent_token
query_start_loc = forward_metadata.query_start_loc
cache_indices = forward_metadata.mamba_cache_indices
retrieve_next_token = forward_metadata.retrieve_next_token
retrieve_next_sibling = forward_metadata.retrieve_next_sibling
retrieve_parent_token = forward_metadata.retrieve_parent_token
mamba_cache_params = self.req_to_token_pool.mamba2_layer_cache(layer_id)
conv_states = mamba_cache_params.conv[0]
@@ -643,6 +973,9 @@ class GDNAttnBackend(MambaAttnBackendBase):
dtype=torch.bool,
device=forward_batch.input_ids.device,
)
intermediate_state_indices = torch.arange(
cache_indices.shape[0], dtype=torch.int32, device=cache_indices.device
)
else:
has_initial_states = forward_batch.extend_prefix_lens > 0
@@ -660,14 +993,30 @@ class GDNAttnBackend(MambaAttnBackendBase):
activation,
conv_state_indices=cache_indices[:batch_size],
intermediate_conv_window=intermediate_conv_window_cache,
intermediate_state_indices=intermediate_state_indices[:batch_size],
retrieve_next_token=retrieve_next_token,
retrieve_next_sibling=retrieve_next_sibling,
retrieve_parent_token=retrieve_parent_token,
)
mixed_qkv = mixed_qkv_processed.transpose(1, 2).view(seq_len, -1)
else:
mixed_qkv = mixed_qkv.transpose(0, 1)
if (
forward_batch.mamba_track_mask is not None
and forward_batch.mamba_track_mask.any()
):
conv_dst = forward_batch.mamba_track_indices
# Gather all slices at once: [:, track_conv_indices] -> [d, num_masked, slice_len]
# track_conv_indices is already filtered and clamped in _init_track_conv_indices
mixed_qkv_to_track = mixed_qkv[
:, forward_metadata.track_conv_indices
].transpose(0, 1)
# Apply mask and assign to destinations
mask_indices = forward_batch.mamba_track_mask.nonzero(as_tuple=True)[0]
conv_states[conv_dst[mask_indices]] = mixed_qkv_to_track
mixed_qkv = causal_conv1d_fn(
mixed_qkv.transpose(0, 1),
mixed_qkv,
conv_weights,
bias,
activation=activation,
@@ -710,12 +1059,13 @@ class GDNAttnBackend(MambaAttnBackendBase):
use_qk_l2norm_in_kernel=True,
disable_state_update=True,
intermediate_states_buffer=intermediate_state_cache,
intermediate_state_indices=intermediate_state_indices,
cache_steps=forward_batch.spec_info.draft_token_num,
retrieve_parent_token=retrieve_parent_token,
)
else:
recurrent_state = ssm_states[cache_indices]
core_attn_out, last_recurrent_state = chunk_gated_delta_rule(
core_attn_out, last_recurrent_state, h = chunk_gated_delta_rule(
q=query,
k=key,
v=value,
@@ -730,6 +1080,10 @@ class GDNAttnBackend(MambaAttnBackendBase):
last_recurrent_state = last_recurrent_state.to(ssm_states.dtype, copy=False)
ssm_states[cache_indices] = last_recurrent_state
self._track_mamba_state_extend(
forward_batch, h, ssm_states, forward_metadata
)
return core_attn_out
@@ -965,14 +1319,23 @@ class HybridLinearAttnBackend(AttentionBackend):
**kwargs,
)
def update_mamba_state_after_mtp_verify(self, accepted_indices, model):
request_number = accepted_indices.shape[0]
def update_mamba_state_after_mtp_verify(
self,
accepted_steps: torch.Tensor,
mamba_track_indices: Optional[torch.Tensor],
mamba_steps_to_track: Optional[torch.Tensor],
model,
):
request_number = accepted_steps.shape[0]
state_indices_tensor = (
self.linear_attn_backend.forward_metadata.mamba_cache_indices[
:request_number
]
)
intermediate_state_indices = torch.arange(
request_number, dtype=torch.int32, device=state_indices_tensor.device
)
mamba_caches = (
self.linear_attn_backend.req_to_token_pool.get_speculative_mamba2_params_all_layers()
@@ -983,19 +1346,41 @@ class HybridLinearAttnBackend(AttentionBackend):
intermediate_state_cache = mamba_caches.intermediate_ssm
intermediate_conv_window_cache = mamba_caches.intermediate_conv_window[0]
# SSM state updates (chunked to reduce peak memory)
valid_mask = accepted_indices >= 0
# Compute common indices once to avoid duplication
valid_state_indices = state_indices_tensor[valid_mask].to(torch.int64) # [N]
last_steps = accepted_indices[valid_mask].to(torch.int64) # [N]
valid_mask = accepted_steps >= 0
dst_state_indices = state_indices_tensor[valid_mask].to(torch.int64) # [N]
src_state_indices = intermediate_state_indices[valid_mask].to(
torch.int64
) # [N]
last_steps = accepted_steps[valid_mask].to(torch.int64) # [N]
# scatter into ssm_states at the chosen cache lines
ssm_states[:, valid_state_indices, :] = intermediate_state_cache[
:, valid_state_indices, last_steps
ssm_states[:, dst_state_indices, :] = intermediate_state_cache[
:, src_state_indices, last_steps
].to(ssm_states.dtype, copy=False)
# Scatter into conv_states at the chosen cache lines
conv_states[:, valid_state_indices, :, :] = intermediate_conv_window_cache[
:, valid_state_indices, last_steps
conv_states[:, dst_state_indices, :] = intermediate_conv_window_cache[
:, src_state_indices, last_steps
].to(conv_states.dtype, copy=False)
# Track indices used for tracking mamba states for prefix cache
if mamba_track_indices is not None:
assert mamba_steps_to_track is not None
track_mask = mamba_steps_to_track >= 0
track_steps = mamba_steps_to_track[track_mask].to(torch.int64) # [N]
if track_steps.numel() == 0:
# No track indices to update
return
dst_track_indices = mamba_track_indices[track_mask].to(torch.int64)
src_track_indices = intermediate_state_indices[track_mask].to(torch.int64)
# scatter into ssm_states at the chosen track states
ssm_states[:, dst_track_indices, :] = intermediate_state_cache[
:, src_track_indices, track_steps
].to(ssm_states.dtype, copy=False)
# scatter into conv_states at the chosen track states
conv_states[:, dst_track_indices, :] = intermediate_conv_window_cache[
:, src_track_indices, track_steps
].to(conv_states.dtype, copy=False)
@@ -578,6 +578,7 @@ def _causal_conv1d_update_kernel(
conv_state_indices_ptr,
num_accepted_tokens_ptr,
intermediate_conv_window_ptr,
intermediate_state_indices_ptr,
retrieve_next_token_ptr,
retrieve_next_sibling_ptr,
retrieve_parent_token_ptr,
@@ -602,6 +603,7 @@ def _causal_conv1d_update_kernel(
stride_inter_step: tl.constexpr,
stride_inter_dim: tl.constexpr,
stride_inter_win: tl.constexpr,
stride_intermediate_state_indices: tl.constexpr,
stride_retrieve_next_token_seq: tl.constexpr,
stride_retrieve_next_token_token: tl.constexpr,
stride_retrieve_next_sibling_seq: tl.constexpr,
@@ -639,6 +641,11 @@ def _causal_conv1d_update_kernel(
conv_state_batch_coord = tl.load(
conv_state_indices_ptr + idx_seq * stride_state_indices
).to(tl.int64)
if SAVE_INTERMEDIATE:
intermediate_state_batch_coord = tl.load(
intermediate_state_indices_ptr
+ idx_seq * stride_intermediate_state_indices
).to(tl.int64)
else:
conv_state_batch_coord = idx_seq
if USE_PAD_SLOT: # noqa
@@ -847,7 +854,7 @@ def _causal_conv1d_update_kernel(
# Layout: [seq(cache line), step, dim, win(K-1)]
base_ptr = (
intermediate_conv_window_ptr
+ conv_state_batch_coord * stride_inter_seq
+ intermediate_state_batch_coord * stride_inter_seq
+ idx_token * stride_inter_step
+ idx_feats * stride_inter_dim
)
@@ -934,7 +941,7 @@ def _causal_conv1d_update_kernel(
# Layout: [seq(cache line), step, dim, win(K-1)]
base_ptr = (
intermediate_conv_window_ptr
+ conv_state_batch_coord * stride_inter_seq
+ intermediate_state_batch_coord * stride_inter_seq
+ idx_token * stride_inter_step
+ idx_feats * stride_inter_dim
)
@@ -980,6 +987,7 @@ def causal_conv1d_update(
conv_state_indices: Optional[torch.Tensor] = None,
num_accepted_tokens: Optional[torch.Tensor] = None,
intermediate_conv_window: Optional[torch.Tensor] = None,
intermediate_state_indices: Optional[torch.Tensor] = None,
retrieve_next_token: Optional[torch.Tensor] = None,
retrieve_next_sibling: Optional[torch.Tensor] = None,
retrieve_parent_token: Optional[torch.Tensor] = None,
@@ -1040,6 +1048,8 @@ def causal_conv1d_update(
assert conv_state.size(0) >= batch
else:
assert (batch,) == conv_state_indices.shape
assert intermediate_state_indices is not None
assert (batch,) == intermediate_state_indices.shape
assert num_cache_lines >= batch
assert weight.stride(1) == 1 # Need this
@@ -1056,6 +1066,11 @@ def causal_conv1d_update(
stride_state_indices = (
conv_state_indices.stride(0) if conv_state_indices is not None else 0
)
stride_intermediate_state_indices = (
intermediate_state_indices.stride(0)
if intermediate_state_indices is not None
else 0
)
if num_accepted_tokens is not None:
state_len = width - 1 + (seqlen - 1) # effective state_len needed
else:
@@ -1117,6 +1132,7 @@ def causal_conv1d_update(
conv_state_indices,
num_accepted_tokens,
intermediate_conv_window if intermediate_conv_window is not None else x,
intermediate_state_indices,
retrieve_next_token,
retrieve_next_sibling,
retrieve_parent_token,
@@ -1141,6 +1157,7 @@ def causal_conv1d_update(
stride_inter_step,
stride_inter_dim,
stride_inter_win,
stride_intermediate_state_indices,
stride_retrieve_next_token_seq,
stride_retrieve_next_token_token,
stride_retrieve_next_sibling_seq,
@@ -27,9 +27,17 @@ from sglang.srt.model_executor.forward_batch_info import ForwardBatch
class ForwardMetadata:
query_start_loc: torch.Tensor
mamba_cache_indices: torch.Tensor
# For topk > 1 eagle
retrieve_next_token: Optional[torch.Tensor] = None
retrieve_next_sibling: Optional[torch.Tensor] = None
retrieve_parent_token: Optional[torch.Tensor] = None
# For prefill radix cache
track_conv_indices: Optional[torch.Tensor] = None
track_ssm_h_src: Optional[torch.Tensor] = None
track_ssm_h_dst: Optional[torch.Tensor] = None
track_ssm_final_src: Optional[torch.Tensor] = None
track_ssm_final_dst: Optional[torch.Tensor] = None
is_target_verify: bool = False
draft_token_num: int = 1
@@ -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
@@ -24,3 +24,5 @@ class CacheInitParams:
enable_metrics: bool = False
enable_kv_cache_events: bool = False
enable_mamba_extra_buffer: bool = False
+12 -2
View File
@@ -19,6 +19,10 @@ from sglang.srt.utils.common import ceil_align
if TYPE_CHECKING:
from sglang.srt.managers.schedule_batch import Req, ScheduleBatch
# 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
MAMBA_STATE_PER_REQ_NO_CACHE = 1
logger = logging.getLogger(__name__)
@@ -300,9 +304,15 @@ def alloc_req_slots(
"""Allocate request slots from the pool."""
if isinstance(req_to_token_pool, HybridReqToTokenPool):
mamba_available_size = req_to_token_pool.mamba_pool.available_size()
if mamba_available_size < num_reqs:
factor = (
MAMBA_STATE_PER_REQ_PREFIX_CACHE
if isinstance(tree_cache, MambaRadixCache)
else 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 isinstance(tree_cache, MambaRadixCache):
mamba_num = max(0, num_reqs - mamba_available_size)
mamba_num = max(0, mamba_state_needed - mamba_available_size)
tree_cache.evict_mamba(mamba_num)
req_pool_indices = req_to_token_pool.alloc(num_reqs, reqs)
else:
+224 -53
View File
@@ -21,16 +21,24 @@ The radix tree data structure for managing the hybrid (full and Mamba) KV cache.
import heapq
from collections import defaultdict
from functools import partial
from typing import TYPE_CHECKING, List, Optional, Tuple
import torch
from numpy import float64
from sglang.srt.mem_cache.allocator import TokenToKVPoolAllocator
from sglang.srt.distributed import get_tensor_model_parallel_rank
from sglang.srt.layers.attention.fla.chunk_delta_h import CHUNK_SIZE as FLA_CHUNK_SIZE
from sglang.srt.mem_cache.allocator import (
PagedTokenToKVPoolAllocator,
TokenToKVPoolAllocator,
)
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache, MatchResult
from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool
from sglang.srt.mem_cache.radix_cache import (
RadixKey,
_key_match_page_size1,
_key_match_paged,
get_child_key,
)
@@ -252,6 +260,30 @@ class LRUList:
return False
return node.id in self.cache
def pretty_print(self, tree_cache: Optional["MambaRadixCache"] = None):
"""
Pretty print the lru list
"""
msg = f"{self.mamba=} LRU list: "
x_lru = self._get_lru()
while x_lru is not None and x_lru.id in self.cache:
msg += f"[{x_lru.id}] {x_lru.last_access_time:f} -> "
x_lru = getattr(x_lru, self.prv)
print(msg)
if not tree_cache:
return
msg = f"{self.mamba=} Nodes (sorted by last_access_time): "
if self.mamba:
nodes = tree_cache._collect_nontombstone_nodes()
else:
nodes = tree_cache._collect_all_nodes()
heapq.heapify(nodes)
while len(nodes):
x = heapq.heappop(nodes)
msg += f"[{x.id}] {x.last_access_time:f} -> "
print(msg)
# Note: this is expensive, only use for debug
def sanity_check_evictable_size(self):
"""
@@ -292,9 +324,13 @@ class LRUList:
if x == tree_cache.root_node:
# root node is not in the lru list
continue
assert (
x_lru is not None and x_lru.id in self.cache
), f"Incorrect LRU list, x_lru is None or not in cache: {x_lru=}, {x.id=}"
assert (
x == x_lru
), f"Incorrect LRU list, {self.mamba=}, x: {x.id=} != x_lru: {x_lru.id=}"
), f"Incorrect LRU list, {self.mamba=}, x: {x.id=} != x_lru: {x_lru.id=}, {x.last_access_time=}, {x_lru.last_access_time=}"
assert (
x_lru.full_lock_ref == 0
), f"x_lru should not be locked when idle, {x_lru.full_lock_ref=}, {x_lru.id=}"
@@ -314,22 +350,33 @@ class LRUList:
evictable_size == lru_list_evictable_size
), f"{self.mamba=}, total nodes: {total_nodes}, total lru: {total_lru}, evictable size: {evictable_size} != lru list evictable size: {lru_list_evictable_size}"
except Exception as e:
msg = f"Mamba Radix tree sanity check failed, ping @yizhang2077: {e}"
logger.error(msg)
raise Exception(msg)
if get_tensor_model_parallel_rank() == 0:
msg = f"Mamba Radix tree sanity check failed, ping @yizhang2077: {e}"
logger.error(msg)
tree_cache.pretty_print()
tree_cache.full_lru_list.pretty_print(tree_cache)
tree_cache.mamba_lru_list.pretty_print(tree_cache)
raise Exception(msg)
class MambaRadixCache(BasePrefixCache):
def __init__(self, params: CacheInitParams):
assert isinstance(params.token_to_kv_pool_allocator, TokenToKVPoolAllocator)
self.req_to_token_pool = params.req_to_token_pool
assert isinstance(
params.token_to_kv_pool_allocator, TokenToKVPoolAllocator
) or isinstance(params.token_to_kv_pool_allocator, PagedTokenToKVPoolAllocator)
self.req_to_token_pool: HybridReqToTokenPool = params.req_to_token_pool
self.token_to_kv_pool_allocator = params.token_to_kv_pool_allocator
assert (
params.page_size == 1
), "Only support page_size=1 in mamba radix cache now."
self.page_size = params.page_size
self.disable = params.disable
self.enable_mamba_extra_buffer = params.enable_mamba_extra_buffer
if not self.enable_mamba_extra_buffer:
assert (
self.page_size == 1
), f"Page size must be 1 for MambaRadixCache v1, got {self.page_size}"
else:
logger.info(f"Mamba extra_buffer is enabled.")
if self.token_to_kv_pool_allocator:
self.device = self.token_to_kv_pool_allocator.device
@@ -339,15 +386,19 @@ class MambaRadixCache(BasePrefixCache):
if params.enable_metrics:
self.init_metrics_collector()
self.key_match_fn = _key_match_page_size1
self.get_child_key_fn = get_child_key
if self.page_size == 1:
self.key_match_fn = _key_match_page_size1
self.get_child_key_fn = get_child_key
else:
self.key_match_fn = partial(_key_match_paged, page_size=self.page_size)
self.get_child_key_fn = partial(get_child_key, page_size=self.page_size)
self.reset()
##### Public API #####
def reset(self) -> None:
self.root_node = TreeNode()
self.root_node.key = []
self.root_node.key = RadixKey([], None)
self.root_node.value = []
self.root_node.full_lock_ref = 1
self.root_node.mamba_lock_ref = 1
@@ -384,7 +435,7 @@ class MambaRadixCache(BasePrefixCache):
last_host_node=self.root_node,
)
value, last_node = self._match_prefix_helper(key)
value, last_node, mamba_branching_seqlen = self._match_prefix_helper(key)
# copy mamba state to req local space if cow is true
if cow_mamba and last_node.mamba_value is not None:
@@ -415,18 +466,26 @@ class MambaRadixCache(BasePrefixCache):
device_indices=value,
last_device_node=last_node,
last_host_node=last_node,
mamba_branching_seqlen=mamba_branching_seqlen,
)
def insert(self, key: RadixKey, value=None, mamba_value=None) -> Tuple[int, bool]:
if self.disable:
return 0
return 0, False
if value is None:
value = torch.tensor([x for x in key.token_ids], dtype=torch.int64)
return self._insert_helper(self.root_node, key, value, mamba_value)
def cache_finished_req(self, req: Req, is_insert: bool = True):
def cache_finished_req(self, req: Req, is_insert: bool = True) -> None:
"""Cache request when it finishes."""
# for abort with prefix cache hit and before alloc is called
if req.req_pool_idx is None:
if req.mamba_pool_idx is not None:
self.req_to_token_pool.mamba_pool.free(req.mamba_pool_idx.unsqueeze(-1))
req.mamba_pool_idx = None
return
kv_committed_len = req.pop_committed_kv_cache()
if self.disable:
@@ -442,56 +501,135 @@ class MambaRadixCache(BasePrefixCache):
req.req_pool_idx, :kv_committed_len
]
page_aligned_len = len(kv_indices)
page_aligned_kv_indices = kv_indices.to(dtype=torch.int64, copy=True)
# Radix Cache takes one ref in memory pool
# insert the token_ids and kv_indices into the radix tree
# Note: the insert function already frees the overlapped kv_indices
mamba_value = req.mamba_pool_idx.unsqueeze(-1).clone()
if is_insert:
cache_len = (
req.mamba_last_track_seqlen
if self.enable_mamba_extra_buffer
else len(token_ids)
)
if cache_len is None:
cache_len = 0
if cache_len != len(token_ids):
cache_end_idx = max(cache_len, req.cache_protected_len)
self.token_to_kv_pool_allocator.free(kv_indices[cache_end_idx:])
token_ids = token_ids[:cache_len]
kv_indices = kv_indices[:cache_len]
if self.page_size != 1:
page_aligned_len = len(kv_indices) // self.page_size * self.page_size
page_aligned_kv_indices = kv_indices[:page_aligned_len].to(
dtype=torch.int64, copy=True
)
else:
page_aligned_len = len(kv_indices)
page_aligned_kv_indices = kv_indices.to(dtype=torch.int64, copy=True)
assert (
cache_len == page_aligned_len
), f"It is required {cache_len=}, {page_aligned_len=}, {kv_committed_len=}, {len(req.origin_input_ids)=}, {len(req.output_ids)=} ping @yizhang2077 if you see this"
# Radix Cache takes one ref in memory pool
# 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
)
)
mamba_value = (
req.mamba_ping_pong_track_buffer[
mamba_ping_pong_track_buffer_to_keep
]
.unsqueeze(-1)
.clone()
)
else:
mamba_value = req.mamba_pool_idx.unsqueeze(-1).clone()
mamba_ping_pong_track_buffer_to_keep = None
new_prefix_len, mamba_exist = self.insert(
RadixKey(token_ids[:page_aligned_len], req.extra_key),
page_aligned_kv_indices,
mamba_value,
)
self.token_to_kv_pool_allocator.free(
kv_indices[len(req.prefix_indices) : new_prefix_len]
kv_indices[req.cache_protected_len : new_prefix_len]
)
else:
self.token_to_kv_pool_allocator.free(
kv_indices[len(req.prefix_indices) : page_aligned_len]
)
self.token_to_kv_pool_allocator.free(kv_indices[req.cache_protected_len :])
mamba_exist = True
if req.req_pool_idx is not None:
self.req_to_token_pool.free(req.req_pool_idx, free_mamba_cache=mamba_exist)
self.dec_lock_ref(req.last_node)
else: # for abort case
self.req_to_token_pool.mamba_pool.free(mamba_value)
if mamba_exist:
mamba_ping_pong_track_buffer_to_keep = None
free_mamba_cache = True if self.enable_mamba_extra_buffer else mamba_exist
self.req_to_token_pool.free(
req.req_pool_idx,
free_mamba_cache=free_mamba_cache,
mamba_ping_pong_track_buffer_to_keep=mamba_ping_pong_track_buffer_to_keep,
)
self.dec_lock_ref(req.last_node)
def cache_unfinished_req(self, req: Req, chunked=False) -> None:
"""Cache request when it is unfinished."""
if self.disable:
def _skip_cache_unfinished_req(req: Req) -> None:
kv_indices = self.req_to_token_pool.req_to_token[
req.req_pool_idx, : len(req.fill_ids)
]
# `req.prefix_indices` will be used in `PrefillAdder::add_chunked_req` later
req.prefix_indices = kv_indices
req.prefix_indices = kv_indices.to(dtype=torch.int64, copy=True)
return
token_ids = req.fill_ids
kv_indices = self.req_to_token_pool.req_to_token[
cache_len = (
req.mamba_last_track_seqlen
if self.enable_mamba_extra_buffer
else len(token_ids)
)
if self.disable or cache_len is None:
return _skip_cache_unfinished_req(req)
kv_indices_orig = self.req_to_token_pool.req_to_token[
req.req_pool_idx, : len(token_ids)
]
page_aligned_len = len(kv_indices)
page_aligned_kv_indices = kv_indices.to(dtype=torch.int64, copy=True)
# kv_indices is the kv indices to be cached
kv_indices = kv_indices_orig[:cache_len]
if self.page_size != 1:
page_aligned_len = len(kv_indices) // self.page_size * self.page_size
page_aligned_kv_indices = kv_indices[:page_aligned_len].to(
dtype=torch.int64, copy=True
)
else:
page_aligned_len = len(kv_indices)
page_aligned_kv_indices = kv_indices.to(dtype=torch.int64, copy=True)
assert page_aligned_len == len(
kv_indices
), f"page_aligned_len != len(kv_indices), {page_aligned_len=}, {len(kv_indices)=}, {cache_len=}, {self.page_size=}, {FLA_CHUNK_SIZE=}"
page_aligned_token_ids = token_ids[:page_aligned_len]
mamba_value = self.req_to_token_pool.get_mamba_indices(
req.req_pool_idx
).unsqueeze(-1)
if self.enable_mamba_extra_buffer:
# copy from the ping pong track 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 = (
req.mamba_ping_pong_track_buffer[mamba_ping_pong_track_buffer_to_keep]
.unsqueeze(-1)
.clone()
)
else:
mamba_value = self.req_to_token_pool.get_mamba_indices(
req.req_pool_idx
).unsqueeze(-1)
# radix tree mamba value is forked from req space
mamba_value_forked = self.req_to_token_pool.mamba_pool.fork_from(mamba_value)
@@ -508,7 +646,7 @@ class MambaRadixCache(BasePrefixCache):
mamba_value_forked,
)
self.token_to_kv_pool_allocator.free(
kv_indices[len(req.prefix_indices) : new_prefix_len]
kv_indices[req.cache_protected_len : new_prefix_len]
)
# there is a mamba cache in radix cache, release it
if mamba_exist:
@@ -526,21 +664,28 @@ class MambaRadixCache(BasePrefixCache):
if not mamba_exist:
assert torch.equal(new_last_node.mamba_value, mamba_value_forked)
assert len(req.prefix_indices) <= len(
assert (
req.cache_protected_len <= len(new_indices) + self.page_size - 1
), f"{req.cache_protected_len=}, {len(new_indices)=}, {len(page_aligned_token_ids)=}, {mamba_exist=}"
assert new_prefix_len <= len(
new_indices
), f"{req.prefix_indices=}, {new_indices=}"
assert new_prefix_len <= len(new_indices), f"{new_prefix_len=}, {new_indices=}"
), f"{new_prefix_len=}, {len(new_indices)=}"
self.req_to_token_pool.write(
(req.req_pool_idx, slice(len(req.prefix_indices), len(new_indices))),
new_indices[len(req.prefix_indices) :],
(req.req_pool_idx, slice(req.cache_protected_len, len(new_indices))),
new_indices[req.cache_protected_len :],
)
self.dec_lock_ref(req.last_node)
self.inc_lock_ref(new_last_node)
# `req.prefix_indices` will be used in `PrefillAdder::add_chunked_req` later
req.prefix_indices = new_indices
# NOTE: this is needed for both page_size == 1 and page_size > 1
req.prefix_indices = torch.cat(
[new_indices, kv_indices_orig[len(new_indices) :]]
)
req.cache_protected_len = len(new_indices)
req.mamba_last_track_seqlen = None
req.last_node = new_last_node
def pretty_print(self) -> None:
@@ -670,7 +815,7 @@ class MambaRadixCache(BasePrefixCache):
It unlocks the mamba_lock_ref for current node if its mamba_value exists.
"""
if self.disable:
return
return None
if node.mamba_value is not None:
assert (
@@ -692,6 +837,8 @@ class MambaRadixCache(BasePrefixCache):
node = node.parent
def sanity_check(self):
if self.disable:
return
self.full_lru_list.sanity_check(self)
self.mamba_lru_list.sanity_check(self)
@@ -734,13 +881,25 @@ class MambaRadixCache(BasePrefixCache):
_dfs_helper(child)
_dfs_helper(self.root_node)
return torch.cat(values)
return torch.cat(values) if len(values) > 0 else torch.tensor([])
def all_mamba_values_flatten(self) -> torch.Tensor:
values = []
def _dfs_helper(node: TreeNode):
if node.mamba_value is not None:
values.append(node.mamba_value)
for _, child in node.children.items():
_dfs_helper(child)
_dfs_helper(self.root_node)
return torch.cat(values) if len(values) > 0 else torch.tensor([])
##### Internal Helper Functions #####
def _match_prefix_helper(
self, key: RadixKey
) -> Tuple[List[torch.Tensor], TreeNode]:
) -> Tuple[List[torch.Tensor], TreeNode, Optional[int]]:
"""
Mamba prefix matching helper. It factors in the sliding window size such that
the matched node is guaranteed to either 1. connected to root without mamba tombstone,
@@ -750,7 +909,7 @@ class MambaRadixCache(BasePrefixCache):
node = self.root_node
child_key = self.get_child_key_fn(key)
value = []
value: List[torch.Tensor] = []
best_value_len = 0
best_last_node = node
while len(key) > 0 and child_key in node.children.keys():
@@ -793,7 +952,19 @@ class MambaRadixCache(BasePrefixCache):
)
node_update = node_update.parent
return value[:best_value_len], best_last_node
# Calculate the branching point. It is defined as the last aligned position that
# does not have a mamba value.
if len(value) > best_value_len:
fla_chunk_aligned_seqlen = (
sum(len(v) for v in value) // FLA_CHUNK_SIZE
) * FLA_CHUNK_SIZE
mamba_branching_seqlen = (
fla_chunk_aligned_seqlen if fla_chunk_aligned_seqlen > 0 else None
)
else:
mamba_branching_seqlen = None
return value[:best_value_len], best_last_node, mamba_branching_seqlen
def _split_node(self, key: RadixKey, child: TreeNode, split_len: int) -> TreeNode:
# new_node -> child
+97 -13
View File
@@ -152,6 +152,7 @@ class MambaPool:
self,
*,
size: int,
spec_state_size: int,
cache_params: BaseLinearStateParams,
device: str,
enable_memory_saver: bool = False,
@@ -198,7 +199,7 @@ class MambaPool:
intermediate_ssm_state_cache = torch.zeros(
size=(
num_mamba_layers,
size + 1,
spec_state_size + 1,
speculative_num_draft_tokens,
temporal_state_shape[0],
temporal_state_shape[1],
@@ -213,7 +214,7 @@ class MambaPool:
torch.zeros(
size=(
num_mamba_layers,
size + 1,
spec_state_size + 1,
speculative_num_draft_tokens,
conv_shape[0],
conv_shape[1],
@@ -267,6 +268,10 @@ class MambaPool:
select_index = self.free_slots[:need_size]
self.free_slots = self.free_slots[need_size:]
# clear at alloc time
for i in range(len(self.mamba_cache.conv)):
self.mamba_cache.conv[i][:, select_index] = 0
self.mamba_cache.temporal[:, select_index] = 0
return select_index
@@ -274,9 +279,6 @@ class MambaPool:
if free_index.numel() == 0:
return
self.free_slots = torch.cat((self.free_slots, free_index))
for i in range(len(self.mamba_cache.conv)):
self.mamba_cache.conv[i][:, free_index] = 0
self.mamba_cache.temporal[:, free_index] = 0
def clear(self):
# Zero the entire mamba cache before resetting free_slots
@@ -333,10 +335,12 @@ class HybridReqToTokenPool(ReqToTokenPool):
*,
size: int,
mamba_size: int,
mamba_spec_state_size: int,
max_context_len: int,
device: str,
enable_memory_saver: bool,
cache_params: BaseLinearStateParams,
enable_mamba_extra_buffer: bool,
speculative_num_draft_tokens: int = None,
):
super().__init__(
@@ -345,23 +349,32 @@ class HybridReqToTokenPool(ReqToTokenPool):
device=device,
enable_memory_saver=enable_memory_saver,
)
self.mamba_ping_pong_track_buffer_size = (
2 if speculative_num_draft_tokens is None else 1
)
self.enable_mamba_extra_buffer = enable_mamba_extra_buffer
self.enable_memory_saver = enable_memory_saver
self._init_mamba_pool(
size=mamba_size,
mamba_spec_state_size=mamba_spec_state_size,
cache_params=cache_params,
device=device,
enable_mamba_extra_buffer=enable_mamba_extra_buffer,
speculative_num_draft_tokens=speculative_num_draft_tokens,
)
def _init_mamba_pool(
self,
size: int,
mamba_spec_state_size: int,
cache_params: BaseLinearStateParams,
device: str,
enable_mamba_extra_buffer: bool,
speculative_num_draft_tokens: int = None,
):
self.mamba_pool = MambaPool(
size=size,
spec_state_size=mamba_spec_state_size,
cache_params=cache_params,
device=device,
enable_memory_saver=self.enable_memory_saver,
@@ -373,32 +386,67 @@ class HybridReqToTokenPool(ReqToTokenPool):
self.req_index_to_mamba_index_mapping: torch.Tensor = torch.zeros(
size, dtype=torch.int32, device=self.device
)
if enable_mamba_extra_buffer:
self.req_index_to_mamba_ping_pong_track_buffer_mapping: torch.Tensor = (
torch.zeros(
(size, self.mamba_ping_pong_track_buffer_size),
dtype=torch.int32,
device=self.device,
)
)
# For chunk prefill req, we do not need to allocate mamba cache,
# We could use allocated mamba cache instead.
def alloc(
self, need_size: int, reqs: Optional[List[Req]] = None
) -> Optional[List[int]]:
def alloc(self, need_size: int, reqs: Optional[List["Req"]]) -> Optional[List[int]]:
assert reqs is not None
select_index = super().alloc(need_size)
if select_index == None:
return None
mamba_index = []
mamba_ping_pong_track_buffer_list = []
for req in reqs:
mid = None
if req.mamba_pool_idx is not None: # for radix cache
mid = req.mamba_pool_idx
else:
mid = self.mamba_pool.alloc(1)[0]
mid = self.mamba_pool.alloc(1)
assert (
mid is not None
), f"Not enough space for mamba cache, try to increase --mamba-full-memory-ratio or --max-mamba-cache-size. {mid=}, {self.mamba_pool.size=}, {self.mamba_pool.available_size()=}, {len(reqs)=}"
mid = mid[0]
req.mamba_pool_idx = mid
if mid is not None:
mamba_index.append(mid)
mamba_index.append(mid)
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
mamba_ping_pong_track_buffer_list.append(
req.mamba_ping_pong_track_buffer.tolist()
)
assert len(select_index) == len(
mamba_index
), f"Not enough space for mamba cache, try to increase --max-mamba-cache-size."
), f"Not enough space for mamba cache, try to increase --mamba-full-memory-ratio or --max-mamba-cache-size."
if self.enable_mamba_extra_buffer:
assert len(select_index) == len(
mamba_ping_pong_track_buffer_list
), f"Not enough space for mamba ping pong idx, try to increase --mamba-full-memory-ratio."
self.req_index_to_mamba_index_mapping[select_index] = torch.tensor(
mamba_index, dtype=torch.int32, device=self.device
)
if self.enable_mamba_extra_buffer:
self.req_index_to_mamba_ping_pong_track_buffer_mapping[select_index] = (
torch.tensor(
mamba_ping_pong_track_buffer_list,
dtype=torch.int32,
device=self.device,
)
)
return select_index
def get_mamba_indices(self, req_indices: torch.Tensor) -> torch.Tensor:
@@ -411,8 +459,19 @@ class HybridReqToTokenPool(ReqToTokenPool):
def get_speculative_mamba2_params_all_layers(self) -> MambaPool.SpeculativeState:
return self.mamba_pool.get_speculative_mamba2_params_all_layers()
def get_mamba_ping_pong_other_idx(self, mamba_next_track_idx: int) -> int:
if self.mamba_ping_pong_track_buffer_size == 2:
return 1 - mamba_next_track_idx
else:
return mamba_next_track_idx
# For chunk prefill, we can not free mamba cache, we need use it in the future
def free(self, free_index: Union[int, List[int]], free_mamba_cache: bool = True):
def free(
self,
free_index: Union[int, List[int]],
free_mamba_cache: bool = True,
mamba_ping_pong_track_buffer_to_keep: Optional[int] = None,
):
if isinstance(free_index, (int,)):
free_index = [free_index]
super().free(free_index)
@@ -420,9 +479,31 @@ class HybridReqToTokenPool(ReqToTokenPool):
mamba_index = self.req_index_to_mamba_index_mapping[free_index]
self.mamba_pool.free(mamba_index)
if self.enable_mamba_extra_buffer:
mamba_ping_pong_track_buffer_to_free = (
self.req_index_to_mamba_ping_pong_track_buffer_mapping[
free_index
].squeeze(0)
)
if mamba_ping_pong_track_buffer_to_keep is not None:
assert mamba_ping_pong_track_buffer_to_keep in [
0,
1,
], f"mamba_ping_pong_track_buffer_to_keep must be 0 or 1, {mamba_ping_pong_track_buffer_to_keep=}"
idx_to_free = list(range(self.mamba_ping_pong_track_buffer_size))
idx_to_free.remove(mamba_ping_pong_track_buffer_to_keep)
mamba_ping_pong_track_buffer_to_free = (
mamba_ping_pong_track_buffer_to_free[idx_to_free]
)
self.mamba_pool.free(mamba_ping_pong_track_buffer_to_free)
def clear(self):
logger.info("Reset HybridReqToTokenPool")
super().clear()
self.mamba_pool.clear()
self.req_index_to_mamba_index_mapping.zero_()
if self.enable_mamba_extra_buffer:
self.req_index_to_mamba_ping_pong_track_buffer_mapping.zero_()
class KVCache(abc.ABC):
@@ -1133,6 +1214,9 @@ class HybridLinearKVPool(KVCache):
cache_v,
)
def move_kv_cache(self, tgt_loc: torch.Tensor, src_loc: torch.Tensor):
self.full_kv_pool.move_kv_cache(tgt_loc, src_loc)
def get_v_head_dim(self):
return self.full_kv_pool.get_value_buffer(0).shape[-1]
@@ -321,6 +321,11 @@ class CudaGraphRunner:
num_tokens_per_bs=self.num_tokens_per_bs,
)
enable_mamba_track = (
self.model_runner.server_args.enable_mamba_extra_buffer()
and self.model_runner.spec_algorithm.is_none()
)
if self.require_gathered_buffer:
assert self.require_mlp_tp_gather or self.require_attn_tp_gather
self.buffers: GraphInputBuffers = GraphInputBuffers.create(
@@ -338,6 +343,7 @@ class CudaGraphRunner:
encoder_len_fill_value=self.encoder_len_fill_value,
num_tokens_per_bs=self.num_tokens_per_bs,
cache_loc_dtype=self._cache_loc_dtype(),
enable_mamba_track=enable_mamba_track,
)
self.tbo_plugin = TboCudaGraphRunnerPlugin()
@@ -537,7 +543,7 @@ class CudaGraphRunner:
def capture_one_batch_size(
self, bs: int, forward: Callable, stream_idx: Optional[int] = None
):
buffers = self.buffers
buffers: GraphInputBuffers = self.buffers
graph = self._create_device_graph()
stream = self.stream
num_tokens = bs * self.num_tokens_per_bs
@@ -611,6 +617,18 @@ class CudaGraphRunner:
else:
lora_ids = None
# mamba state tracking
mamba_track_indices = (
buffers.mamba_track_indices[:bs]
if buffers.mamba_track_indices is not None
else None
)
mamba_track_mask = (
buffers.mamba_track_mask[:bs]
if buffers.mamba_track_mask is not None
else None
)
if stream_idx is None:
attn_backend = self.model_runner.attn_backend
else:
@@ -631,6 +649,9 @@ class CudaGraphRunner:
attn_backend=attn_backend,
out_cache_loc=out_cache_loc,
seq_lens_sum=seq_lens.sum().item(),
mamba_track_indices=mamba_track_indices,
mamba_track_mask=mamba_track_mask,
mamba_track_seqlens=None, # Prefill only
encoder_lens=encoder_lens,
return_logprob=False,
positions=positions,
@@ -249,6 +249,12 @@ class ForwardBatch:
# The indices of output tokens in the token_to_kv_pool_swa
# TODO(shiyang, biao): integrate out_cache_loc_swa into multiple attention backends
out_cache_loc_swa: Optional[torch.Tensor] = None
# The indices to track mamba state with
mamba_track_indices: Optional[torch.Tensor] = None # shape: [b], int64
# The mask to track mamba state if needed
mamba_track_mask: Optional[torch.Tensor] = None # shape: [b], bool
# The seqlens to track mamba state if masked, prefill only.
mamba_track_seqlens: Optional[torch.Tensor] = None # shape: [b], int64
# Optional seq_lens on cpu
seq_lens_cpu: Optional[torch.Tensor] = None
@@ -398,6 +404,9 @@ class ForwardBatch:
req_pool_indices=batch.req_pool_indices,
seq_lens=batch.seq_lens,
out_cache_loc=batch.out_cache_loc,
mamba_track_indices=batch.mamba_track_indices,
mamba_track_mask=batch.mamba_track_mask,
mamba_track_seqlens=batch.mamba_track_seqlens,
mm_inputs=batch.multimodal_inputs,
encoder_cached=batch.encoder_cached,
encoder_lens=batch.encoder_lens,
@@ -881,6 +890,16 @@ class ForwardBatch:
if self.encoder_lens is not None:
self.encoder_lens = self._pad_tensor_to_size(self.encoder_lens, bs)
self.positions = self._pad_tensor_to_size(self.positions, num_tokens)
if self.mamba_track_indices is not None:
self.mamba_track_indices = self._pad_tensor_to_size(
self.mamba_track_indices, bs
)
if self.mamba_track_mask is not None:
self.mamba_track_mask = self._pad_tensor_to_size(self.mamba_track_mask, bs)
if self.mamba_track_seqlens is not None:
self.mamba_track_seqlens = self._pad_tensor_to_size(
self.mamba_track_seqlens, bs
)
if self.mrope_positions is not None:
self.mrope_positions = self._pad_tensor_to_size(self.mrope_positions, bs)
@@ -25,6 +25,8 @@ class GraphInputBuffers:
num_token_non_padded: torch.Tensor
custom_mask: torch.Tensor
next_token_logits_buffer: torch.Tensor
mamba_track_indices: Optional[torch.Tensor]
mamba_track_mask: Optional[torch.Tensor]
global_num_tokens_gpu: torch.Tensor
global_num_tokens_for_logprob_gpu: torch.Tensor
encoder_lens: Optional[torch.Tensor]
@@ -48,6 +50,7 @@ class GraphInputBuffers:
encoder_len_fill_value: int,
num_tokens_per_bs: int,
cache_loc_dtype: torch.dtype,
enable_mamba_track: bool,
) -> "GraphInputBuffers":
with torch.device(device):
input_ids = torch.zeros((max_num_token,), dtype=torch.int64)
@@ -66,6 +69,14 @@ class GraphInputBuffers:
(max_num_token, vocab_size),
dtype=torch.float,
)
mamba_track_indices = (
torch.zeros((max_bs,), dtype=torch.int64)
if enable_mamba_track
else None
)
mamba_track_mask = (
torch.zeros((max_bs,), dtype=torch.bool) if enable_mamba_track else None
)
if pp_size > 1:
pp_proxy_tensors = {
@@ -111,6 +122,8 @@ class GraphInputBuffers:
num_token_non_padded=num_token_non_padded,
custom_mask=custom_mask,
next_token_logits_buffer=next_token_logits_buffer,
mamba_track_indices=mamba_track_indices,
mamba_track_mask=mamba_track_mask,
encoder_lens=encoder_lens,
global_num_tokens_gpu=global_num_tokens_gpu,
global_num_tokens_for_logprob_gpu=global_num_tokens_for_logprob_gpu,
@@ -134,6 +147,10 @@ class GraphInputBuffers:
if bs != raw_bs:
self.seq_lens.fill_(seq_len_fill_value)
self.out_cache_loc.zero_()
if self.mamba_track_indices is not None:
self.mamba_track_indices.zero_()
if self.mamba_track_mask is not None:
self.mamba_track_mask.fill_(False)
# Common inputs
self.input_ids[:raw_num_token].copy_(forward_batch.input_ids)
@@ -142,6 +159,17 @@ class GraphInputBuffers:
self.out_cache_loc[:raw_num_token].copy_(forward_batch.out_cache_loc)
self.positions[:raw_num_token].copy_(forward_batch.positions)
if (
self.mamba_track_indices is not None
and forward_batch.mamba_track_indices is not None
):
self.mamba_track_indices[:raw_bs].copy_(forward_batch.mamba_track_indices)
if (
self.mamba_track_mask is not None
and forward_batch.mamba_track_mask is not None
):
self.mamba_track_mask[:raw_bs].copy_(forward_batch.mamba_track_mask)
seq_lens_cpu: Optional[torch.Tensor] = None
if forward_batch.seq_lens_cpu is not None:
if bs != raw_bs:
@@ -238,8 +238,10 @@ def add_chunked_prefix_cache_attention_backend(backend_name):
# Detect stragger ranks in model loading
UNBALANCED_MODEL_LOADING_TIMEOUT_S = 480 # leave more time for post data processing
# the ratio of mamba cache pool size to max_running_requests, it will be safe when it is larger than 2 (yizhang2077)
# 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_NO_OVERLAP = 1
logger = logging.getLogger(__name__)
@@ -1446,14 +1448,9 @@ class ModelRunner:
server_args = self.server_args
assert config is not None
speculativa_ratio = (
0
if server_args.speculative_num_draft_tokens is None
else server_args.speculative_num_draft_tokens
)
if (
server_args.disable_radix_cache
or config.mamba2_cache_params.mamba_cache_per_req == 0
or server_args.max_mamba_cache_size is not None
):
# with disable radix cache, sets the max_mamba_cache_size based on the max_running_requests
if server_args.max_mamba_cache_size is None:
@@ -1461,7 +1458,25 @@ class ModelRunner:
server_args.max_mamba_cache_size = server_args.max_running_requests
else:
server_args.max_mamba_cache_size = 512
server_args.max_mamba_cache_size = server_args.max_mamba_cache_size // (
server_args.dp_size if server_args.enable_dp_attention else 1
)
else:
assert config.mamba2_cache_params.mamba_cache_per_req > 0
# reserve the memory for the intermediate mamba states used for spec dec
if not self.spec_algorithm.is_none():
assert server_args.speculative_num_draft_tokens is not None
assert server_args.max_running_requests is not None
mamba_state_intermediate_size = (
config.mamba2_cache_params.mamba_cache_per_req
* server_args.max_running_requests
* server_args.speculative_num_draft_tokens
)
total_rest_memory = total_rest_memory - (
mamba_state_intermediate_size / (1 << 30)
)
# allocate the memory based on the ratio between mamba state memory vs. full kv cache memory
# solve the equations:
# 1. mamba_state_memory + full_kv_cache_memory == total_rest_memory
@@ -1475,21 +1490,22 @@ class ModelRunner:
server_args.max_mamba_cache_size = int(
(mamba_state_memory_raw * (1 << 30))
// config.mamba2_cache_params.mamba_cache_per_req
// (1 + speculativa_ratio)
)
if self.hybrid_gdn_config is not None:
server_args.max_mamba_cache_size = server_args.max_mamba_cache_size // (
server_args.dp_size if server_args.enable_dp_attention else 1
)
mamba_state_memory = (
server_args.max_mamba_cache_size
* config.mamba2_cache_params.mamba_cache_per_req
* (1 + speculativa_ratio)
/ (1 << 30)
)
return total_rest_memory - mamba_state_memory
@property
def qwen3_next_config(self):
config = self.model_config.hf_config
if isinstance(config, Qwen3NextConfig):
return config
return None
@property
def hybrid_gdn_config(self):
config = self.model_config.hf_config
@@ -1683,11 +1699,18 @@ class ModelRunner:
)
if self.mambaish_config is not None:
ratio = (
MAMBA_CACHE_SIZE_MAX_RUNNING_REQUESTS_RATIO
if not self.server_args.disable_radix_cache
else 1
)
additional_ratio = 0
if (
self.server_args.enable_mamba_extra_buffer()
and not self.spec_algorithm.is_none()
):
additional_ratio = MAMBA_CACHE_V2_ADDITIONAL_RATIO_NO_OVERLAP
else:
additional_ratio = MAMBA_CACHE_V2_ADDITIONAL_RATIO_OVERLAP
if self.server_args.disable_radix_cache:
ratio = 1
else:
ratio = MAMBA_CACHE_SIZE_MAX_RUNNING_REQUESTS_RATIO + additional_ratio
max_num_reqs = min(
max_num_reqs, self.server_args.max_mamba_cache_size // ratio
)
@@ -1789,11 +1812,13 @@ class ModelRunner:
self.req_to_token_pool = HybridReqToTokenPool(
size=max_num_reqs,
mamba_size=self.server_args.max_mamba_cache_size,
mamba_spec_state_size=max_num_reqs,
max_context_len=self.model_config.context_len
+ extra_max_context_len,
device=self.device,
enable_memory_saver=self.server_args.enable_memory_saver,
cache_params=config.mamba2_cache_params,
enable_mamba_extra_buffer=self.server_args.enable_mamba_extra_buffer(),
speculative_num_draft_tokens=self.server_args.speculative_num_draft_tokens,
)
else:
+64 -8
View File
@@ -30,6 +30,7 @@ import orjson
from sglang.srt.connector import ConnectorType
from sglang.srt.environ import ToolStrictLevel, envs
from sglang.srt.function_call.function_call_parser import FunctionCallParser
from sglang.srt.layers.attention.fla.chunk_delta_h import CHUNK_SIZE as FLA_CHUNK_SIZE
from sglang.srt.lora.lora_registry import LoRARef
from sglang.srt.parser.reasoning_parser import ReasoningParser
from sglang.srt.utils.common import (
@@ -185,6 +186,8 @@ FP8_GEMM_RUNNER_BACKEND_CHOICES = [
MAMBA_SSM_DTYPE_CHOICES = ["float32", "bfloat16"]
mamba_scheduler_strategy_CHOICES = ["auto", "no_buffer", "extra_buffer"]
# Allow external code to add more choices
def add_load_format_choices(choices):
@@ -466,6 +469,8 @@ class ServerArgs:
max_mamba_cache_size: Optional[int] = None
mamba_ssm_dtype: str = "float32"
mamba_full_memory_ratio: float = 0.9
mamba_scheduler_strategy: str = "auto"
mamba_track_interval: int = 256
# Hierarchical cache
enable_hierarchical_cache: bool = False
@@ -737,6 +742,10 @@ class ServerArgs:
self.random_seed = random.randint(0, 1 << 30)
if self.mm_process_config is None:
self.mm_process_config = {}
if self.mamba_scheduler_strategy == "auto":
# TODO: when extra_buffer is more verified, we can set the default path based on
# [overlap, non-overlap]
self.mamba_scheduler_strategy = "no_buffer"
# Handle ModelScope model downloads
if get_bool_env_var("SGLANG_USE_MODELSCOPE"):
@@ -1333,12 +1342,6 @@ class ServerArgs:
f"{model_arch}"
)
elif model_arch in ["Qwen3NextForCausalLM"]:
if not self.disable_radix_cache:
logger.warning(
"Disabling overlap schedule since MambaRadixCache is not compatible with "
"overlap schedule currently, try to use --disable-radix-cache if overlap schedule is necessary"
)
self.disable_overlap_schedule = True
if is_sm100_supported():
quantization_config = getattr(hf_config, "quantization_config", None)
quant_method = (
@@ -1372,15 +1375,52 @@ class ServerArgs:
)
self.disable_radix_cache = True
self.disable_overlap_schedule = False
# Mamba radix cache v2
if self.enable_mamba_extra_buffer():
assert (
is_cuda()
), "Mamba extra_buffer is only supported on CUDA devices with FLA backend"
assert (
self.disaggregation_mode == "null"
), "Mamba extra_buffer is not compatible with disaggregation mode yet."
if self.speculative_num_draft_tokens is not None:
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}"
if self.page_size is not None:
assert (
self.mamba_track_interval % self.page_size == 0
), f"mamba_track_interval {self.mamba_track_interval} must be divisible by page_size {self.page_size}"
assert (
FLA_CHUNK_SIZE % self.page_size == 0
), f"Page size for hybrid GDN model must be divisible by {FLA_CHUNK_SIZE}, got {self.page_size}"
if self.speculative_algorithm is not None:
logger.info(
f"Disable overlap schedule for {model_arch} model speculative decoding."
)
self.disable_overlap_schedule = True
elif not self.disable_radix_cache:
logger.warning(
"Disabling overlap schedule since MambaRadixCache no_buffer is not compatible with "
"overlap schedule currently, try to use --mamba-scheduler-strategy extra_buffer to enable overlap schedule"
)
self.disable_overlap_schedule = True
elif model_arch in [
"FalconH1ForCausalLM",
"JetNemotronForCausalLM",
"JetVLMForConditionalGeneration",
]:
assert (
not self.enable_mamba_extra_buffer()
), f"mamba extra_buffer is not supported for {model_arch} model"
if not self.disable_radix_cache:
logger.warning(
"Disabling overlap schedule since MambaRadixCache is not compatible with "
"overlap schedule currently, try to use --disable-radix-cache if overlap schedule is necessary"
"Disabling overlap schedule since mamba no_buffer is not compatible with "
"overlap schedule, try to use --disable-radix-cache if overlap schedule is necessary"
)
self.disable_overlap_schedule = True
if is_sm100_supported():
@@ -3535,6 +3575,19 @@ class ServerArgs:
default=ServerArgs.mamba_full_memory_ratio,
help="The ratio of mamba state memory to full kv cache memory.",
)
parser.add_argument(
"--mamba-scheduler-strategy",
type=str,
choices=mamba_scheduler_strategy_CHOICES,
default=ServerArgs.mamba_scheduler_strategy,
help="The strategy to use for mamba radix cache.",
)
parser.add_argument(
"--mamba-track-interval",
type=int,
default=ServerArgs.mamba_track_interval,
help="The interval to track the mamba state during decode.",
)
# Hierarchical cache
parser.add_argument(
@@ -4326,6 +4379,9 @@ class ServerArgs:
model_config = self.get_model_config()
return model_config.attention_arch == AttentionArch.MLA
def enable_mamba_extra_buffer(self) -> bool:
return self.mamba_scheduler_strategy == "extra_buffer"
def check_server_args(self):
# Check parallel size constraints
assert (
@@ -147,6 +147,16 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin):
bs,
)
if get_global_server_args().enable_mamba_extra_buffer():
batch.mamba_track_indices = torch.tensor(
[
req.mamba_ping_pong_track_buffer[req.mamba_next_track_idx]
for req in batch.reqs
],
dtype=torch.int64,
device=batch.device,
)
def generate_attn_arg_prefill(
self,
req_pool_indices: torch.Tensor,
+83 -38
View File
@@ -316,6 +316,7 @@ class EAGLEWorker(TpModelWorker):
logits_output=logits_output,
next_token_ids=verify_output.verified_id,
num_accepted_tokens=sum(verify_output.accept_length_per_req_cpu),
accept_length_per_req_cpu=verify_output.accept_length_per_req_cpu,
can_run_cuda_graph=can_run_cuda_graph,
)
@@ -669,6 +670,7 @@ class EAGLEWorker(TpModelWorker):
pass
def verify(self, batch: ScheduleBatch, spec_info: EagleVerifyInput):
seq_lens_pre_verify = batch.seq_lens.clone()
spec_info.prepare_for_verify(batch, self.page_size)
spec_info.num_tokens_per_batch = self.speculative_num_steps + 1
batch.return_hidden_states = False
@@ -743,44 +745,8 @@ class EAGLEWorker(TpModelWorker):
self.target_worker.model_runner.hybrid_gdn_config is not None
or self.target_worker.model_runner.mamba2_config is not None
):
accepted_length = (
torch.tensor(
res.accept_length_per_req_cpu,
device=logits_output.hidden_states.device,
dtype=torch.int64,
)
+ 1
)
# If topk > 1, we need to use retrieve_next_token and retrieve_next_sibling to handle the eagle tree custom attention mask
# res.accepted_indices.shape[0] > 0 skips DP attn idle batch
if spec_info.topk > 1 and res.accepted_indices.shape[0] > 0:
# accepted_indices=[0,2,3,4,5,7,9,10,11], accepted_length=[4, 3, 2], cumulative_accepted_lengths=[4, 7, 9]
# first_token_indices_per_req=prepend(0, accepted_indices[cumulative_accepted_lengths[:-1]]) = [0, 5, 10]
# last_token_indices_per_req=accepted_indices[cumulative_accepted_lengths - 1] = [4, 9, 11] (last token ID of each req)
# max_relative_indices_per_req = [4,4,1]; those are the per-req spec-decoding step offsets that contain the correct mamba caches
cumulative_accepted_lengths = torch.cumsum(accepted_length, dim=0)
req_start_positions = torch.cat(
[
torch.zeros(
1,
dtype=cumulative_accepted_lengths.dtype,
device=cumulative_accepted_lengths.device,
),
cumulative_accepted_lengths[:-1],
]
)
first_token_indices_per_req = res.accepted_indices[req_start_positions]
last_token_indices_per_req = res.accepted_indices[
cumulative_accepted_lengths - 1
]
max_relative_indices_per_req = (
last_token_indices_per_req - first_token_indices_per_req
)
else:
max_relative_indices_per_req = accepted_length - 1
self.target_worker.model_runner.attn_backend.update_mamba_state_after_mtp_verify(
max_relative_indices_per_req, self.target_worker.model_runner.model
self._mamba_verify_update(
batch, res, logits_output, spec_info, seq_lens_pre_verify
)
if batch.return_logprob:
@@ -794,6 +760,85 @@ class EAGLEWorker(TpModelWorker):
return logits_output, res, model_worker_batch, can_run_cuda_graph
def _mamba_verify_update(
self,
batch: ScheduleBatch,
res: EagleVerifyOutput,
logits_output: LogitsProcessorOutput,
spec_info: EagleVerifyInput,
seq_lens_pre_verify: torch.Tensor,
):
accepted_length = (
torch.tensor(
res.accept_length_per_req_cpu,
device=logits_output.hidden_states.device,
dtype=torch.int64,
)
+ 1
)
cumulative_accepted_lengths = torch.cumsum(accepted_length, dim=0)
# prepend 0 to the cumulative_accepted_lengths
accepted_indices_start = torch.cat(
[
torch.zeros(
1,
dtype=cumulative_accepted_lengths.dtype,
device=cumulative_accepted_lengths.device,
),
cumulative_accepted_lengths[:-1],
]
)
accepted_indices_offset = torch.arange(
0,
len(batch.seq_lens) * batch.spec_info.draft_token_num,
step=batch.spec_info.draft_token_num,
dtype=accepted_indices_start.dtype,
device=accepted_indices_start.device,
)
# If topk > 1, we need to use retrieve_next_token and retrieve_next_sibling to handle the eagle tree custom attention mask
# res.accepted_indices.shape[0] > 0 skips DP attn idle batch
if spec_info.topk > 1 and res.accepted_indices.shape[0] > 0:
# accepted_indices=[0,2,3,4,5,7,9,10,11], accepted_length=[4, 3, 2], cumulative_accepted_lengths=[4, 7, 9]
# first_token_indices_per_req=prepend(0, accepted_indices[cumulative_accepted_lengths[:-1]]) = [0, 5, 10]
# last_token_indices_per_req=accepted_indices[cumulative_accepted_lengths - 1] = [4, 9, 11] (last token ID of each req)
# max_relative_indices_per_req = [4,4,1]; those are the per-req spec-decoding step offsets that contain the correct mamba caches
# first_token_indices_per_req = res.accepted_indices[accepted_indices_start]
accepted_steps = (
res.accepted_indices[cumulative_accepted_lengths - 1]
- accepted_indices_offset
)
else:
accepted_steps = accepted_length - 1
if batch.mamba_track_indices is not None:
# If after verify, the request's seq_lens has crossed a mamba track interval,
# we need to update the mamba state for the request at the crossing point.
mamba_track_interval = self.server_args.mamba_track_interval
to_track_mask = (
seq_lens_pre_verify // mamba_track_interval
!= batch.seq_lens // mamba_track_interval
)
tracking_point = (
batch.seq_lens // mamba_track_interval * mamba_track_interval
)
to_track_ith = torch.clamp(tracking_point - seq_lens_pre_verify - 1, min=0)
mamba_steps_to_track = torch.where(
to_track_mask,
res.accepted_indices[to_track_ith + accepted_indices_start]
- accepted_indices_offset,
-1,
)
else:
mamba_steps_to_track = None
self.target_worker.model_runner.attn_backend.update_mamba_state_after_mtp_verify(
accepted_steps=accepted_steps,
mamba_track_indices=batch.mamba_track_indices,
mamba_steps_to_track=mamba_steps_to_track,
model=self.target_worker.model_runner.model,
)
def add_logprob_values(
self,
batch: ScheduleBatch,
+4 -2
View File
@@ -1,6 +1,7 @@
import inspect
import json
import os
import random
import numpy as np
import requests
@@ -79,8 +80,9 @@ def get_input_ids(
break
text = format_longbench_v2_example(example)
tokens = tokenizer.encode(text)
# Truncate to max_tokens
input_ids.append(tokens[:max_prompt_tokens])
# Truncate to a random length between 0.5x and 1.5x of max_prompt_tokens
truncate_len = int(max_prompt_tokens * random.uniform(0.5, 1.5))
input_ids.append(tokens[:truncate_len])
# Save to local cache
with open(cache_file, "w") as f: