diff --git a/python/sglang/srt/model_executor/piecewise_cuda_graph_runner.py b/python/sglang/srt/model_executor/piecewise_cuda_graph_runner.py index 2dcd59ceb..556ddcfdb 100644 --- a/python/sglang/srt/model_executor/piecewise_cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/piecewise_cuda_graph_runner.py @@ -127,6 +127,13 @@ def set_torch_compile_config(): class PiecewiseCudaGraphRunner: """A PiecewiseCudaGraphRunner runs the forward pass of a model with cuda graph and torch.compile.""" + def is_mamba_track_enabled(self): + return ( + self.model_runner.server_args.enable_mamba_extra_buffer() + and not self.model_runner.server_args.disable_radix_cache + and self.model_runner.spec_algorithm.is_none() + ) + def __init__(self, model_runner: ModelRunner): # Parse args self.model_runner = model_runner @@ -175,8 +182,10 @@ class PiecewiseCudaGraphRunner: self.capture_hidden_mode = CaptureHiddenMode.FULL self.max_num_tokens = max(self.capture_num_tokens) + self.max_bs = model_runner.req_to_token_pool.size self.is_multimodal = model_runner.is_multimodal + self.mamba_track_enabled = self.is_mamba_track_enabled() # Graph inputs with torch.device(self.device): @@ -189,6 +198,21 @@ class PiecewiseCudaGraphRunner: if model_runner.is_hybrid_swa else None ) + self.mamba_track_indices = ( + torch.zeros((self.max_bs,), dtype=torch.int64) + if self.mamba_track_enabled + else None + ) + self.mamba_track_mask = ( + torch.zeros((self.max_bs,), dtype=torch.bool) + if self.mamba_track_enabled + else None + ) + self.mamba_track_seqlens = ( + torch.zeros((self.max_bs,), dtype=torch.int32) + if self.mamba_track_enabled + else None + ) self.positions = torch.zeros((self.max_num_tokens,), dtype=torch.int64) self.tbo_plugin = TboCudaGraphRunnerPlugin() @@ -269,6 +293,19 @@ class PiecewiseCudaGraphRunner: if self.out_cache_loc_swa is not None else None ) + mamba_track_indices = ( + self.mamba_track_indices[:1] + if self.mamba_track_indices is not None + else None + ) + mamba_track_mask = ( + self.mamba_track_mask[:1] if self.mamba_track_mask is not None else None + ) + mamba_track_seqlens = ( + self.mamba_track_seqlens[:1] + if self.mamba_track_seqlens is not None + else None + ) with torch.device(self.device): forward_batch = ForwardBatch( forward_mode=ForwardMode.EXTEND, @@ -286,6 +323,9 @@ class PiecewiseCudaGraphRunner: out_cache_loc=out_cache_loc, out_cache_loc_swa=out_cache_loc_swa, seq_lens_sum=num_tokens, + mamba_track_indices=mamba_track_indices, + mamba_track_mask=mamba_track_mask, + mamba_track_seqlens=mamba_track_seqlens, encoder_lens=None, return_logprob=False, extend_num_tokens=num_tokens, @@ -386,6 +426,19 @@ class PiecewiseCudaGraphRunner: if self.out_cache_loc_swa is not None else None ) + mamba_track_indices = ( + self.mamba_track_indices[:bs] + if self.mamba_track_indices is not None + else None + ) + mamba_track_mask = ( + self.mamba_track_mask[:bs] if self.mamba_track_mask is not None else None + ) + mamba_track_seqlens = ( + self.mamba_track_seqlens[:bs] + if self.mamba_track_seqlens is not None + else None + ) positions = self.positions[:num_tokens] mrope_positions = ( self.mrope_positions[:, :num_tokens] if self.is_multimodal else None @@ -417,6 +470,9 @@ class PiecewiseCudaGraphRunner: out_cache_loc=out_cache_loc, out_cache_loc_swa=out_cache_loc_swa, seq_lens_sum=num_tokens, + mamba_track_indices=mamba_track_indices, + mamba_track_mask=mamba_track_mask, + mamba_track_seqlens=mamba_track_seqlens, encoder_lens=None, return_logprob=False, extend_num_tokens=num_tokens, @@ -504,6 +560,23 @@ class PiecewiseCudaGraphRunner: forward_batch.out_cache_loc ) ) + + if ( + self.mamba_track_indices is not None + and forward_batch.mamba_track_indices is not None + ): + self.mamba_track_indices[: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[:bs].copy_(forward_batch.mamba_track_mask) + if ( + self.mamba_track_seqlens is not None + and forward_batch.mamba_track_seqlens is not None + ): + self.mamba_track_seqlens[:bs].copy_(forward_batch.mamba_track_seqlens) + input_ids = self.input_ids[:static_num_tokens] positions = self.positions[:static_num_tokens] out_cache_loc = self.out_cache_loc[:static_num_tokens] @@ -514,6 +587,19 @@ class PiecewiseCudaGraphRunner: else None ) + mamba_track_indices = ( + self.mamba_track_indices[:bs] + if self.mamba_track_indices is not None + else None + ) + mamba_track_mask = ( + self.mamba_track_mask[:bs] if self.mamba_track_mask is not None else None + ) + mamba_track_seqlens = ( + self.mamba_track_seqlens[:bs] + if self.mamba_track_seqlens is not None + else None + ) if forward_batch.mrope_positions is not None: self.mrope_positions[:, :num_tokens].copy_(forward_batch.mrope_positions) @@ -546,6 +632,9 @@ class PiecewiseCudaGraphRunner: out_cache_loc=out_cache_loc, out_cache_loc_swa=out_cache_loc_swa, seq_lens_sum=forward_batch.seq_lens_sum, + mamba_track_indices=mamba_track_indices, + mamba_track_mask=mamba_track_mask, + mamba_track_seqlens=mamba_track_seqlens, encoder_lens=forward_batch.encoder_lens, return_logprob=False, extend_seq_lens=forward_batch.extend_seq_lens,