diff --git a/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py b/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py index afe767a6c..c61dd84a6 100644 --- a/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py +++ b/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py @@ -362,6 +362,29 @@ class AscendAttnBackend(AttentionBackend): ): pass + def init_forward_metadata_out_graph( + self, + forward_batch: ForwardBatch, + in_capture: bool = False, + ): + bs = forward_batch.batch_size + if in_capture: + self._init_cuda_graph_metadata( + bs, forward_batch.forward_mode, forward_batch.seq_lens + ) + self._apply_cuda_graph_metadata( + bs=bs, + req_pool_indices=forward_batch.req_pool_indices, + seq_lens=forward_batch.seq_lens, + seq_lens_cpu=( + forward_batch.seq_lens.cpu() + if in_capture + else forward_batch.seq_lens_cpu + ), + forward_mode=forward_batch.forward_mode, + spec_info=forward_batch.spec_info, + ) + def init_forward_metadata(self, forward_batch: ForwardBatch): """Init the metadata for a forward pass.""" self.forward_metadata = ForwardMetadata() @@ -538,39 +561,19 @@ class AscendAttnBackend(AttentionBackend): self.graph_metadata[bs] = metadata return metadata - def init_forward_metadata_capture_cuda_graph( + def _apply_cuda_graph_metadata( self, bs: int, - num_tokens: int, req_pool_indices: torch.Tensor, seq_lens: torch.Tensor, - encoder_lens: Optional[torch.Tensor], + seq_lens_cpu: torch.Tensor, forward_mode: ForwardMode, spec_info: Optional[SpecInput], ): - self._init_cuda_graph_metadata(bs, forward_mode, seq_lens) - self.init_forward_metadata_replay_cuda_graph( - bs=bs, - req_pool_indices=req_pool_indices, - seq_lens=seq_lens, - seq_lens_sum=None, - encoder_lens=encoder_lens, - forward_mode=forward_mode, - spec_info=spec_info, - seq_lens_cpu=seq_lens.cpu(), - ) + """Shared capture+replay body for the cuda-graph init path. - def init_forward_metadata_replay_cuda_graph( - self, - bs: int, - req_pool_indices: torch.Tensor, - seq_lens: torch.Tensor, - seq_lens_sum: int, - encoder_lens: Optional[torch.Tensor], - forward_mode: ForwardMode, - spec_info: Optional[SpecInput], - seq_lens_cpu: Optional[torch.Tensor], - ): + Public entry: :py:meth:`init_forward_metadata_out_graph`. + """ metadata = self.graph_metadata[bs] max_len = seq_lens_cpu[:bs].max().item() if forward_mode.is_target_verify(): @@ -2395,6 +2398,32 @@ class AscendAttnMultiStepDraftBackend: for i in range(self.speculative_num_steps - 1): call_fn(i, forward_batch) + def init_forward_metadata_out_graph( + self, + forward_batch: ForwardBatch, + in_capture: bool = False, + ): + from sglang.srt.model_executor.forward_batch_info import build_inner_fb_view + + inner_fb = build_inner_fb_view( + forward_batch, + bs=forward_batch.batch_size, + forward_mode=ForwardMode.DECODE, + ) + + def call_fn(i, _forward_batch): + self.attn_backends[i].init_forward_metadata_out_graph( + inner_fb, in_capture=in_capture + ) + + self.common_template(forward_batch, call_fn) + + def init_forward_metadata_in_graph(self, forward_batch: ForwardBatch) -> None: + def call_fn(i, _forward_batch): + self.attn_backends[i].init_forward_metadata_in_graph(forward_batch) + + self.common_template(forward_batch, call_fn) + def init_forward_metadata(self, forward_batch: ForwardBatch): def call_fn(i, forward_batch): assert forward_batch.spec_info is not None @@ -2405,34 +2434,3 @@ class AscendAttnMultiStepDraftBackend: def init_cuda_graph_state(self, max_bs, max_num_tokens): for i in range(self.speculative_num_steps): self.attn_backends[i].init_cuda_graph_state(max_bs, max_num_tokens) - - def init_forward_metadata_capture_cuda_graph(self, forward_batch: ForwardBatch): - def call_fn(i, forward_batch): - self.attn_backends[i].init_forward_metadata_capture_cuda_graph( - forward_batch.batch_size, - forward_batch.batch_size * self.topk, - forward_batch.req_pool_indices, - forward_batch.seq_lens, - encoder_lens=None, - forward_mode=ForwardMode.DECODE, - spec_info=forward_batch.spec_info, - ) - - self.common_template(forward_batch, call_fn) - - def init_forward_metadata_replay_cuda_graph( - self, forward_batch: ForwardBatch, bs: int - ): - def call_fn(i, forward_batch): - self.attn_backends[i].init_forward_metadata_replay_cuda_graph( - bs, - forward_batch.req_pool_indices, - forward_batch.seq_lens, - seq_lens_sum=-1, - encoder_lens=None, - forward_mode=ForwardMode.DECODE, - spec_info=forward_batch.spec_info, - seq_lens_cpu=forward_batch.seq_lens_cpu, - ) - - self.common_template(forward_batch, call_fn) diff --git a/python/sglang/srt/hardware_backend/npu/attention/ascend_gdn_backend.py b/python/sglang/srt/hardware_backend/npu/attention/ascend_gdn_backend.py index 109d474c1..419992437 100644 --- a/python/sglang/srt/hardware_backend/npu/attention/ascend_gdn_backend.py +++ b/python/sglang/srt/hardware_backend/npu/attention/ascend_gdn_backend.py @@ -72,6 +72,21 @@ class AscendGDNAttnBackend(AscendMambaAttnBackendBase): else: self.ssm_state_indices = cache_indices + def init_forward_metadata_out_graph( + self, + forward_batch: ForwardBatch, + in_capture: bool = False, + ): + if forward_batch.forward_mode.is_draft_extend(True): + return + super().init_forward_metadata_out_graph(forward_batch, in_capture=in_capture) + self.prepare_gdn_inputs( + forward_batch.batch_size, + forward_batch.forward_mode, + forward_batch.spec_info, + ) + self.graph_mode = True + def init_forward_metadata(self, forward_batch: ForwardBatch): if forward_batch.forward_mode.is_draft_extend(True): return @@ -83,53 +98,6 @@ class AscendGDNAttnBackend(AscendMambaAttnBackendBase): ) self.graph_mode = False - def init_forward_metadata_capture_cuda_graph( - self, - bs: int, - num_tokens: int, - req_pool_indices: torch.Tensor, - seq_lens: torch.Tensor, - encoder_lens: Optional[torch.Tensor], - forward_mode: ForwardMode, - spec_info: Optional[Union[EagleDraftInput, EagleVerifyInput]], - ): - self.init_forward_metadata_replay_cuda_graph( - bs=bs, - req_pool_indices=req_pool_indices, - seq_lens=seq_lens, - seq_lens_sum=None, - encoder_lens=encoder_lens, - forward_mode=forward_mode, - spec_info=spec_info, - seq_lens_cpu=seq_lens.cpu(), - ) - - def init_forward_metadata_replay_cuda_graph( - self, - bs: int, - req_pool_indices: torch.Tensor, - seq_lens: torch.Tensor, - seq_lens_sum: int, - encoder_lens: Optional[torch.Tensor], - forward_mode: ForwardMode, - spec_info: Optional[Union[EagleDraftInput, EagleVerifyInput]], - seq_lens_cpu: Optional[torch.Tensor], - ): - if forward_mode.is_draft_extend(True): - return - super().init_forward_metadata_replay_cuda_graph( - bs, - req_pool_indices, - seq_lens, - seq_lens_sum, - encoder_lens, - forward_mode, - spec_info, - seq_lens_cpu, - ) - self.prepare_gdn_inputs(bs, forward_mode, spec_info) - self.graph_mode = True - def forward_decode( self, layer: RadixLinearAttention, diff --git a/python/sglang/srt/hardware_backend/npu/attention/ascend_hybrid_linear_attn_backend.py b/python/sglang/srt/hardware_backend/npu/attention/ascend_hybrid_linear_attn_backend.py index de7cf58fa..f11aecc11 100644 --- a/python/sglang/srt/hardware_backend/npu/attention/ascend_hybrid_linear_attn_backend.py +++ b/python/sglang/srt/hardware_backend/npu/attention/ascend_hybrid_linear_attn_backend.py @@ -131,9 +131,13 @@ class AscendMambaAttnBackendBase(MambaAttnBackendBase): spec_info: Optional[SpecInput], seq_lens_cpu: Optional[torch.Tensor], ): - num_padding = torch.count_nonzero( - seq_lens_cpu == self.get_cuda_graph_seq_len_fill_value() - ) + # out_graph passes seq_lens_cpu=None at capture; mirror the base guard. + if seq_lens_cpu is None: + num_padding = 0 + else: + num_padding = torch.count_nonzero( + seq_lens_cpu == self.get_cuda_graph_seq_len_fill_value() + ) # Make sure forward metadata is correctly handled for padding reqs req_pool_indices[bs - num_padding :] = 0 mamba_indices = self.req_to_token_pool.get_mamba_indices(req_pool_indices) diff --git a/python/sglang/srt/layers/attention/aiter_backend.py b/python/sglang/srt/layers/attention/aiter_backend.py index 9b0fcbf42..61f0c20aa 100755 --- a/python/sglang/srt/layers/attention/aiter_backend.py +++ b/python/sglang/srt/layers/attention/aiter_backend.py @@ -820,6 +820,25 @@ class AiterAttnBackend(AttentionBackend): ) return output + def init_forward_metadata_out_graph( + self, + forward_batch: ForwardBatch, + in_capture: bool = False, + ): + seq_lens_cpu = ( + forward_batch.seq_lens.cpu() if in_capture else forward_batch.seq_lens_cpu + ) + self._apply_cuda_graph_metadata( + bs=forward_batch.batch_size, + req_pool_indices=forward_batch.req_pool_indices, + seq_lens=forward_batch.seq_lens, + seq_lens_sum=None if in_capture else forward_batch.seq_lens_sum, + encoder_lens=forward_batch.encoder_lens, + forward_mode=forward_batch.forward_mode, + spec_info=forward_batch.spec_info, + seq_lens_cpu=seq_lens_cpu, + ) + def init_forward_metadata(self, forward_batch: ForwardBatch): """Init auxiliary variables for aiter attention backend.""" @@ -1482,28 +1501,7 @@ class AiterAttnBackend(AttentionBackend): device=self.device, ) - def init_forward_metadata_capture_cuda_graph( - self, - bs: int, - num_tokens: int, - req_pool_indices: torch.Tensor, - seq_lens: torch.Tensor, - encoder_lens: Optional[torch.Tensor], - forward_mode: ForwardMode, - spec_info: Optional[SpecInput], - ): - self.init_forward_metadata_replay_cuda_graph( - bs=bs, - req_pool_indices=req_pool_indices, - seq_lens=seq_lens, - seq_lens_sum=None, - encoder_lens=encoder_lens, - forward_mode=forward_mode, - spec_info=spec_info, - seq_lens_cpu=seq_lens.cpu(), - ) - - def init_forward_metadata_replay_cuda_graph( + def _apply_cuda_graph_metadata( self, bs: int, req_pool_indices: torch.Tensor, @@ -2866,33 +2864,26 @@ class AiterMultiStepDraftBackend: max_bs, max_num_tokens, kv_indices_buf=self.cuda_graph_kv_indices[i] ) - def init_forward_metadata_capture_cuda_graph(self, forward_batch: ForwardBatch): - def call_fn(i, forward_batch): - self.attn_backends[i].init_forward_metadata_capture_cuda_graph( - forward_batch.batch_size, - forward_batch.batch_size * self.topk, - forward_batch.req_pool_indices, - forward_batch.seq_lens, - encoder_lens=None, - forward_mode=ForwardMode.DECODE, - spec_info=forward_batch.spec_info, - ) - - self.common_template(forward_batch, self.cuda_graph_kv_indices, call_fn) - - def init_forward_metadata_replay_cuda_graph( - self, forward_batch: ForwardBatch, bs: int + def init_forward_metadata_out_graph( + self, + forward_batch: ForwardBatch, + in_capture: bool = False, ): - def call_fn(i, forward_batch): - self.attn_backends[i].init_forward_metadata_replay_cuda_graph( - bs, - forward_batch.req_pool_indices, - forward_batch.seq_lens, - seq_lens_sum=-1, - encoder_lens=None, - forward_mode=ForwardMode.DECODE, - spec_info=forward_batch.spec_info, - seq_lens_cpu=forward_batch.seq_lens_cpu, + from sglang.srt.model_executor.forward_batch_info import build_inner_fb_view + + inner_fb = build_inner_fb_view( + forward_batch, + bs=forward_batch.batch_size, + forward_mode=ForwardMode.DECODE, + ) + + def call_fn(i, _forward_batch): + self.attn_backends[i].init_forward_metadata_out_graph( + inner_fb, in_capture=in_capture ) self.common_template(forward_batch, self.cuda_graph_kv_indices, call_fn) + + def init_forward_metadata_in_graph(self, forward_batch: ForwardBatch) -> None: + for attn_backend in self.attn_backends: + attn_backend.init_forward_metadata_in_graph(forward_batch) diff --git a/python/sglang/srt/layers/attention/base_attn_backend.py b/python/sglang/srt/layers/attention/base_attn_backend.py index d757ce79a..49cf92ba7 100644 --- a/python/sglang/srt/layers/attention/base_attn_backend.py +++ b/python/sglang/srt/layers/attention/base_attn_backend.py @@ -1,6 +1,6 @@ from __future__ import annotations -from abc import ABC, abstractmethod +from abc import ABC from typing import TYPE_CHECKING, Optional import torch @@ -11,52 +11,84 @@ from sglang.srt.utils.common import is_npu if TYPE_CHECKING: from sglang.srt.layers.attention.dsa.dsa_indexer import BaseIndexerMetadata from sglang.srt.layers.radix_attention import RadixAttention - from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode + from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.speculative.spec_info import SpecInput class AttentionBackend(ABC): - """The base class of attention backends""" + """The base class of attention backends. + + Forward-data init contract (3 methods): + + - ``init_forward_metadata(fb)`` — eager entry point. Default is a wrapper + that calls ``_out_graph(fb)`` then ``_in_graph(fb)``. Backends may + override to keep an independent eager body. + - ``init_forward_metadata_out_graph(fb, in_capture=False)`` — per-iter + metadata prep, runs outside ``with graph.capture():``. Capture + sites pass ``in_capture=True``; replay/eager use the default + ``False``. Backends read ``in_capture`` only when capture / replay + bodies diverge. + - ``init_forward_metadata_in_graph(fb)`` — graph-recordable static-shape + GPU op, runs inside ``with graph.capture():`` at capture time and + is auto-replayed by ``graph.replay()``. Default is no-op. + + The legacy ``init_forward_metadata_capture_cuda_graph`` and + ``init_forward_metadata_replay_cuda_graph`` overrides are fully + deprecated and removed from the ABC: out-of-tree backends overriding + those must migrate to ``init_forward_metadata_out_graph(fb, in_capture)``. + """ + + def init_forward_metadata(self, forward_batch: ForwardBatch): + """Eager entry point. Default = ``_out_graph(fb) + _in_graph(fb)``. + + Backends may override to keep an independent eager body. + """ + self.init_forward_metadata_out_graph(forward_batch) + self.init_forward_metadata_in_graph(forward_batch) + + def init_forward_metadata_out_graph( + self, + forward_batch: ForwardBatch, + in_capture: bool = False, + ): + """Per-iter metadata prep — runs outside ``with graph.capture():``. + + Called at: + * capture: before ``with graph.capture():`` (caller passes + ``in_capture=True``). + * replay: before ``graph.replay()`` (``in_capture=False``). + * eager: via :py:meth:`init_forward_metadata` default wrapper + (``in_capture=False``). + + Backends read ``in_capture`` only when capture / replay bodies + diverge (e.g., snapshot metadata, swap buffer pointers, install + temp workspace). Host op / dynamic-shape / non-graph-recordable + logic lives here. + + Default: no-op. + """ + + def init_forward_metadata_in_graph(self, forward_batch: ForwardBatch): + """Graph-recordable static-shape GPU op. + + Runs inside ``with graph.capture():`` at capture time; recorded + ops auto-execute at replay via ``graph.replay()``. + + Lint contract for overrides: body must NOT call ``.item()`` / + ``.cpu()`` / ``.tolist()`` / dynamic-shape ``torch.empty()``. + Such ops belong in :py:meth:`init_forward_metadata_out_graph`; they + cannot be recorded into a cuda graph. + + Default: no-op. + """ # Opt out only when this backend never reads seq_lens_cpu / seq_lens_sum. needs_cpu_seq_lens: bool = True - @abstractmethod - def init_forward_metadata(self, forward_batch: ForwardBatch): - """Init the metadata for a forward pass.""" - raise NotImplementedError() - def init_cuda_graph_state(self, max_bs: int, max_num_tokens: int): """Init the global shared states for cuda graph.""" raise NotImplementedError() - def init_forward_metadata_capture_cuda_graph( - self, - bs: int, - num_tokens: int, - req_pool_indices: torch.Tensor, - seq_lens: torch.Tensor, - encoder_lens: Optional[torch.Tensor], - forward_mode: ForwardMode, - spec_info: Optional[SpecInput], - ): - """Init the metadata for a forward pass for capturing a cuda graph.""" - raise NotImplementedError() - - def init_forward_metadata_replay_cuda_graph( - self, - bs: int, - req_pool_indices: torch.Tensor, - seq_lens: torch.Tensor, - seq_lens_sum: int, - encoder_lens: Optional[torch.Tensor], - forward_mode: ForwardMode, - spec_info: Optional[SpecInput], - seq_lens_cpu: Optional[torch.Tensor], - ): - """Init the metadata for a forward pass for replaying a cuda graph.""" - raise NotImplementedError() - def get_cuda_graph_seq_len_fill_value(self): """Get the fill value for padded seq lens. Typically, it is 0 or 1.""" raise NotImplementedError() diff --git a/python/sglang/srt/layers/attention/cutlass_mla_backend.py b/python/sglang/srt/layers/attention/cutlass_mla_backend.py index 22781284d..a5c8d409f 100644 --- a/python/sglang/srt/layers/attention/cutlass_mla_backend.py +++ b/python/sglang/srt/layers/attention/cutlass_mla_backend.py @@ -14,13 +14,12 @@ import triton from sglang.srt.layers.attention.flashinfer_mla_backend import FlashInferMLAAttnBackend from sglang.srt.layers.attention.utils import create_flashmla_kv_indices_triton from sglang.srt.layers.dp_attention import get_attention_tp_size -from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode +from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.utils import is_cuda if TYPE_CHECKING: from sglang.srt.layers.radix_attention import RadixAttention from sglang.srt.model_executor.model_runner import ModelRunner - from sglang.srt.speculative.spec_info import SpecInput _is_cuda = is_cuda() if _is_cuda: @@ -79,6 +78,37 @@ class CutlassMLABackend(FlashInferMLAAttnBackend): self.q_data_type = model_runner.dtype self.kv_cache_dim = self.kv_lora_rank + self.qk_rope_head_dim + def init_forward_metadata_out_graph( + self, + forward_batch: ForwardBatch, + in_capture: bool = False, + ): + bs = forward_batch.batch_size + forward_mode = forward_batch.forward_mode + spec_info = forward_batch.spec_info + + if forward_mode.is_decode_or_idle() and spec_info is None: + create_flashmla_kv_indices_triton[(bs,)]( + self.req_to_token, + forward_batch.req_pool_indices[:bs], + forward_batch.seq_lens[:bs], + None, + self.cuda_graph_kv_indices, + self.req_to_token.stride(0), + self.cuda_graph_kv_indices.stride(0), + PAGED_SIZE=PAGE_SIZE, + ) + if in_capture: + max_seqlen_pad = self.cuda_graph_kv_indices.shape[1] + self.forward_metadata = CutlassMLADecodeMetadata( + self.cuda_graph_mla_workspace, + self.cuda_graph_kv_indices[:bs, :max_seqlen_pad], + ) + else: + super().init_forward_metadata_out_graph( + forward_batch, in_capture=in_capture + ) + def init_forward_metadata(self, forward_batch: ForwardBatch): bs = forward_batch.batch_size @@ -143,77 +173,6 @@ class CutlassMLABackend(FlashInferMLAAttnBackend): ) self.cuda_graph_kv_indices = cuda_graph_kv_indices - def init_forward_metadata_capture_cuda_graph( - self, - bs: int, - num_tokens: int, - req_pool_indices: torch.Tensor, - seq_lens: torch.Tensor, - encoder_lens: Optional[torch.Tensor], - forward_mode: ForwardMode, - spec_info: Optional[SpecInput], - ): - if forward_mode.is_decode_or_idle() and spec_info is None: - self.init_forward_metadata_replay_cuda_graph( - bs=bs, - req_pool_indices=req_pool_indices, - seq_lens=seq_lens, - seq_lens_sum=None, - encoder_lens=encoder_lens, - forward_mode=forward_mode, - spec_info=spec_info, - seq_lens_cpu=None, - ) - max_seqlen_pad = self.cuda_graph_kv_indices.shape[1] - self.forward_metadata = CutlassMLADecodeMetadata( - self.cuda_graph_mla_workspace, - self.cuda_graph_kv_indices[:bs, :max_seqlen_pad], - ) - else: - super().init_forward_metadata_capture_cuda_graph( - bs, - num_tokens, - req_pool_indices, - seq_lens, - encoder_lens, - forward_mode, - spec_info, - ) - - def init_forward_metadata_replay_cuda_graph( - self, - bs: int, - req_pool_indices: torch.Tensor, - seq_lens: torch.Tensor, - seq_lens_sum: int, - encoder_lens: Optional[torch.Tensor], - forward_mode: ForwardMode, - spec_info: Optional[SpecInput], - seq_lens_cpu: Optional[torch.Tensor], - ): - if forward_mode.is_decode_or_idle(): - create_flashmla_kv_indices_triton[(bs,)]( - self.req_to_token, - req_pool_indices[:bs], - seq_lens[:bs], - None, - self.cuda_graph_kv_indices, - self.req_to_token.stride(0), - self.cuda_graph_kv_indices.stride(0), - PAGED_SIZE=PAGE_SIZE, - ) - else: - super().init_forward_metadata_replay_cuda_graph( - bs, - req_pool_indices, - seq_lens, - seq_lens_sum, - encoder_lens, - forward_mode, - spec_info, - seq_lens_cpu, - ) - def get_cuda_graph_seq_len_fill_value(self): return 1 diff --git a/python/sglang/srt/layers/attention/deepseek_v4_backend.py b/python/sglang/srt/layers/attention/deepseek_v4_backend.py index d2d4ed487..f9b3e973f 100644 --- a/python/sglang/srt/layers/attention/deepseek_v4_backend.py +++ b/python/sglang/srt/layers/attention/deepseek_v4_backend.py @@ -54,7 +54,6 @@ from sglang.srt.layers.dp_attention import ( from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.speculative.eagle_utils import per_step_draft_out_cache_loc -from sglang.srt.speculative.spec_info import SpecInput from sglang.srt.utils import ceil_align from sglang.srt.utils.common import is_sm120_supported @@ -387,7 +386,6 @@ class DeepseekV4AttnBackend( DSV4RawVerifyMetadata, DSV4RawDecodeMetadata, ] = None - self._replay_forward_batch: Optional[ForwardBatch] = None # FIXME: out-of-band def _move_to_device(self, x: List[int]) -> torch.Tensor: pin_tensor = torch.tensor(x, dtype=torch.int32, pin_memory=True) @@ -671,6 +669,136 @@ class DeepseekV4AttnBackend( use_prefill_cuda_graph=use_prefill_cuda_graph, ) + def init_forward_metadata_in_graph(self, forward_batch: ForwardBatch) -> None: + # Upgrade Raw->Full so the c4/c128 compress + core_attn + indexer + # materialization is recorded inside the cuda graph; a no-op (Full + # already) when PREP_IN_CUDA_GRAPH=0. + if isinstance(self.forward_metadata, DSV4RawVerifyMetadata): + self.forward_metadata = self.make_forward_metadata_from_raw_verify( + raw_metadata=self.forward_metadata, + ) + elif isinstance(self.forward_metadata, DSV4RawDecodeMetadata): + self.forward_metadata = self.make_forward_metadata_from_raw_decode( + raw_metadata=self.forward_metadata, + ) + + def init_forward_metadata_out_graph( + self, + forward_batch: ForwardBatch, + in_capture: bool = False, + ) -> None: + bucket = _GraphBucket.of(forward_batch.forward_mode) + bs = forward_batch.batch_size + req_pool_indices = forward_batch.req_pool_indices + seq_lens = forward_batch.seq_lens + + if in_capture: + # Captured graph does no real cache writes, so synthesize a dummy + # out_cache_loc per bucket (replay supplies the real value). + assert req_pool_indices.size(0) == bs + assert seq_lens.size(0) == bs + num_tokens = forward_batch.positions.numel() + if bucket == _GraphBucket.DECODE_OR_IDLE: + out_cache_loc = torch.zeros_like(seq_lens) + elif bucket == _GraphBucket.TARGET_VERIFY: + out_cache_loc = torch.zeros(num_tokens, **self.cuda_int32_kwargs) + else: + out_cache_loc = None + actual_forward_mode = forward_batch.forward_mode + seq_lens_sum = int(seq_lens.sum().item()) + seq_lens_cpu = seq_lens.cpu() + else: + out_cache_loc = forward_batch.out_cache_loc + actual_forward_mode = getattr( + forward_batch, "actual_forward_mode", forward_batch.forward_mode + ) + seq_lens_sum = forward_batch.seq_lens_sum + seq_lens_cpu = forward_batch.seq_lens_cpu + + if actual_forward_mode == ForwardMode.IDLE: + logger.debug( + f"[IDLE replay] bs={bs}, " + f"local_seq_lens_len={len(seq_lens)}, " + f"has_graph={bs in self.cuda_graph_metadata_of_bucket_and_bs[_GraphBucket.DECODE_OR_IDLE]}" + ) + device = seq_lens.device + seq_lens = torch.ones(bs, dtype=seq_lens.dtype, device=device) + seq_lens_cpu = torch.ones(bs, dtype=torch.int64) + seq_lens_sum = bs + req_pool_indices = torch.zeros( + bs, dtype=req_pool_indices.dtype, device=device + ) + out_cache_loc = torch.zeros(bs, dtype=torch.int64, device=device) + + assert seq_lens_cpu is not None + seq_lens = seq_lens[:bs] + seq_lens_cpu = seq_lens_cpu[:bs] + req_pool_indices = req_pool_indices[:bs] + + actual_max_seq_len = seq_lens_cpu.max().item() + chosen_max_seq_len = self.MAX_SEQ_LEN_FOR_CAPTURE + assert actual_max_seq_len <= chosen_max_seq_len + + if bucket == _GraphBucket.DECODE_OR_IDLE: + assert out_cache_loc is not None + assert len(out_cache_loc.shape) == 1, f"{out_cache_loc.shape=}" + out_cache_loc_padded = torch.nn.functional.pad( + out_cache_loc, + pad=(0, bs - len(out_cache_loc)), + mode="constant", + value=0, + ) + temp_metadata = self.init_forward_metadata_decode( + max_seq_len=chosen_max_seq_len, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + out_cache_loc=out_cache_loc_padded, + ) + elif bucket == _GraphBucket.TARGET_VERIFY: + assert out_cache_loc is not None + num_tokens_v = self.speculative_num_draft_tokens * bs + out_cache_loc_padded = torch.nn.functional.pad( + out_cache_loc, + pad=(0, num_tokens_v - len(out_cache_loc)), + mode="constant", + value=0, + ) + temp_metadata = self.init_forward_metadata_target_verify( + max_seq_len=chosen_max_seq_len, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + out_cache_loc=out_cache_loc_padded, + use_prefill_cuda_graph=True, + ) + elif bucket == _GraphBucket.DRAFT_EXTEND: + num_tokens_per_bs = self.draft_extend_num_tokens_per_bs + temp_metadata = self.init_forward_metadata_draft_extend( + max_seq_len=chosen_max_seq_len, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + seq_lens_cpu=seq_lens_cpu.tolist(), + num_tokens_per_bs=num_tokens_per_bs, + use_prefill_cuda_graph=True, + ) + else: + raise NotImplementedError + + self.replay_cuda_graph_metadata_from( + bs=bs, temp_metadata=temp_metadata, bucket=bucket + ) + + if in_capture: + # Preserve _current_capture_raw for on_after_cuda_graph_warmup + metadata = self.forward_metadata + self._current_capture_raw = ( + metadata + if isinstance( + metadata, + (DSV4RawDecodeMetadata, DSV4RawVerifyMetadata), + ) + else None + ) + def init_forward_metadata(self, forward_batch: ForwardBatch) -> None: if self.mtp_enabled and forward_batch.forward_mode.is_idle(): return @@ -686,8 +814,7 @@ class DeepseekV4AttnBackend( if forward_batch.forward_mode.is_decode_or_idle(): # DSv4 bakes this step's KV write target (c4/c128) into metadata, - # so slice the shared multi-step out_cache_loc now rather than at - # forward time. + # so slice the shared multi-step out_cache_loc now, not at forward time. out_cache_loc = forward_batch.out_cache_loc if self.topk > 0 and self.speculative_num_steps > 1: out_cache_loc = per_step_draft_out_cache_loc( @@ -734,154 +861,24 @@ class DeepseekV4AttnBackend( raise NotImplementedError(f"unsupported mode {forward_batch.forward_mode=}") self.forward_metadata = metadata + self.init_forward_metadata_in_graph(forward_batch) def init_cuda_graph_state(self, max_bs: int, max_num_tokens: int) -> None: self.cuda_graph_metadata_of_bucket_and_bs: Dict[ _GraphBucket, Dict[ int, - Union[DSV4Metadata, DSV4RawDecodeMetadata, DSV4RawVerifyMetadata], + Union[ + DSV4Metadata, + DSV4RawDecodeMetadata, + DSV4RawVerifyMetadata, + ], ], ] = {bucket: {} for bucket in _GraphBucket} self.draft_extend_num_tokens_per_bs = ( max_num_tokens // max_bs if max_bs > 0 else 1 ) - def init_forward_metadata_capture_cuda_graph( - self, - bs: int, - num_tokens: int, - req_pool_indices: torch.Tensor, - seq_lens: torch.Tensor, - encoder_lens: Optional[torch.Tensor], - forward_mode: ForwardMode, - spec_info: Optional[SpecInput], - ) -> None: - from types import SimpleNamespace - - assert req_pool_indices.size(0) == bs - assert seq_lens.size(0) == bs - - bucket = _GraphBucket.of(forward_mode) - if bucket == _GraphBucket.DECODE_OR_IDLE: - dummy_cache_loc = torch.zeros_like(seq_lens) - elif bucket == _GraphBucket.TARGET_VERIFY: - dummy_cache_loc = torch.zeros(num_tokens, **self.cuda_int32_kwargs) - else: - dummy_cache_loc = None - - self._replay_forward_batch = SimpleNamespace( - out_cache_loc=dummy_cache_loc, - forward_mode=forward_mode, - ) - self.init_forward_metadata_replay_cuda_graph( - bs=bs, - req_pool_indices=req_pool_indices, - seq_lens=seq_lens, - seq_lens_sum=int(seq_lens.sum().item()), - encoder_lens=encoder_lens, - forward_mode=forward_mode, - spec_info=spec_info, - seq_lens_cpu=seq_lens.cpu(), - ) - # Preserve _current_capture_raw for on_after_cuda_graph_warmup - metadata = self.forward_metadata - self._current_capture_raw = ( - metadata - if isinstance(metadata, (DSV4RawDecodeMetadata, DSV4RawVerifyMetadata)) - else None - ) - - def init_forward_metadata_replay_cuda_graph( - self, - bs: int, - req_pool_indices: torch.Tensor, - seq_lens: torch.Tensor, - seq_lens_sum: int, - encoder_lens: Optional[torch.Tensor], - forward_mode: ForwardMode, - spec_info: Optional[SpecInput], - seq_lens_cpu: Optional[torch.Tensor], - ) -> None: - bucket = _GraphBucket.of(forward_mode) - - # FIXME: see cuda_graph_runner — this attribute is set out-of-band. - fb = self._replay_forward_batch - out_cache_loc = fb.out_cache_loc - actual_forward_mode = fb.forward_mode - - if actual_forward_mode == ForwardMode.IDLE: - logger.debug( - f"[IDLE replay] bs={bs}, " - f"local_seq_lens_len={len(seq_lens)}, " - f"has_graph={bs in self.cuda_graph_metadata_of_bucket_and_bs[_GraphBucket.DECODE_OR_IDLE]}" - ) - device = seq_lens.device - seq_lens = torch.ones(bs, dtype=seq_lens.dtype, device=device) - seq_lens_cpu = torch.ones(bs, dtype=torch.int64) - seq_lens_sum = bs - req_pool_indices = torch.zeros( - bs, dtype=req_pool_indices.dtype, device=device - ) - out_cache_loc = torch.zeros(bs, dtype=torch.int64, device=device) - - assert seq_lens_cpu is not None - seq_lens = seq_lens[:bs] - seq_lens_cpu = seq_lens_cpu[:bs] - req_pool_indices = req_pool_indices[:bs] - - actual_max_seq_len = seq_lens_cpu.max().item() - chosen_max_seq_len = self.MAX_SEQ_LEN_FOR_CAPTURE - assert actual_max_seq_len <= chosen_max_seq_len - - if bucket == _GraphBucket.DECODE_OR_IDLE: - assert out_cache_loc is not None - assert len(out_cache_loc.shape) == 1, f"{out_cache_loc.shape=}" - out_cache_loc_padded = torch.nn.functional.pad( - out_cache_loc, - pad=(0, bs - len(out_cache_loc)), - mode="constant", - value=0, - ) - temp_metadata = self.init_forward_metadata_decode( - max_seq_len=chosen_max_seq_len, - req_pool_indices=req_pool_indices, - seq_lens=seq_lens, - out_cache_loc=out_cache_loc_padded, - ) - elif bucket == _GraphBucket.TARGET_VERIFY: - assert out_cache_loc is not None - num_tokens = self.speculative_num_draft_tokens * bs - out_cache_loc_padded = torch.nn.functional.pad( - out_cache_loc, - pad=(0, num_tokens - len(out_cache_loc)), - mode="constant", - value=0, - ) - temp_metadata = self.init_forward_metadata_target_verify( - max_seq_len=chosen_max_seq_len, - req_pool_indices=req_pool_indices, - seq_lens=seq_lens, - out_cache_loc=out_cache_loc_padded, - use_prefill_cuda_graph=True, - ) - elif bucket == _GraphBucket.DRAFT_EXTEND: - num_tokens_per_bs = self.draft_extend_num_tokens_per_bs - temp_metadata = self.init_forward_metadata_draft_extend( - max_seq_len=chosen_max_seq_len, - req_pool_indices=req_pool_indices, - seq_lens=seq_lens, - seq_lens_cpu=seq_lens_cpu.tolist(), - num_tokens_per_bs=num_tokens_per_bs, - use_prefill_cuda_graph=True, - ) - else: - raise NotImplementedError - - self.replay_cuda_graph_metadata_from( - bs=bs, temp_metadata=temp_metadata, bucket=bucket - ) - def replay_cuda_graph_metadata_from( self, bs: int, @@ -938,24 +935,6 @@ class DeepseekV4AttnBackend( cache_nope_fp8_rope_bf16_pack=swa_k_pack, ) - def _maybe_upgrade_forward_metadata(self) -> None: - # With SGLANG_PREP_IN_CUDA_GRAPH=1, init_forward_metadata_* - # returns a Raw metadata that only carries a few tensors. The - # full DSV4Metadata (including c4/c128 compress + core_attn + - # indexer metadata) must be materialized before any caller that - # touches those fields. For 1.6T the first two layers have - # compress_ratio=128, so forward_core_compressor / forward_c4_indexer - # can fire before attn_backend.forward(), and must trigger the - # upgrade themselves. - if isinstance(self.forward_metadata, DSV4RawVerifyMetadata): - self.forward_metadata = self.make_forward_metadata_from_raw_verify( - raw_metadata=self.forward_metadata, - ) - elif isinstance(self.forward_metadata, DSV4RawDecodeMetadata): - self.forward_metadata = self.make_forward_metadata_from_raw_decode( - raw_metadata=self.forward_metadata, - ) - def forward( self, q: torch.Tensor, @@ -968,8 +947,6 @@ class DeepseekV4AttnBackend( attn_sink: Optional[torch.Tensor] = None, **_, ) -> torch.Tensor: - self._maybe_upgrade_forward_metadata() - if self.mtp_enabled and forward_batch.forward_mode.is_idle(): return q.new_empty(q.shape[0], q.shape[1], layer.v_head_dim) @@ -1233,6 +1210,52 @@ class DeepseekV4MultiStepBackend(DeepseekV4AttnBackend): ) ) + def init_forward_metadata_in_graph(self, forward_batch: ForwardBatch) -> None: + for attn_backend in self.attn_backends: + attn_backend.init_forward_metadata_in_graph(forward_batch) + + def init_forward_metadata_out_graph( + self, + forward_batch: ForwardBatch, + in_capture: bool = False, + ): + from types import SimpleNamespace + + inner_fb = SimpleNamespace( + batch_size=forward_batch.batch_size, + forward_mode=ForwardMode.DECODE, + # Propagate the real runtime mode so inner backends can detect IDLE + # and apply their idle substitution. + actual_forward_mode=getattr( + forward_batch, "actual_forward_mode", forward_batch.forward_mode + ), + input_ids=getattr(forward_batch, "input_ids", None), + positions=getattr(forward_batch, "positions", None), + req_pool_indices=forward_batch.req_pool_indices, + seq_lens=forward_batch.seq_lens, + seq_lens_sum=forward_batch.seq_lens_sum, + seq_lens_cpu=forward_batch.seq_lens_cpu, + encoder_lens=None, + out_cache_loc=getattr(forward_batch, "out_cache_loc", None), + spec_info=forward_batch.spec_info, + ) + if in_capture: + for i in range(self.speculative_num_steps): + self.attn_backends[i].init_forward_metadata_out_graph( + inner_fb, in_capture=True + ) + else: + if self.speculative_num_steps == 1: + return + self.attn_backends[0].init_forward_metadata_out_graph(inner_fb) + temp_metadata = self.attn_backends[0].forward_metadata + for i in range(1, self.speculative_num_steps - 1): + self.attn_backends[i].replay_cuda_graph_metadata_from( + bs=forward_batch.batch_size, + temp_metadata=temp_metadata, + bucket=_GraphBucket.DECODE_OR_IDLE, + ) + def init_forward_metadata(self, forward_batch: ForwardBatch): for i in range(self.speculative_num_steps - 1): self.attn_backends[i].init_forward_metadata(forward_batch) @@ -1241,49 +1264,10 @@ class DeepseekV4MultiStepBackend(DeepseekV4AttnBackend): for i in range(self.speculative_num_steps): self.attn_backends[i].init_cuda_graph_state(max_bs, max_num_tokens) - def init_forward_metadata_capture_cuda_graph(self, forward_batch: ForwardBatch): - for i in range(self.speculative_num_steps): - self.attn_backends[i].init_forward_metadata_capture_cuda_graph( - forward_batch.batch_size, - forward_batch.batch_size * self.topk, - forward_batch.req_pool_indices, - forward_batch.seq_lens, - encoder_lens=None, - forward_mode=ForwardMode.DECODE, - spec_info=forward_batch.spec_info, - ) - def on_after_cuda_graph_warmup(self): for backend in self.attn_backends: backend.on_after_cuda_graph_warmup() - def init_forward_metadata_replay_cuda_graph( - self, forward_batch: ForwardBatch, bs: int - ): - if self.speculative_num_steps == 1: - return - - self.attn_backends[0]._replay_forward_batch = forward_batch - self.attn_backends[0].init_forward_metadata_replay_cuda_graph( - bs=bs, - req_pool_indices=forward_batch.req_pool_indices, - seq_lens=forward_batch.seq_lens, - seq_lens_sum=forward_batch.seq_lens_sum, - encoder_lens=None, - forward_mode=ForwardMode.DECODE, - spec_info=forward_batch.spec_info, - seq_lens_cpu=forward_batch.seq_lens_cpu, - ) - self.attn_backends[0]._replay_forward_batch = None - temp_metadata = self.attn_backends[0].forward_metadata - - for i in range(1, self.speculative_num_steps - 1): - self.attn_backends[i].replay_cuda_graph_metadata_from( - bs=bs, - temp_metadata=temp_metadata, - bucket=_GraphBucket.DECODE_OR_IDLE, - ) - def _pad_tensor_to_size(tensor: torch.Tensor, size: int, *, value: int = 0): if value == 0: diff --git a/python/sglang/srt/layers/attention/deepseek_v4_backend_hip_radix.py b/python/sglang/srt/layers/attention/deepseek_v4_backend_hip_radix.py index f400764d5..c42c8d678 100644 --- a/python/sglang/srt/layers/attention/deepseek_v4_backend_hip_radix.py +++ b/python/sglang/srt/layers/attention/deepseek_v4_backend_hip_radix.py @@ -52,7 +52,6 @@ from sglang.srt.layers.dp_attention import ( from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.speculative.eagle_utils import per_step_draft_out_cache_loc -from sglang.srt.speculative.spec_info import SpecInput from sglang.srt.utils import ceil_align if TYPE_CHECKING: @@ -381,7 +380,6 @@ class DeepseekV4HipRadixBackend( DSV4RawVerifyMetadata, DSV4RawDecodeMetadata, ] = None - self._replay_forward_batch: Optional[ForwardBatch] = None # FIXME: out-of-band def _move_to_device(self, x: List[int]) -> torch.Tensor: pin_tensor = torch.tensor(x, dtype=torch.int32, pin_memory=True) @@ -661,6 +659,133 @@ class DeepseekV4HipRadixBackend( use_prefill_cuda_graph=use_prefill_cuda_graph, ) + def init_forward_metadata_in_graph(self, forward_batch: ForwardBatch) -> None: + # Upgrade Raw->Full so the c4/c128 compress + core_attn + indexer + # materialization is recorded inside the cuda graph; a no-op (Full + # already) when PREP_IN_CUDA_GRAPH=0. + if isinstance(self.forward_metadata, DSV4RawVerifyMetadata): + self.forward_metadata = self.make_forward_metadata_from_raw_verify( + raw_metadata=self.forward_metadata, + ) + elif isinstance(self.forward_metadata, DSV4RawDecodeMetadata): + self.forward_metadata = self.make_forward_metadata_from_raw_decode( + raw_metadata=self.forward_metadata, + ) + + def init_forward_metadata_out_graph( + self, + forward_batch: ForwardBatch, + in_capture: bool = False, + ) -> None: + bucket = _GraphBucket.of(forward_batch.forward_mode) + bs = forward_batch.batch_size + req_pool_indices = forward_batch.req_pool_indices + seq_lens = forward_batch.seq_lens + + if in_capture: + assert req_pool_indices.size(0) == bs + assert seq_lens.size(0) == bs + num_tokens = forward_batch.positions.numel() + if bucket == _GraphBucket.DECODE_OR_IDLE: + out_cache_loc = torch.zeros_like(seq_lens) + elif bucket == _GraphBucket.TARGET_VERIFY: + out_cache_loc = torch.zeros(num_tokens, **self.cuda_int32_kwargs) + else: + out_cache_loc = None + actual_forward_mode = forward_batch.forward_mode + seq_lens_sum = int(seq_lens.sum().item()) + seq_lens_cpu = seq_lens.cpu() + else: + out_cache_loc = forward_batch.out_cache_loc + actual_forward_mode = getattr( + forward_batch, "actual_forward_mode", forward_batch.forward_mode + ) + seq_lens_sum = forward_batch.seq_lens_sum + seq_lens_cpu = forward_batch.seq_lens_cpu + + if actual_forward_mode == ForwardMode.IDLE: + logger.debug( + f"[IDLE replay] bs={bs}, " + f"local_seq_lens_len={len(seq_lens)}, " + f"has_graph={bs in self.cuda_graph_metadata_of_bucket_and_bs[_GraphBucket.DECODE_OR_IDLE]}" + ) + device = seq_lens.device + seq_lens = torch.ones(bs, dtype=seq_lens.dtype, device=device) + seq_lens_cpu = torch.ones(bs, dtype=torch.int64) + seq_lens_sum = bs + req_pool_indices = torch.zeros( + bs, dtype=req_pool_indices.dtype, device=device + ) + out_cache_loc = torch.zeros(bs, dtype=torch.int64, device=device) + + assert seq_lens_cpu is not None + seq_lens = seq_lens[:bs] + seq_lens_cpu = seq_lens_cpu[:bs] + req_pool_indices = req_pool_indices[:bs] + + actual_max_seq_len = seq_lens_cpu.max().item() + chosen_max_seq_len = self.MAX_SEQ_LEN_FOR_CAPTURE + assert actual_max_seq_len <= chosen_max_seq_len + + if bucket == _GraphBucket.DECODE_OR_IDLE: + assert out_cache_loc is not None + assert len(out_cache_loc.shape) == 1, f"{out_cache_loc.shape=}" + out_cache_loc_padded = torch.nn.functional.pad( + out_cache_loc, + pad=(0, bs - len(out_cache_loc)), + mode="constant", + value=0, + ) + temp_metadata = self.init_forward_metadata_decode( + max_seq_len=chosen_max_seq_len, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + out_cache_loc=out_cache_loc_padded, + ) + elif bucket == _GraphBucket.TARGET_VERIFY: + assert out_cache_loc is not None + num_tokens_v = self.speculative_num_draft_tokens * bs + out_cache_loc_padded = torch.nn.functional.pad( + out_cache_loc, + pad=(0, num_tokens_v - len(out_cache_loc)), + mode="constant", + value=0, + ) + temp_metadata = self.init_forward_metadata_target_verify( + max_seq_len=chosen_max_seq_len, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + out_cache_loc=out_cache_loc_padded, + use_prefill_cuda_graph=True, + ) + elif bucket == _GraphBucket.DRAFT_EXTEND: + num_tokens_per_bs = self.draft_extend_num_tokens_per_bs + temp_metadata = self.init_forward_metadata_draft_extend( + max_seq_len=chosen_max_seq_len, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + seq_lens_cpu=seq_lens_cpu.tolist(), + num_tokens_per_bs=num_tokens_per_bs, + use_prefill_cuda_graph=True, + ) + else: + raise NotImplementedError + + self.replay_cuda_graph_metadata_from( + bs=bs, temp_metadata=temp_metadata, bucket=bucket + ) + + if in_capture: + metadata = self.forward_metadata + self._current_capture_raw = ( + metadata + if isinstance( + metadata, + (DSV4RawDecodeMetadata, DSV4RawVerifyMetadata), + ) + else None + ) + def init_forward_metadata(self, forward_batch: ForwardBatch) -> None: if self.mtp_enabled and forward_batch.forward_mode.is_idle(): return @@ -676,8 +801,7 @@ class DeepseekV4HipRadixBackend( if forward_batch.forward_mode.is_decode_or_idle(): # DSv4 bakes this step's KV write target (c4/c128) into metadata, - # so slice the shared multi-step out_cache_loc now rather than at - # forward time. + # so slice the shared multi-step out_cache_loc now, not at forward time. out_cache_loc = forward_batch.out_cache_loc if self.topk > 0 and self.speculative_num_steps > 1: out_cache_loc = per_step_draft_out_cache_loc( @@ -725,154 +849,24 @@ class DeepseekV4HipRadixBackend( raise NotImplementedError(f"unsupported mode {forward_batch.forward_mode=}") self.forward_metadata = metadata + self.init_forward_metadata_in_graph(forward_batch) def init_cuda_graph_state(self, max_bs: int, max_num_tokens: int) -> None: self.cuda_graph_metadata_of_bucket_and_bs: Dict[ _GraphBucket, Dict[ int, - Union[DSV4Metadata, DSV4RawDecodeMetadata, DSV4RawVerifyMetadata], + Union[ + DSV4Metadata, + DSV4RawDecodeMetadata, + DSV4RawVerifyMetadata, + ], ], ] = {bucket: {} for bucket in _GraphBucket} self.draft_extend_num_tokens_per_bs = ( max_num_tokens // max_bs if max_bs > 0 else 1 ) - def init_forward_metadata_capture_cuda_graph( - self, - bs: int, - num_tokens: int, - req_pool_indices: torch.Tensor, - seq_lens: torch.Tensor, - encoder_lens: Optional[torch.Tensor], - forward_mode: ForwardMode, - spec_info: Optional[SpecInput], - ) -> None: - from types import SimpleNamespace - - assert req_pool_indices.size(0) == bs - assert seq_lens.size(0) == bs - - bucket = _GraphBucket.of(forward_mode) - if bucket == _GraphBucket.DECODE_OR_IDLE: - dummy_cache_loc = torch.zeros_like(seq_lens) - elif bucket == _GraphBucket.TARGET_VERIFY: - dummy_cache_loc = torch.zeros(num_tokens, **self.cuda_int32_kwargs) - else: - dummy_cache_loc = None - - self._replay_forward_batch = SimpleNamespace( - out_cache_loc=dummy_cache_loc, - forward_mode=forward_mode, - ) - self.init_forward_metadata_replay_cuda_graph( - bs=bs, - req_pool_indices=req_pool_indices, - seq_lens=seq_lens, - seq_lens_sum=int(seq_lens.sum().item()), - encoder_lens=encoder_lens, - forward_mode=forward_mode, - spec_info=spec_info, - seq_lens_cpu=seq_lens.cpu(), - ) - # Preserve _current_capture_raw for on_after_cuda_graph_warmup - metadata = self.forward_metadata - self._current_capture_raw = ( - metadata - if isinstance(metadata, (DSV4RawDecodeMetadata, DSV4RawVerifyMetadata)) - else None - ) - - def init_forward_metadata_replay_cuda_graph( - self, - bs: int, - req_pool_indices: torch.Tensor, - seq_lens: torch.Tensor, - seq_lens_sum: int, - encoder_lens: Optional[torch.Tensor], - forward_mode: ForwardMode, - spec_info: Optional[SpecInput], - seq_lens_cpu: Optional[torch.Tensor], - ) -> None: - bucket = _GraphBucket.of(forward_mode) - - # FIXME: see cuda_graph_runner — this attribute is set out-of-band. - fb = self._replay_forward_batch - out_cache_loc = fb.out_cache_loc - actual_forward_mode = fb.forward_mode - - if actual_forward_mode == ForwardMode.IDLE: - logger.debug( - f"[IDLE replay] bs={bs}, " - f"local_seq_lens_len={len(seq_lens)}, " - f"has_graph={bs in self.cuda_graph_metadata_of_bucket_and_bs[_GraphBucket.DECODE_OR_IDLE]}" - ) - device = seq_lens.device - seq_lens = torch.ones(bs, dtype=seq_lens.dtype, device=device) - seq_lens_cpu = torch.ones(bs, dtype=torch.int64) - seq_lens_sum = bs - req_pool_indices = torch.zeros( - bs, dtype=req_pool_indices.dtype, device=device - ) - out_cache_loc = torch.zeros(bs, dtype=torch.int64, device=device) - - assert seq_lens_cpu is not None - seq_lens = seq_lens[:bs] - seq_lens_cpu = seq_lens_cpu[:bs] - req_pool_indices = req_pool_indices[:bs] - - actual_max_seq_len = seq_lens_cpu.max().item() - chosen_max_seq_len = self.MAX_SEQ_LEN_FOR_CAPTURE - assert actual_max_seq_len <= chosen_max_seq_len - - if bucket == _GraphBucket.DECODE_OR_IDLE: - assert out_cache_loc is not None - assert len(out_cache_loc.shape) == 1, f"{out_cache_loc.shape=}" - out_cache_loc_padded = torch.nn.functional.pad( - out_cache_loc, - pad=(0, bs - len(out_cache_loc)), - mode="constant", - value=0, - ) - temp_metadata = self.init_forward_metadata_decode( - max_seq_len=chosen_max_seq_len, - req_pool_indices=req_pool_indices, - seq_lens=seq_lens, - out_cache_loc=out_cache_loc_padded, - ) - elif bucket == _GraphBucket.TARGET_VERIFY: - assert out_cache_loc is not None - num_tokens = self.speculative_num_draft_tokens * bs - out_cache_loc_padded = torch.nn.functional.pad( - out_cache_loc, - pad=(0, num_tokens - len(out_cache_loc)), - mode="constant", - value=0, - ) - temp_metadata = self.init_forward_metadata_target_verify( - max_seq_len=chosen_max_seq_len, - req_pool_indices=req_pool_indices, - seq_lens=seq_lens, - out_cache_loc=out_cache_loc_padded, - use_prefill_cuda_graph=True, - ) - elif bucket == _GraphBucket.DRAFT_EXTEND: - num_tokens_per_bs = self.draft_extend_num_tokens_per_bs - temp_metadata = self.init_forward_metadata_draft_extend( - max_seq_len=chosen_max_seq_len, - req_pool_indices=req_pool_indices, - seq_lens=seq_lens, - seq_lens_cpu=seq_lens_cpu.tolist(), - num_tokens_per_bs=num_tokens_per_bs, - use_prefill_cuda_graph=True, - ) - else: - raise NotImplementedError - - self.replay_cuda_graph_metadata_from( - bs=bs, temp_metadata=temp_metadata, bucket=bucket - ) - def replay_cuda_graph_metadata_from( self, bs: int, @@ -929,24 +923,6 @@ class DeepseekV4HipRadixBackend( cache_nope_fp8_rope_bf16_pack=swa_k_pack, ) - def _maybe_upgrade_forward_metadata(self) -> None: - # With SGLANG_PREP_IN_CUDA_GRAPH=1, init_forward_metadata_* - # returns a Raw metadata that only carries a few tensors. The - # full DSV4Metadata (including c4/c128 compress + core_attn + - # indexer metadata) must be materialized before any caller that - # touches those fields. For 1.6T the first two layers have - # compress_ratio=128, so forward_core_compressor / forward_c4_indexer - # can fire before attn_backend.forward(), and must trigger the - # upgrade themselves. - if isinstance(self.forward_metadata, DSV4RawVerifyMetadata): - self.forward_metadata = self.make_forward_metadata_from_raw_verify( - raw_metadata=self.forward_metadata, - ) - elif isinstance(self.forward_metadata, DSV4RawDecodeMetadata): - self.forward_metadata = self.make_forward_metadata_from_raw_decode( - raw_metadata=self.forward_metadata, - ) - def forward( self, q: torch.Tensor, @@ -959,8 +935,6 @@ class DeepseekV4HipRadixBackend( attn_sink: Optional[torch.Tensor] = None, **_, ) -> torch.Tensor: - self._maybe_upgrade_forward_metadata() - if self.mtp_enabled and forward_batch.forward_mode.is_idle(): return q.new_empty(q.shape[0], q.shape[1], layer.v_head_dim) @@ -1212,6 +1186,52 @@ class DeepseekV4MultiStepBackend(DeepseekV4HipRadixBackend): ) ) + def init_forward_metadata_in_graph(self, forward_batch: ForwardBatch) -> None: + for attn_backend in self.attn_backends: + attn_backend.init_forward_metadata_in_graph(forward_batch) + + def init_forward_metadata_out_graph( + self, + forward_batch: ForwardBatch, + in_capture: bool = False, + ): + from types import SimpleNamespace + + inner_fb = SimpleNamespace( + batch_size=forward_batch.batch_size, + forward_mode=ForwardMode.DECODE, + # Propagate the real runtime mode so inner backends can detect IDLE + # and apply their idle substitution. + actual_forward_mode=getattr( + forward_batch, "actual_forward_mode", forward_batch.forward_mode + ), + input_ids=getattr(forward_batch, "input_ids", None), + positions=getattr(forward_batch, "positions", None), + req_pool_indices=forward_batch.req_pool_indices, + seq_lens=forward_batch.seq_lens, + seq_lens_sum=forward_batch.seq_lens_sum, + seq_lens_cpu=forward_batch.seq_lens_cpu, + encoder_lens=None, + out_cache_loc=getattr(forward_batch, "out_cache_loc", None), + spec_info=forward_batch.spec_info, + ) + if in_capture: + for i in range(self.speculative_num_steps): + self.attn_backends[i].init_forward_metadata_out_graph( + inner_fb, in_capture=True + ) + else: + if self.speculative_num_steps == 1: + return + self.attn_backends[0].init_forward_metadata_out_graph(inner_fb) + temp_metadata = self.attn_backends[0].forward_metadata + for i in range(1, self.speculative_num_steps - 1): + self.attn_backends[i].replay_cuda_graph_metadata_from( + bs=forward_batch.batch_size, + temp_metadata=temp_metadata, + bucket=_GraphBucket.DECODE_OR_IDLE, + ) + def init_forward_metadata(self, forward_batch: ForwardBatch): for i in range(self.speculative_num_steps - 1): self.attn_backends[i].init_forward_metadata(forward_batch) @@ -1220,49 +1240,10 @@ class DeepseekV4MultiStepBackend(DeepseekV4HipRadixBackend): for i in range(self.speculative_num_steps): self.attn_backends[i].init_cuda_graph_state(max_bs, max_num_tokens) - def init_forward_metadata_capture_cuda_graph(self, forward_batch: ForwardBatch): - for i in range(self.speculative_num_steps): - self.attn_backends[i].init_forward_metadata_capture_cuda_graph( - forward_batch.batch_size, - forward_batch.batch_size * self.topk, - forward_batch.req_pool_indices, - forward_batch.seq_lens, - encoder_lens=None, - forward_mode=ForwardMode.DECODE, - spec_info=forward_batch.spec_info, - ) - def on_after_cuda_graph_warmup(self): for backend in self.attn_backends: backend.on_after_cuda_graph_warmup() - def init_forward_metadata_replay_cuda_graph( - self, forward_batch: ForwardBatch, bs: int - ): - if self.speculative_num_steps == 1: - return - - self.attn_backends[0]._replay_forward_batch = forward_batch - self.attn_backends[0].init_forward_metadata_replay_cuda_graph( - bs=bs, - req_pool_indices=forward_batch.req_pool_indices, - seq_lens=forward_batch.seq_lens, - seq_lens_sum=forward_batch.seq_lens_sum, - encoder_lens=None, - forward_mode=ForwardMode.DECODE, - spec_info=forward_batch.spec_info, - seq_lens_cpu=forward_batch.seq_lens_cpu, - ) - self.attn_backends[0]._replay_forward_batch = None - temp_metadata = self.attn_backends[0].forward_metadata - - for i in range(1, self.speculative_num_steps - 1): - self.attn_backends[i].replay_cuda_graph_metadata_from( - bs=bs, - temp_metadata=temp_metadata, - bucket=_GraphBucket.DECODE_OR_IDLE, - ) - def _pad_tensor_to_size(tensor: torch.Tensor, size: int, *, value: int = 0): if value == 0: diff --git a/python/sglang/srt/layers/attention/dsa_backend.py b/python/sglang/srt/layers/attention/dsa_backend.py index 11b305516..e3b8eebd5 100644 --- a/python/sglang/srt/layers/attention/dsa_backend.py +++ b/python/sglang/srt/layers/attention/dsa_backend.py @@ -408,6 +408,25 @@ class DeepseekSparseAttnBackend( ) return page_table[:, strided_indices] // page_size + def init_forward_metadata_out_graph( + self, + forward_batch: ForwardBatch, + in_capture: bool = False, + ): + seq_lens_cpu = ( + forward_batch.seq_lens.cpu() if in_capture else forward_batch.seq_lens_cpu + ) + self._apply_cuda_graph_metadata( + bs=forward_batch.batch_size, + req_pool_indices=forward_batch.req_pool_indices, + seq_lens=forward_batch.seq_lens, + seq_lens_cpu=seq_lens_cpu, + forward_mode=forward_batch.forward_mode, + spec_info=forward_batch.spec_info, + out_cache_loc=getattr(forward_batch, "out_cache_loc", None), + actual_forward_mode=getattr(forward_batch, "actual_forward_mode", None), + ) + def init_forward_metadata(self, forward_batch: ForwardBatch): """Init the metadata for a forward pass.""" batch_size = forward_batch.batch_size @@ -971,42 +990,23 @@ class DeepseekSparseAttnBackend( self.decode_cuda_graph_metadata[bs] = metadata self.forward_metadata = metadata - def init_forward_metadata_capture_cuda_graph( - self, - bs: int, - num_tokens: int, - req_pool_indices: torch.Tensor, - seq_lens: torch.Tensor, - encoder_lens: Optional[torch.Tensor], - forward_mode: ForwardMode, - spec_info: Optional[SpecInput], - ): - """Initialize forward metadata for capturing CUDA graph.""" - self.init_forward_metadata_replay_cuda_graph( - bs=bs, - req_pool_indices=req_pool_indices, - seq_lens=seq_lens, - seq_lens_sum=None, - encoder_lens=encoder_lens, - forward_mode=forward_mode, - spec_info=spec_info, - seq_lens_cpu=seq_lens.cpu(), - ) - - def init_forward_metadata_replay_cuda_graph( + def _apply_cuda_graph_metadata( self, bs: int, req_pool_indices: torch.Tensor, seq_lens: torch.Tensor, - seq_lens_sum: int, - encoder_lens: Optional[torch.Tensor], + seq_lens_cpu: torch.Tensor, forward_mode: ForwardMode, spec_info: Optional[SpecInput], - seq_lens_cpu: Optional[torch.Tensor], out_cache_loc: Optional[torch.Tensor] = None, actual_forward_mode: Optional[ForwardMode] = None, ): - """Initialize forward metadata for replaying CUDA graph.""" + """Shared capture+replay body for the cuda-graph init path. + + Public entry: :py:meth:`init_forward_metadata_out_graph`. Spec runners + also call this directly via _apply_cuda_graph_metadata when they + need to pass out_cache_loc / actual_forward_mode explicitly. + """ assert seq_lens_cpu is not None if bs not in self.decode_cuda_graph_metadata: @@ -2370,21 +2370,26 @@ class DeepseekSparseAttnMultiStepBackend: for i in range(self.speculative_num_steps - 1): self.attn_backends[i].init_cuda_graph_state(max_bs, max_num_tokens) - def init_forward_metadata_capture_cuda_graph(self, forward_batch: ForwardBatch): - for i in range(self.speculative_num_steps - 1): - self.attn_backends[i].init_forward_metadata_capture_cuda_graph( - forward_batch.batch_size, - forward_batch.batch_size * self.topk, - forward_batch.req_pool_indices, - forward_batch.seq_lens, - encoder_lens=None, - forward_mode=ForwardMode.DECODE, - spec_info=forward_batch.spec_info, - ) - - def init_forward_metadata_replay_cuda_graph( - self, forward_batch: ForwardBatch, bs: int + def init_forward_metadata_out_graph( + self, + forward_batch: ForwardBatch, + in_capture: bool = False, ): + from sglang.srt.model_executor.forward_batch_info import build_inner_fb_view + + if in_capture: + inner_fb = build_inner_fb_view( + forward_batch, + bs=forward_batch.batch_size, + forward_mode=ForwardMode.DECODE, + ) + for i in range(self.speculative_num_steps - 1): + self.attn_backends[i].init_forward_metadata_out_graph( + inner_fb, in_capture=True + ) + return + + bs = forward_batch.batch_size if envs.SGLANG_DSA_ENABLE_MTP_PRECOMPUTE_METADATA.get(): # Precompute metadata once (shared across all backends) precomputed = self.attn_backends[0]._precompute_replay_metadata( @@ -2542,20 +2547,21 @@ class DeepseekSparseAttnMultiStepBackend: forward_mode=ForwardMode.DECODE, ) else: - # Fallback: compute metadata separately for each backend for i in range(self.speculative_num_steps - 1): - self.attn_backends[i].init_forward_metadata_replay_cuda_graph( + self.attn_backends[i]._apply_cuda_graph_metadata( bs=bs, req_pool_indices=forward_batch.req_pool_indices, seq_lens=forward_batch.seq_lens, - seq_lens_sum=forward_batch.seq_lens_sum, - encoder_lens=None, + seq_lens_cpu=forward_batch.seq_lens_cpu, forward_mode=ForwardMode.DECODE, spec_info=forward_batch.spec_info, - seq_lens_cpu=forward_batch.seq_lens_cpu, out_cache_loc=None, ) + def init_forward_metadata_in_graph(self, forward_batch: ForwardBatch) -> None: + for i in range(self.speculative_num_steps - 1): + self.attn_backends[i].init_forward_metadata_in_graph(forward_batch) + # Backward-compat aliases (deprecated: use DSA class names) DeepseekSparseAttnBackend = DeepseekSparseAttnBackend diff --git a/python/sglang/srt/layers/attention/dsv4/compressor.py b/python/sglang/srt/layers/attention/dsv4/compressor.py index f20de6bfb..dd7dbf1e1 100644 --- a/python/sglang/srt/layers/attention/dsv4/compressor.py +++ b/python/sglang/srt/layers/attention/dsv4/compressor.py @@ -57,9 +57,6 @@ class CompressorBackendMixin: assert isinstance(metadata, FusedCompressMetadata) return metadata - def _maybe_upgrade_forward_metadata(self) -> None: - pass - def forward_compress( self, *, @@ -153,11 +150,6 @@ class CompressorBackendMixin: ) -> None: if forward_batch.forward_mode.is_idle(): return - # PREP_IN_CG lazy upgrade: the concrete backend (DeepseekV4AttnBackend) - # owns this helper. MQALayer._forward_prepare calls us before - # attn_backend.forward(), so Raw -> DSV4Metadata must happen here too - # (e.g. 1.6T layer 0 has compress_ratio=128 and needs cX_compress_metadata). - self._maybe_upgrade_forward_metadata() token_to_kv_pool = self.token_to_kv_pool if TYPE_CHECKING: assert isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool) @@ -187,8 +179,6 @@ class CompressorBackendMixin: compressor: Compressor, ) -> None: assert is_overlap_compress(compressor.ratio) - # PREP_IN_CG lazy upgrade (see forward_core_compressor for rationale). - self._maybe_upgrade_forward_metadata() token_to_kv_pool = self.token_to_kv_pool if TYPE_CHECKING: assert isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool) diff --git a/python/sglang/srt/layers/attention/dsv4/compressor_v2.py b/python/sglang/srt/layers/attention/dsv4/compressor_v2.py index bd5293900..b48bd985c 100644 --- a/python/sglang/srt/layers/attention/dsv4/compressor_v2.py +++ b/python/sglang/srt/layers/attention/dsv4/compressor_v2.py @@ -404,9 +404,6 @@ class CompressorBackendMixin: super().__init__() self.forward_metadata: DSV4Metadata - # NOTE: Will be overridden - def _maybe_upgrade_forward_metadata(self): ... - def _get_paged_compress_metadata(self, compress_ratio: int) -> CompressMetadata: attr_name = f"c{compress_ratio}_compress_metadata" return getattr(self.forward_metadata, attr_name) @@ -483,7 +480,6 @@ class CompressorBackendMixin: if forward_batch.forward_mode.is_idle(): return - self._maybe_upgrade_forward_metadata() token_to_kv_pool = self.token_to_kv_pool token_to_kv_pool = cast("DeepSeekV4TokenToKVPool", token_to_kv_pool) kv_score_input = compressor.compute_kv_score(x, forward_batch) diff --git a/python/sglang/srt/layers/attention/dsv4/indexer.py b/python/sglang/srt/layers/attention/dsv4/indexer.py index 4c9e21bf8..933185031 100644 --- a/python/sglang/srt/layers/attention/dsv4/indexer.py +++ b/python/sglang/srt/layers/attention/dsv4/indexer.py @@ -443,9 +443,6 @@ class C4IndexerBackendMixin: ) -> None: if forward_batch.forward_mode.is_idle(): return - # PREP_IN_CG lazy upgrade: this runs from MQALayer._forward_prepare, - # before attn_backend.forward() would trigger the upgrade. - self._maybe_upgrade_forward_metadata() token_to_kv_pool = self.token_to_kv_pool if TYPE_CHECKING: diff --git a/python/sglang/srt/layers/attention/dual_chunk_flashattention_backend.py b/python/sglang/srt/layers/attention/dual_chunk_flashattention_backend.py index 00dc6e169..29dea5eec 100644 --- a/python/sglang/srt/layers/attention/dual_chunk_flashattention_backend.py +++ b/python/sglang/srt/layers/attention/dual_chunk_flashattention_backend.py @@ -21,7 +21,9 @@ from sglang.jit_kernel.flash_attention import ( ) from sglang.srt.distributed.parallel_state import get_tensor_model_parallel_rank from sglang.srt.layers.attention.base_attn_backend import AttentionBackend -from sglang.srt.layers.attention.flashattention_backend import FlashAttentionMetadata +from sglang.srt.layers.attention.flashattention_backend import ( + FlashAttentionMetadata, +) from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode if TYPE_CHECKING: @@ -172,6 +174,35 @@ class DualChunkFlashAttentionBackend(AttentionBackend): end_head = start_head + self.num_heads return [layer_sparse_attention_config[i] for i in range(start_head, end_head)] + def init_forward_metadata_out_graph( + self, + forward_batch: ForwardBatch, + in_capture: bool = False, + ): + bs = forward_batch.batch_size + req_pool_indices = forward_batch.req_pool_indices + seq_lens = forward_batch.seq_lens + forward_mode = forward_batch.forward_mode + + if in_capture: + self._bind_metadata_buffers(bs, req_pool_indices, forward_mode) + + self._apply_cuda_graph_metadata( + bs=bs, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + forward_mode=forward_mode, + ) + + if in_capture and forward_mode.is_decode_or_idle(): + # Restore max_seq_len scalars — replay sets actual values but CUDA + # graph needs the safe upper bound baked in at capture time. + md = self.forward_metadata + md.max_seq_len = self.max_context_len + md.max_seq_len_intra = self.max_context_len + md.max_seq_len_succ = self.max_context_len + md.max_seq_len_inter = self.max_context_len + def init_forward_metadata(self, forward_batch: ForwardBatch): """Initialize forward metadata hence all layers in the forward pass can reuse it.""" @@ -577,49 +608,17 @@ class DualChunkFlashAttentionBackend(AttentionBackend): self.forward_metadata = metadata - def init_forward_metadata_capture_cuda_graph( + def _apply_cuda_graph_metadata( self, bs: int, - num_tokens: int, req_pool_indices: torch.Tensor, seq_lens: torch.Tensor, - encoder_lens: Optional[torch.Tensor], forward_mode: ForwardMode, - spec_info: Optional[None], ): - self._bind_metadata_buffers(bs, req_pool_indices, forward_mode) - self.init_forward_metadata_replay_cuda_graph( - bs=bs, - req_pool_indices=req_pool_indices, - seq_lens=seq_lens, - seq_lens_sum=None, - encoder_lens=encoder_lens, - forward_mode=forward_mode, - spec_info=spec_info, - seq_lens_cpu=seq_lens.cpu(), - ) - # Restore max_seq_len scalars — replay sets actual values but CUDA graph - # needs the safe upper bound baked in at capture time. - if forward_mode.is_decode_or_idle(): - md = self.forward_metadata - md.max_seq_len = self.max_context_len - md.max_seq_len_intra = self.max_context_len - md.max_seq_len_succ = self.max_context_len - md.max_seq_len_inter = self.max_context_len + """Shared capture+replay body for the cuda-graph init path. - def init_forward_metadata_replay_cuda_graph( - self, - bs: int, - req_pool_indices: torch.Tensor, - seq_lens: torch.Tensor, - seq_lens_sum: int, - encoder_lens: Optional[torch.Tensor], - forward_mode: ForwardMode, - spec_info: Optional[None], - seq_lens_cpu: Optional[torch.Tensor], - out_cache_loc: torch.Tensor = None, - ): - """Initialize forward metadata for replaying CUDA graph.""" + Public entry: :py:meth:`init_forward_metadata_out_graph`. + """ assert forward_mode.is_decode() seq_lens = seq_lens[:bs] req_pool_indices = req_pool_indices[:bs] diff --git a/python/sglang/srt/layers/attention/flashattention_backend.py b/python/sglang/srt/layers/attention/flashattention_backend.py index dce97babf..a5b637c1f 100644 --- a/python/sglang/srt/layers/attention/flashattention_backend.py +++ b/python/sglang/srt/layers/attention/flashattention_backend.py @@ -273,6 +273,102 @@ class FlashAttentionBackend(AttentionBackend): num_splits=self.num_splits, ) + def init_forward_metadata_out_graph( + self, + forward_batch: ForwardBatch, + in_capture: bool = False, + ): + bs = forward_batch.batch_size + req_pool_indices = forward_batch.req_pool_indices + seq_lens = forward_batch.seq_lens + encoder_lens = forward_batch.encoder_lens + forward_mode = forward_batch.forward_mode + spec_info = forward_batch.spec_info + out_cache_loc = getattr(forward_batch, "out_cache_loc", None) + + if in_capture: + num_tokens = forward_batch.positions.numel() + seq_lens_cpu = seq_lens.cpu() + self._bind_metadata_buffers( + bs, + num_tokens, + encoder_lens, + forward_mode, + spec_info, + seq_lens.device, + ) + + if ( + forward_mode.is_decode_or_idle() + and spec_info is not None + and self.topk > 1 + ): + # topk>1 draft decode: replay needs out_cache_loc which capture doesn't have; + # set forward_metadata directly and let actual CUDA graph replay fill data. + self.forward_metadata = self.draft_decode_metadata_topk_normal[bs] + self.forward_metadata_spec_decode_expand = ( + self.draft_decode_metadata_topk_expand[bs] + ) + return + + if forward_mode.is_target_verify() and self.topk > 1: + # topk>1 target verify: replay needs spec_info.positions and .custom_mask + # which are not populated at capture time. + self.forward_metadata = self.target_verify_metadata_topk_normal[bs] + self.forward_metadata_spec_decode_expand = ( + self.target_verify_metadata_topk_expand[bs] + ) + return + + self._apply_cuda_graph_metadata( + bs=bs, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + seq_lens_sum=None, + encoder_lens=encoder_lens, + forward_mode=forward_mode, + spec_info=spec_info, + seq_lens_cpu=seq_lens_cpu, + out_cache_loc=out_cache_loc, + ) + + if forward_mode.is_decode_or_idle() and spec_info is None: + # Local attention and scheduler metadata require capture-time slice sizing. + # Both depend on data already filled by replay above. + metadata = self.decode_cuda_graph_metadata[bs] + self._maybe_update_local_attn_metadata_for_capture(metadata, bs) + if self._sched_meta_buf is not None: + sched = self._compute_scheduler_metadata( + bs, + max(metadata.max_seq_len_k, 1), + metadata.cache_seqlens_int32, + metadata.cu_seqlens_q, + ) + if sched is not None: + n = sched.shape[0] + self._sched_meta_buf[:n] = sched + self._sched_meta_buf[n:] = 0 + metadata.scheduler_metadata = self._sched_meta_buf[:n] + + if forward_mode.is_draft_extend(include_v2=True): + # CUDA graph bakes max_seq_len_q as a constant. replay() sets it to + # max(num_accept_tokens_cpu) which is None/empty at capture time, + # falling back to 1. Restore the correct upper bound so the kernel + # sees num_tokens_per_bs (not 1) for all replays of this graph. + self.forward_metadata.max_seq_len_q = num_tokens // bs + else: + self._apply_cuda_graph_metadata( + bs=bs, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + seq_lens_sum=forward_batch.seq_lens_sum, + encoder_lens=encoder_lens, + forward_mode=forward_mode, + spec_info=spec_info, + seq_lens_cpu=forward_batch.seq_lens_cpu, + out_cache_loc=out_cache_loc, + ) + def init_forward_metadata(self, forward_batch: ForwardBatch): """Initialize forward metadata hence all layers in the forward pass can reuse it.""" metadata = FlashAttentionMetadata() @@ -1901,77 +1997,7 @@ class FlashAttentionBackend(AttentionBackend): return metadata, metadata_expand - def init_forward_metadata_capture_cuda_graph( - self, - bs: int, - num_tokens: int, - req_pool_indices: torch.Tensor, - seq_lens: torch.Tensor, - encoder_lens: Optional[torch.Tensor], - forward_mode: ForwardMode, - spec_info: Optional[SpecInput], - ): - """Initialize forward metadata for capturing CUDA graph.""" - seq_lens_cpu = seq_lens.cpu() - self._bind_metadata_buffers( - bs, num_tokens, encoder_lens, forward_mode, spec_info, seq_lens.device - ) - - if forward_mode.is_decode_or_idle() and spec_info is not None and self.topk > 1: - # topk>1 draft decode: replay needs out_cache_loc which capture doesn't have; - # set forward_metadata directly and let actual CUDA graph replay fill data. - self.forward_metadata = self.draft_decode_metadata_topk_normal[bs] - self.forward_metadata_spec_decode_expand = ( - self.draft_decode_metadata_topk_expand[bs] - ) - return - - if forward_mode.is_target_verify() and self.topk > 1: - # topk>1 target verify: replay needs spec_info.positions and .custom_mask - # which are not populated at capture time. - self.forward_metadata = self.target_verify_metadata_topk_normal[bs] - self.forward_metadata_spec_decode_expand = ( - self.target_verify_metadata_topk_expand[bs] - ) - return - - self.init_forward_metadata_replay_cuda_graph( - bs=bs, - req_pool_indices=req_pool_indices, - seq_lens=seq_lens, - seq_lens_sum=None, - encoder_lens=encoder_lens, - forward_mode=forward_mode, - spec_info=spec_info, - seq_lens_cpu=seq_lens_cpu, - ) - - if forward_mode.is_decode_or_idle() and spec_info is None: - # Local attention and scheduler metadata require capture-time slice sizing. - # Both depend on data already filled by replay above. - metadata = self.decode_cuda_graph_metadata[bs] - self._maybe_update_local_attn_metadata_for_capture(metadata, bs) - if self._sched_meta_buf is not None: - sched = self._compute_scheduler_metadata( - bs, - max(metadata.max_seq_len_k, 1), - metadata.cache_seqlens_int32, - metadata.cu_seqlens_q, - ) - if sched is not None: - n = sched.shape[0] - self._sched_meta_buf[:n] = sched - self._sched_meta_buf[n:] = 0 - metadata.scheduler_metadata = self._sched_meta_buf[:n] - - if forward_mode.is_draft_extend(include_v2=True): - # CUDA graph bakes max_seq_len_q as a constant. replay() sets it to - # max(num_accept_tokens_cpu) which is None/empty at capture time, - # falling back to 1. Restore the correct upper bound so the kernel - # sees num_tokens_per_bs (not 1) for all replays of this graph. - self.forward_metadata.max_seq_len_q = num_tokens // bs - - def init_forward_metadata_replay_cuda_graph( + def _apply_cuda_graph_metadata( self, bs: int, req_pool_indices: torch.Tensor, @@ -1983,7 +2009,13 @@ class FlashAttentionBackend(AttentionBackend): seq_lens_cpu: Optional[torch.Tensor], out_cache_loc: Optional[torch.Tensor] = None, ): - """Initialize forward metadata for replaying CUDA graph.""" + """Shared capture+replay body for the cuda-graph init path. + + Public entry: :py:meth:`init_forward_metadata_out_graph`. This helper + formerly lived as the legacy ``init_forward_metadata_replay_cuda_graph``; + the capture path used to wrap it. Both legacy method overrides + are gone. + """ seq_lens = seq_lens[:bs] seq_lens_cpu = seq_lens_cpu[:bs] req_pool_indices = req_pool_indices[:bs] @@ -2324,6 +2356,16 @@ class FlashAttentionBackend(AttentionBackend): ) metadata.page_table[:, :max_seq_pages].copy_(page_indices // self.page_size) + else: + raise ValueError( + f"FA3 `_apply_cuda_graph_metadata` only supports the modes the " + f"full cuda-graph runner captures (decode / idle / target_verify " + f"/ draft_extend / draft_extend_v2). Got {forward_mode=}. " + f"Piecewise / breakable capture must route through " + f"`init_forward_metadata(fb)` (the eager entry) instead of " + f"`init_forward_metadata_out_graph(fb, in_capture=True)`." + ) + if encoder_lens is not None: # Per-request varlen encoder support (e.g. MossVL different images). metadata.encoder_max_seq_len_k = int(encoder_lens.max().item()) @@ -2353,7 +2395,10 @@ class FlashAttentionBackend(AttentionBackend): return 1 def _maybe_init_local_attn_metadata( - self, forwardbatch: ForwardBatch, metadata: FlashAttentionMetadata, device + self, + forwardbatch: ForwardBatch, + metadata: FlashAttentionMetadata, + device, ): """Centralized utility to initialize local_attn_metadata if chunked attention is enabled.""" if not self.has_local_attention: @@ -2628,45 +2673,33 @@ class FlashAttentionMultiStepBackend: for i in range(self.speculative_num_steps - 1): self.attn_backends[i].init_cuda_graph_state(max_bs, max_num_tokens) - def init_forward_metadata_capture_cuda_graph( + def init_forward_metadata_out_graph( self, forward_batch: ForwardBatch, + in_capture: bool = False, ): + from sglang.srt.model_executor.forward_batch_info import build_inner_fb_view + assert forward_batch.spec_info is not None assert forward_batch.spec_info.is_draft_input() - for i in range(self.speculative_num_steps - 1): - self.attn_backends[i].init_forward_metadata_capture_cuda_graph( - forward_batch.batch_size, - forward_batch.batch_size * self.topk, - forward_batch.req_pool_indices, - forward_batch.seq_lens, - encoder_lens=forward_batch.encoder_lens, - forward_mode=ForwardMode.DECODE, - spec_info=forward_batch.spec_info, - ) - - def init_forward_metadata_replay_cuda_graph( - self, forward_batch: ForwardBatch, bs: int - ): - assert forward_batch.spec_info is not None - assert forward_batch.spec_info.is_draft_input() - + inner_fb = build_inner_fb_view( + forward_batch, + bs=forward_batch.batch_size, + forward_mode=ForwardMode.DECODE, + encoder_lens=forward_batch.encoder_lens, + ) for i in range(self.speculative_num_steps - 1): # TODO: incrementally update the metadata for the later steps, # so that they do not need to recompute everything from scratch. - self.attn_backends[i].init_forward_metadata_replay_cuda_graph( - bs, - forward_batch.req_pool_indices, - forward_batch.seq_lens, - forward_batch.seq_lens_sum, - encoder_lens=forward_batch.encoder_lens, - forward_mode=ForwardMode.DECODE, - spec_info=forward_batch.spec_info, - seq_lens_cpu=forward_batch.seq_lens_cpu, - out_cache_loc=forward_batch.out_cache_loc, + self.attn_backends[i].init_forward_metadata_out_graph( + inner_fb, in_capture=in_capture ) + def init_forward_metadata_in_graph(self, forward_batch: ForwardBatch) -> None: + for i in range(self.speculative_num_steps - 1): + self.attn_backends[i].init_forward_metadata_in_graph(forward_batch) + @torch.compile(dynamic=True, backend=get_compiler_backend()) def draft_decode_set_expand_metadata( diff --git a/python/sglang/srt/layers/attention/flashinfer_backend.py b/python/sglang/srt/layers/attention/flashinfer_backend.py index d055d5cff..93943f533 100644 --- a/python/sglang/srt/layers/attention/flashinfer_backend.py +++ b/python/sglang/srt/layers/attention/flashinfer_backend.py @@ -461,6 +461,69 @@ class FlashInferAttnBackend(AttentionBackend): ), ) + def init_forward_metadata_out_graph( + self, + forward_batch: ForwardBatch, + in_capture: bool = False, + ): + bs = forward_batch.batch_size + req_pool_indices = forward_batch.req_pool_indices + seq_lens = forward_batch.seq_lens + seq_lens_cpu = forward_batch.seq_lens_cpu + seq_lens_sum = forward_batch.seq_lens_sum + encoder_lens = forward_batch.encoder_lens + forward_mode = forward_batch.forward_mode + spec_info = forward_batch.spec_info + + if in_capture: + num_tokens = forward_batch.positions.numel() + self._prepare_cuda_graph_metadata(bs, num_tokens, forward_mode, spec_info) + + if forward_mode.is_decode_or_idle(): + self.indices_updater_decode.update( + req_pool_indices[:bs], + seq_lens[:bs], + seq_lens_cpu[:bs] if seq_lens_cpu is not None else None, + seq_lens_sum, + decode_wrappers=self.decode_cuda_graph_metadata[bs], + encoder_lens=encoder_lens[:bs] if encoder_lens is not None else None, + spec_info=spec_info, + fixed_split_size=None, + disable_split_kv=self.disable_cuda_graph_kv_split, + ) + elif forward_mode.is_target_verify() or forward_mode.is_draft_extend(): + self.indices_updater_prefill.update( + req_pool_indices[:bs], + seq_lens[:bs], + seq_lens_cpu[:bs] if seq_lens_cpu is not None else None, + seq_lens_sum, + prefix_lens=None, + prefill_wrappers=self.prefill_cuda_graph_metadata[bs], + use_ragged=False, + encoder_lens=encoder_lens[:bs] if encoder_lens is not None else None, + spec_info=spec_info, + ) + elif forward_mode.is_dllm_extend(): + self.indices_updater_prefill.update( + req_pool_indices[:bs], + seq_lens[:bs], + seq_lens_cpu[:bs] if seq_lens_cpu is not None else None, + seq_lens_sum, + prefix_lens=seq_lens - self.dllm_config.block_size, + prefill_wrappers=self.prefill_cuda_graph_metadata[bs], + use_ragged=not self.use_paged, + encoder_lens=encoder_lens[:bs] if encoder_lens is not None else None, + spec_info=None, + ) + else: + raise ValueError("Invalid forward mode") + + if in_capture and forward_mode.is_decode_or_idle(): + # fast_decode_plan needs _cached_module from the initial begin_forward + # above, so install it only after that first plan has run. + for w in self.decode_cuda_graph_metadata[bs]: + w.begin_forward = partial(fast_decode_plan, w) + def init_forward_metadata(self, forward_batch: ForwardBatch): if forward_batch.forward_mode.is_decode_or_idle(): self.indices_updater_decode.update( @@ -662,85 +725,6 @@ class FlashInferAttnBackend(AttentionBackend): else: raise ValueError(f"Invalid mode: {forward_mode=}") - def init_forward_metadata_capture_cuda_graph( - self, - bs: int, - num_tokens: int, - req_pool_indices: torch.Tensor, - seq_lens: torch.Tensor, - encoder_lens: Optional[torch.Tensor], - forward_mode: ForwardMode, - spec_info: Optional[SpecInput], - ): - seq_lens_sum = seq_lens.sum().item() - seq_lens_cpu = seq_lens.cpu() - self._prepare_cuda_graph_metadata(bs, num_tokens, forward_mode, spec_info) - self.init_forward_metadata_replay_cuda_graph( - bs=bs, - req_pool_indices=req_pool_indices, - seq_lens=seq_lens, - seq_lens_sum=seq_lens_sum, - encoder_lens=encoder_lens, - forward_mode=forward_mode, - spec_info=spec_info, - seq_lens_cpu=seq_lens_cpu, - ) - # fast_decode_plan requires _cached_module set by the initial full - # begin_forward call above; install it only after that first plan runs. - if forward_mode.is_decode_or_idle(): - for w in self.decode_cuda_graph_metadata[bs]: - w.begin_forward = partial(fast_decode_plan, w) - - def init_forward_metadata_replay_cuda_graph( - self, - bs: int, - req_pool_indices: torch.Tensor, - seq_lens: torch.Tensor, - seq_lens_sum: int, - encoder_lens: Optional[torch.Tensor], - forward_mode: ForwardMode, - spec_info: Optional[SpecInput], - seq_lens_cpu: Optional[torch.Tensor], - ): - if forward_mode.is_decode_or_idle(): - self.indices_updater_decode.update( - req_pool_indices[:bs], - seq_lens[:bs], - seq_lens_cpu[:bs] if seq_lens_cpu is not None else None, - seq_lens_sum, - decode_wrappers=self.decode_cuda_graph_metadata[bs], - encoder_lens=encoder_lens[:bs] if encoder_lens is not None else None, - spec_info=spec_info, - fixed_split_size=None, - disable_split_kv=self.disable_cuda_graph_kv_split, - ) - elif forward_mode.is_target_verify() or forward_mode.is_draft_extend(): - self.indices_updater_prefill.update( - req_pool_indices[:bs], - seq_lens[:bs], - seq_lens_cpu[:bs] if seq_lens_cpu is not None else None, - seq_lens_sum, - prefix_lens=None, - prefill_wrappers=self.prefill_cuda_graph_metadata[bs], - use_ragged=False, - encoder_lens=encoder_lens[:bs] if encoder_lens is not None else None, - spec_info=spec_info, - ) - elif forward_mode.is_dllm_extend(): - self.indices_updater_prefill.update( - req_pool_indices[:bs], - seq_lens[:bs], - seq_lens_cpu[:bs] if seq_lens_cpu is not None else None, - seq_lens_sum, - prefix_lens=seq_lens - self.dllm_config.block_size, - prefill_wrappers=self.prefill_cuda_graph_metadata[bs], - use_ragged=not self.use_paged, - encoder_lens=encoder_lens[:bs] if encoder_lens is not None else None, - spec_info=None, - ) - else: - raise ValueError("Invalid forward mode") - def get_cuda_graph_seq_len_fill_value(self): return 1 @@ -1668,37 +1652,27 @@ class FlashInferMultiStepDraftBackend: max_bs, max_num_tokens, kv_indices_buf=self.cuda_graph_kv_indices[i] ) - def init_forward_metadata_capture_cuda_graph(self, forward_batch: ForwardBatch): - def call_fn(i, forward_batch): - self.attn_backends[i].init_forward_metadata_capture_cuda_graph( - forward_batch.batch_size, - forward_batch.batch_size * self.topk, - forward_batch.req_pool_indices, - forward_batch.seq_lens, - encoder_lens=None, - forward_mode=ForwardMode.DECODE, - spec_info=forward_batch.spec_info, - ) - - self.common_template(forward_batch, self.cuda_graph_kv_indices, call_fn) - - def init_forward_metadata_replay_cuda_graph( - self, forward_batch: ForwardBatch, bs: int + def init_forward_metadata_out_graph( + self, + forward_batch: ForwardBatch, + in_capture: bool = False, ): - def call_fn(i, forward_batch): - self.attn_backends[i].init_forward_metadata_replay_cuda_graph( - bs, - forward_batch.req_pool_indices, - forward_batch.seq_lens, - seq_lens_sum=-1, - encoder_lens=None, - forward_mode=ForwardMode.DECODE, - spec_info=forward_batch.spec_info, - seq_lens_cpu=forward_batch.seq_lens_cpu, + from sglang.srt.model_executor.forward_batch_info import build_inner_fb_view + + bs = forward_batch.batch_size + + def call_fn(i, fb): + inner_fb = build_inner_fb_view(fb, bs=bs, forward_mode=ForwardMode.DECODE) + self.attn_backends[i].init_forward_metadata_out_graph( + inner_fb, in_capture=in_capture ) self.common_template(forward_batch, self.cuda_graph_kv_indices, call_fn) + def init_forward_metadata_in_graph(self, forward_batch: ForwardBatch) -> None: + for attn_backend in self.attn_backends: + attn_backend.init_forward_metadata_in_graph(forward_batch) + def should_use_tensor_core( kv_cache_dtype: torch.dtype, diff --git a/python/sglang/srt/layers/attention/flashinfer_mla_backend.py b/python/sglang/srt/layers/attention/flashinfer_mla_backend.py index cda57acb4..a1b6f2981 100644 --- a/python/sglang/srt/layers/attention/flashinfer_mla_backend.py +++ b/python/sglang/srt/layers/attention/flashinfer_mla_backend.py @@ -289,6 +289,80 @@ class FlashInferMLAAttnBackend(AttentionBackend): self.decode_cuda_graph_metadata = {} self.prefill_cuda_graph_metadata = {} # For verify + def init_forward_metadata_out_graph( + self, + forward_batch: ForwardBatch, + in_capture: bool = False, + ): + bs = forward_batch.batch_size + req_pool_indices = forward_batch.req_pool_indices + seq_lens = forward_batch.seq_lens + forward_mode = forward_batch.forward_mode + spec_info = forward_batch.spec_info + + if in_capture: + num_tokens = forward_batch.positions.numel() + seq_lens_sum = seq_lens.sum().item() + seq_lens_cpu = seq_lens.cpu() + + if forward_mode.is_decode_or_idle(): + decode_wrapper = BatchMLAPagedAttentionWrapper( + self.workspace_buffer, + use_cuda_graph=True, + qo_indptr=self.cuda_graph_qo_indptr[: num_tokens + 1], + kv_indptr=self.cuda_graph_kv_indptr[: num_tokens + 1], + kv_indices=self.cuda_graph_kv_indices, + kv_len_arr=self.cuda_graph_kv_lens[:num_tokens], + backend="auto", + ) + self.indices_updater_decode.update( + req_pool_indices, + seq_lens, + seq_lens_sum, + decode_wrapper=decode_wrapper, + init_metadata_replay=False, + spec_info=spec_info, + ) + self.decode_cuda_graph_metadata[bs] = decode_wrapper + self.forward_metadata = DecodeMetadata(decode_wrapper) + # fast_mla_decode_plan needs _cached_module from the initial + # begin_forward above, so install it only after that call completes. + decode_wrapper.plan = partial(fast_mla_decode_plan, decode_wrapper) + elif forward_mode.is_target_verify() or forward_mode.is_draft_extend(): + prefill_wrapper = BatchMLAPagedAttentionWrapper( + self.workspace_buffer, + use_cuda_graph=True, + qo_indptr=self.cuda_graph_qo_indptr[: bs + 1], + kv_indptr=self.cuda_graph_kv_indptr[: bs + 1], + kv_indices=self.cuda_graph_kv_indices, + kv_len_arr=self.cuda_graph_kv_lens[:bs], + backend="auto", + ) + self.prefill_cuda_graph_metadata[bs] = prefill_wrapper + self.forward_metadata = PrefillMetadata(prefill_wrapper, False) + else: + raise ValueError(f"Invalid mode: {forward_mode=}") + + self._apply_cuda_graph_metadata( + bs=bs, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + seq_lens_sum=seq_lens_sum, + forward_mode=forward_mode, + spec_info=spec_info, + seq_lens_cpu=seq_lens_cpu, + ) + else: + self._apply_cuda_graph_metadata( + bs=bs, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + seq_lens_sum=forward_batch.seq_lens_sum, + forward_mode=forward_mode, + spec_info=spec_info, + seq_lens_cpu=forward_batch.seq_lens_cpu, + ) + def init_forward_metadata(self, forward_batch: ForwardBatch): if forward_batch.forward_mode.is_decode_or_idle(): self.indices_updater_decode.update( @@ -374,83 +448,20 @@ class FlashInferMLAAttnBackend(AttentionBackend): "kv_indices": self.cuda_graph_kv_indices, } - def init_forward_metadata_capture_cuda_graph( - self, - bs: int, - num_tokens: int, - req_pool_indices: torch.Tensor, - seq_lens: torch.Tensor, - encoder_lens: Optional[torch.Tensor], - forward_mode: ForwardMode, - spec_info: Optional[SpecInput], - ): - seq_lens_sum = seq_lens.sum().item() - seq_lens_cpu = seq_lens.cpu() - - if forward_mode.is_decode_or_idle(): - # Decode: create wrapper, run the initial full begin_forward (False), - # then install the fast plan. After that, call replay so the - # data-update path (update(True)) is also exercised during capture. - decode_wrapper = BatchMLAPagedAttentionWrapper( - self.workspace_buffer, - use_cuda_graph=True, - qo_indptr=self.cuda_graph_qo_indptr[: num_tokens + 1], - kv_indptr=self.cuda_graph_kv_indptr[: num_tokens + 1], - kv_indices=self.cuda_graph_kv_indices, - kv_len_arr=self.cuda_graph_kv_lens[:num_tokens], - backend="auto", - ) - self.indices_updater_decode.update( - req_pool_indices, - seq_lens, - seq_lens_sum, - decode_wrapper=decode_wrapper, - init_metadata_replay=False, - spec_info=spec_info, - ) - self.decode_cuda_graph_metadata[bs] = decode_wrapper - self.forward_metadata = DecodeMetadata(decode_wrapper) - # fast_mla_decode_plan requires _cached_module set by the initial - # begin_forward above; install it only after that call completes. - decode_wrapper.plan = partial(fast_mla_decode_plan, decode_wrapper) - elif forward_mode.is_target_verify() or forward_mode.is_draft_extend(): - # Prefill: create wrapper and store — replay handles the update call. - prefill_wrapper = BatchMLAPagedAttentionWrapper( - self.workspace_buffer, - use_cuda_graph=True, - qo_indptr=self.cuda_graph_qo_indptr[: bs + 1], - kv_indptr=self.cuda_graph_kv_indptr[: bs + 1], - kv_indices=self.cuda_graph_kv_indices, - kv_len_arr=self.cuda_graph_kv_lens[:bs], - backend="auto", - ) - self.prefill_cuda_graph_metadata[bs] = prefill_wrapper - self.forward_metadata = PrefillMetadata(prefill_wrapper, False) - else: - raise ValueError(f"Invalid mode: {forward_mode=}") - - self.init_forward_metadata_replay_cuda_graph( - bs=bs, - req_pool_indices=req_pool_indices, - seq_lens=seq_lens, - seq_lens_sum=seq_lens_sum, - encoder_lens=encoder_lens, - forward_mode=forward_mode, - spec_info=spec_info, - seq_lens_cpu=seq_lens_cpu, - ) - - def init_forward_metadata_replay_cuda_graph( + def _apply_cuda_graph_metadata( self, bs: int, req_pool_indices: torch.Tensor, seq_lens: torch.Tensor, seq_lens_sum: int, - encoder_lens: Optional[torch.Tensor], forward_mode: ForwardMode, spec_info: Optional[SpecInput], seq_lens_cpu: Optional[torch.Tensor], ): + """Shared capture+replay body for the cuda-graph init path. + + Public entry: :py:meth:`init_forward_metadata_out_graph`. + """ if forward_mode.is_decode_or_idle(): assert seq_lens_cpu is not None kv_len_arr_cpu = seq_lens_cpu[:bs] @@ -993,37 +1004,30 @@ class FlashInferMLAMultiStepDraftBackend: max_bs, max_num_tokens, kv_indices_buf=self.cuda_graph_kv_indices[i] ) - def init_forward_metadata_capture_cuda_graph(self, forward_batch: ForwardBatch): - def call_fn(i, forward_batch): - self.attn_backends[i].init_forward_metadata_capture_cuda_graph( - forward_batch.batch_size, - forward_batch.batch_size * self.topk, - forward_batch.req_pool_indices, - forward_batch.seq_lens, - encoder_lens=None, - forward_mode=ForwardMode.DECODE, - spec_info=forward_batch.spec_info, - ) - - self.common_template(forward_batch, self.cuda_graph_kv_indices, call_fn) - - def init_forward_metadata_replay_cuda_graph( - self, forward_batch: ForwardBatch, bs: int + def init_forward_metadata_out_graph( + self, + forward_batch: ForwardBatch, + in_capture: bool = False, ): - def call_fn(i, forward_batch): - self.attn_backends[i].init_forward_metadata_replay_cuda_graph( - bs, - forward_batch.req_pool_indices, - forward_batch.seq_lens, - seq_lens_sum=-1, - encoder_lens=None, - forward_mode=ForwardMode.DECODE, - spec_info=forward_batch.spec_info, - seq_lens_cpu=forward_batch.seq_lens_cpu, + from sglang.srt.model_executor.forward_batch_info import build_inner_fb_view + + inner_fb = build_inner_fb_view( + forward_batch, + bs=forward_batch.batch_size, + forward_mode=ForwardMode.DECODE, + ) + + def call_fn(i, _forward_batch): + self.attn_backends[i].init_forward_metadata_out_graph( + inner_fb, in_capture=in_capture ) self.common_template(forward_batch, self.cuda_graph_kv_indices, call_fn) + def init_forward_metadata_in_graph(self, forward_batch: ForwardBatch) -> None: + for attn_backend in self.attn_backends: + attn_backend.init_forward_metadata_in_graph(forward_batch) + def fast_mla_decode_plan( self, diff --git a/python/sglang/srt/layers/attention/flashmla_backend.py b/python/sglang/srt/layers/attention/flashmla_backend.py index ff77ed927..375fcebe3 100644 --- a/python/sglang/srt/layers/attention/flashmla_backend.py +++ b/python/sglang/srt/layers/attention/flashmla_backend.py @@ -21,7 +21,6 @@ from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMo if TYPE_CHECKING: from sglang.srt.layers.radix_attention import RadixAttention from sglang.srt.model_executor.model_runner import ModelRunner - from sglang.srt.speculative.spec_info import SpecInput logger = logging.getLogger(__name__) @@ -86,6 +85,25 @@ class FlashMLABackend(FlashInferMLAAttnBackend): self.cuda_graph_mla_metadata_view = None self.cuda_graph_num_splits_view = None + def init_forward_metadata_out_graph( + self, + forward_batch: ForwardBatch, + in_capture: bool = False, + ): + forward_mode = forward_batch.forward_mode + if forward_mode.is_decode_or_idle() or forward_mode.is_target_verify(): + self._apply_decode_target_verify_metadata( + bs=forward_batch.batch_size, + req_pool_indices=forward_batch.req_pool_indices, + seq_lens=forward_batch.seq_lens, + seq_lens_cpu=forward_batch.seq_lens_cpu, + forward_mode=forward_mode, + ) + else: + super().init_forward_metadata_out_graph( + forward_batch, in_capture=in_capture + ) + def init_forward_metadata(self, forward_batch: ForwardBatch): bs = forward_batch.batch_size if forward_batch.forward_mode.is_decode_or_idle(): @@ -185,50 +203,21 @@ class FlashMLABackend(FlashInferMLAAttnBackend): self.cuda_graph_mla_metadata_view = None self.cuda_graph_num_splits_view = None - def init_forward_metadata_capture_cuda_graph( - self, - bs: int, - num_tokens: int, - req_pool_indices: torch.Tensor, - seq_lens: torch.Tensor, - encoder_lens: Optional[torch.Tensor], - forward_mode: ForwardMode, - spec_info: Optional[SpecInput], - ): - if forward_mode.is_decode_or_idle() or forward_mode.is_target_verify(): - self.init_forward_metadata_replay_cuda_graph( - bs=bs, - req_pool_indices=req_pool_indices, - seq_lens=seq_lens, - seq_lens_sum=None, - encoder_lens=encoder_lens, - forward_mode=forward_mode, - spec_info=spec_info, - seq_lens_cpu=None, - ) - else: - super().init_forward_metadata_capture_cuda_graph( - bs, - num_tokens, - req_pool_indices, - seq_lens, - encoder_lens, - forward_mode, - spec_info, - ) - - def init_forward_metadata_replay_cuda_graph( + def _apply_decode_target_verify_metadata( self, bs: int, req_pool_indices: torch.Tensor, seq_lens: torch.Tensor, - seq_lens_sum: int, - encoder_lens: Optional[torch.Tensor], - forward_mode: ForwardMode, - spec_info: Optional[SpecInput], seq_lens_cpu: Optional[torch.Tensor], + forward_mode: ForwardMode, ): - if forward_mode.is_decode_or_idle() or forward_mode.is_target_verify(): + """Shared decode/target-verify capture+replay body. + + Public entry: :py:meth:`init_forward_metadata_out_graph` (which routes + to this helper for decode/target-verify and falls back to the + FlashInferMLA parent for prefill/draft-extend). + """ + if True: seq_lens = seq_lens[:bs] seq_lens_cpu = seq_lens_cpu[:bs] if seq_lens_cpu is not None else None @@ -295,17 +284,6 @@ class FlashMLABackend(FlashInferMLAAttnBackend): self.cuda_graph_num_splits_view, self.cuda_graph_kv_indices[:bs, :max_seqlen_pad], ) - else: - super().init_forward_metadata_replay_cuda_graph( - bs, - req_pool_indices, - seq_lens, - seq_lens_sum, - encoder_lens, - forward_mode, - spec_info, - seq_lens_cpu, - ) def get_cuda_graph_seq_len_fill_value(self): return 1 @@ -516,39 +494,29 @@ class FlashMLAMultiStepDraftBackend: max_bs, max_num_tokens, block_kv_indices=None ) - def init_forward_metadata_capture_cuda_graph(self, forward_batch: ForwardBatch): - def call_fn(i, forward_batch): - # EAGLE draft worker uses DECODE mode for draft steps - from sglang.srt.model_executor.forward_batch_info import ForwardMode - - # Create a dummy forward_mode for draft step - self.attn_backends[i].init_forward_metadata_capture_cuda_graph( - forward_batch.batch_size, - forward_batch.batch_size * self.topk, - forward_batch.req_pool_indices, - forward_batch.seq_lens, - encoder_lens=None, - forward_mode=ForwardMode.DECODE, - spec_info=forward_batch.spec_info, - ) - - self.common_template(forward_batch, call_fn) - - def init_forward_metadata_replay_cuda_graph( - self, forward_batch: ForwardBatch, bs: int + def init_forward_metadata_out_graph( + self, + forward_batch: ForwardBatch, + in_capture: bool = False, ): - def call_fn(i, forward_batch): - from sglang.srt.model_executor.forward_batch_info import ForwardMode + from sglang.srt.model_executor.forward_batch_info import ( + ForwardMode, + build_inner_fb_view, + ) - self.attn_backends[i].init_forward_metadata_replay_cuda_graph( - bs, - forward_batch.req_pool_indices, - forward_batch.seq_lens, - seq_lens_sum=-1, - encoder_lens=None, - forward_mode=ForwardMode.DECODE, - spec_info=forward_batch.spec_info, - seq_lens_cpu=forward_batch.seq_lens_cpu, + inner_fb = build_inner_fb_view( + forward_batch, + bs=forward_batch.batch_size, + forward_mode=ForwardMode.DECODE, + ) + + def call_fn(i, _forward_batch): + self.attn_backends[i].init_forward_metadata_out_graph( + inner_fb, in_capture=in_capture ) self.common_template(forward_batch, call_fn) + + def init_forward_metadata_in_graph(self, forward_batch: ForwardBatch) -> None: + for attn_backend in self.attn_backends: + attn_backend.init_forward_metadata_in_graph(forward_batch) diff --git a/python/sglang/srt/layers/attention/hybrid_attn_backend.py b/python/sglang/srt/layers/attention/hybrid_attn_backend.py index df0c70dc5..92e354dbe 100644 --- a/python/sglang/srt/layers/attention/hybrid_attn_backend.py +++ b/python/sglang/srt/layers/attention/hybrid_attn_backend.py @@ -7,7 +7,6 @@ from sglang.srt.layers.attention.dsa.dsa_indexer import BaseIndexerMetadata from sglang.srt.layers.radix_attention import RadixAttention from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.model_executor.model_runner import ModelRunner -from sglang.srt.speculative.spec_info import SpecInput class HybridAttnBackend(AttentionBackend): @@ -52,6 +51,14 @@ class HybridAttnBackend(AttentionBackend): else: return self.prefill_backend + def init_forward_metadata_out_graph( + self, + forward_batch: ForwardBatch, + in_capture: bool = False, + ): + backend = self._select_backend(forward_batch.forward_mode) + backend.init_forward_metadata_out_graph(forward_batch, in_capture=in_capture) + def init_forward_metadata(self, forward_batch: ForwardBatch): backend = self._select_backend(forward_batch.forward_mode) backend.init_forward_metadata(forward_batch) @@ -66,50 +73,6 @@ class HybridAttnBackend(AttentionBackend): # that will be used for target_verify. self.prefill_backend.init_cuda_graph_state(max_bs, max_num_tokens) - def init_forward_metadata_capture_cuda_graph( - self, - bs: int, - num_tokens: int, - req_pool_indices: torch.Tensor, - seq_lens: torch.Tensor, - encoder_lens: Optional[torch.Tensor], - forward_mode: ForwardMode, - spec_info: Optional[SpecInput], - ): - backend = self._select_backend(forward_mode) - backend.init_forward_metadata_capture_cuda_graph( - bs, - num_tokens, - req_pool_indices, - seq_lens, - encoder_lens, - forward_mode, - spec_info, - ) - - def init_forward_metadata_replay_cuda_graph( - self, - bs: int, - req_pool_indices: torch.Tensor, - seq_lens: torch.Tensor, - seq_lens_sum: int, - encoder_lens: Optional[torch.Tensor], - forward_mode: ForwardMode, - spec_info: Optional[SpecInput], - seq_lens_cpu: Optional[torch.Tensor], - ): - backend = self._select_backend(forward_mode) - backend.init_forward_metadata_replay_cuda_graph( - bs, - req_pool_indices, - seq_lens, - seq_lens_sum, - encoder_lens, - forward_mode, - spec_info, - seq_lens_cpu, - ) - def get_cuda_graph_seq_len_fill_value(self): return self.decode_backend.get_cuda_graph_seq_len_fill_value() diff --git a/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py b/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py index eaa78565d..71248dcf2 100644 --- a/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py +++ b/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py @@ -258,6 +258,21 @@ class MambaAttnBackendBase(AttentionBackend): has_mamba_track_mask=has_mamba_track_mask, ) + def init_forward_metadata_out_graph( + self, + forward_batch: ForwardBatch, + in_capture: bool = False, + ): + # seq_lens_cpu is unused by _replay_metadata for the non-target-verify + # case but kept in the contract for compatibility. + self.forward_metadata = self._replay_metadata( + forward_batch.batch_size, + forward_batch.req_pool_indices, + forward_batch.forward_mode, + forward_batch.spec_info, + forward_batch.seq_lens_cpu if not in_capture else None, + ) + def init_forward_metadata(self, forward_batch: ForwardBatch): self._execute_deferred_mamba_cow_and_clear(forward_batch) self.forward_metadata = self._forward_metadata(forward_batch) @@ -393,42 +408,6 @@ class MambaAttnBackendBase(AttentionBackend): track_ssm_final_dst.to(self.device, non_blocking=True), ) - def init_forward_metadata_capture_cuda_graph( - self, - bs: int, - num_tokens: int, - req_pool_indices: torch.Tensor, - seq_lens: torch.Tensor, - encoder_lens: Optional[torch.Tensor], - forward_mode: ForwardMode, - spec_info: Optional[Union[EagleDraftInput, EagleVerifyInput]], - ): - self.init_forward_metadata_replay_cuda_graph( - bs=bs, - req_pool_indices=req_pool_indices, - seq_lens=seq_lens, - seq_lens_sum=None, - encoder_lens=encoder_lens, - forward_mode=forward_mode, - spec_info=spec_info, - seq_lens_cpu=None, - ) - - def init_forward_metadata_replay_cuda_graph( - self, - bs: int, - req_pool_indices: torch.Tensor, - seq_lens: torch.Tensor, - seq_lens_sum: int, - encoder_lens: Optional[torch.Tensor], - forward_mode: ForwardMode, - spec_info: Optional[Union[EagleDraftInput, EagleVerifyInput]], - seq_lens_cpu: Optional[torch.Tensor], - ): - self.forward_metadata = self._replay_metadata( - bs, req_pool_indices, forward_mode, spec_info, seq_lens_cpu - ) - def init_forward_metadata_capture_cpu_graph( self, bs: int, @@ -698,6 +677,27 @@ class Mamba2AttnBackend(MambaAttnBackendBase): model_runner.server_args.mamba_track_interval >= self.mamba_chunk_size ), f"mamba_track_interval ({model_runner.server_args.mamba_track_interval}) must be >= mamba_chunk_size ({self.mamba_chunk_size})" + def init_forward_metadata_out_graph( + self, + forward_batch: ForwardBatch, + in_capture: bool = False, + ): + metadata = self._replay_metadata( + forward_batch.batch_size, + forward_batch.req_pool_indices, + forward_batch.forward_mode, + forward_batch.spec_info, + forward_batch.seq_lens_cpu if not in_capture else None, + ) + spec_info = forward_batch.spec_info + draft_token_num = spec_info.draft_token_num if spec_info is not None else 1 + self.forward_metadata = Mamba2Metadata.prepare_decode( + metadata, + forward_batch.seq_lens, + is_target_verify=forward_batch.forward_mode.is_target_verify(), + draft_token_num=draft_token_num, + ) + def init_forward_metadata(self, forward_batch: ForwardBatch): self._execute_deferred_mamba_cow_and_clear(forward_batch) metadata = self._forward_metadata(forward_batch) @@ -707,49 +707,6 @@ class Mamba2AttnBackend(MambaAttnBackendBase): forward_batch, ) - def init_forward_metadata_capture_cuda_graph( - self, - bs: int, - num_tokens: int, - req_pool_indices: torch.Tensor, - seq_lens: torch.Tensor, - encoder_lens: Optional[torch.Tensor], - forward_mode: ForwardMode, - spec_info: Optional[Union[EagleDraftInput, EagleVerifyInput]], - ): - self.init_forward_metadata_replay_cuda_graph( - bs=bs, - req_pool_indices=req_pool_indices, - seq_lens=seq_lens, - seq_lens_sum=None, - encoder_lens=encoder_lens, - forward_mode=forward_mode, - spec_info=spec_info, - seq_lens_cpu=None, - ) - - def init_forward_metadata_replay_cuda_graph( - self, - bs: int, - req_pool_indices: torch.Tensor, - seq_lens: torch.Tensor, - seq_lens_sum: int, - encoder_lens: Optional[torch.Tensor], - forward_mode: ForwardMode, - spec_info: Optional[Union[EagleDraftInput, EagleVerifyInput]], - seq_lens_cpu: Optional[torch.Tensor], - ): - metadata = self._replay_metadata( - bs, req_pool_indices, forward_mode, spec_info, seq_lens_cpu - ) - draft_token_num = spec_info.draft_token_num if spec_info is not None else 1 - self.forward_metadata = Mamba2Metadata.prepare_decode( - metadata, - seq_lens, - is_target_verify=forward_mode.is_target_verify(), - draft_token_num=draft_token_num, - ) - def forward( self, mixer: MambaMixer2, @@ -847,6 +804,16 @@ class HybridLinearAttnBackend(AttentionBackend): assert layer_id is not None, "either layer or layer_id must be provided" return layer_id in self.full_attn_layers + def init_forward_metadata_out_graph( + self, + forward_batch: ForwardBatch, + in_capture: bool = False, + ): + for attn_backend in self.attn_backend_list: + attn_backend.init_forward_metadata_out_graph( + forward_batch, in_capture=in_capture + ) + def init_forward_metadata(self, forward_batch: ForwardBatch): if forward_batch.forward_mode.is_draft_extend_v2(): # DRAFT_EXTEND_V2 only runs full-attn layers in the draft model, @@ -864,27 +831,6 @@ class HybridLinearAttnBackend(AttentionBackend): for attn_backend in self.attn_backend_list: attn_backend.init_cpu_graph_state(max_bs, max_num_tokens) - def init_forward_metadata_capture_cuda_graph( - self, - bs: int, - num_tokens: int, - req_pool_indices: torch.Tensor, - seq_lens: torch.Tensor, - encoder_lens: Optional[torch.Tensor], - forward_mode: ForwardMode, - spec_info: Optional[SpecInput], - ): - for attn_backend in self.attn_backend_list: - attn_backend.init_forward_metadata_capture_cuda_graph( - bs, - num_tokens, - req_pool_indices, - seq_lens, - encoder_lens, - forward_mode, - spec_info, - ) - def init_forward_metadata_capture_cpu_graph( self, bs: int, @@ -906,29 +852,6 @@ class HybridLinearAttnBackend(AttentionBackend): spec_info, ) - def init_forward_metadata_replay_cuda_graph( - self, - bs: int, - req_pool_indices: torch.Tensor, - seq_lens: torch.Tensor, - seq_lens_sum: int, - encoder_lens: Optional[torch.Tensor], - forward_mode: ForwardMode, - spec_info: Optional[SpecInput], - seq_lens_cpu: Optional[torch.Tensor], - ): - for attn_backend in self.attn_backend_list: - attn_backend.init_forward_metadata_replay_cuda_graph( - bs, - req_pool_indices, - seq_lens, - seq_lens_sum, - encoder_lens, - forward_mode, - spec_info, - seq_lens_cpu, - ) - def get_cuda_graph_seq_len_fill_value(self): return self.full_attn_backend.get_cuda_graph_seq_len_fill_value() diff --git a/python/sglang/srt/layers/attention/linear/lightning_backend.py b/python/sglang/srt/layers/attention/linear/lightning_backend.py index bc63db161..4a98de6cf 100644 --- a/python/sglang/srt/layers/attention/linear/lightning_backend.py +++ b/python/sglang/srt/layers/attention/linear/lightning_backend.py @@ -1,6 +1,5 @@ import logging import math -from typing import Optional, Union import torch @@ -9,12 +8,13 @@ from sglang.srt.layers.attention.linear.lightning_attn import ( BailingLinearKernel, linear_decode_forward_triton, ) -from sglang.srt.layers.attention.linear.linear_metadata import BailingLinearMetadata +from sglang.srt.layers.attention.linear.linear_metadata import ( + BailingLinearMetadata, +) from sglang.srt.layers.attention.linear.seg_la import SegLaMeta, seg_la_fwd from sglang.srt.layers.radix_attention import RadixAttention -from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode +from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.model_runner import ModelRunner -from sglang.srt.speculative.eagle_info import EagleDraftInput, EagleVerifyInput logger = logging.getLogger(__name__) @@ -71,6 +71,28 @@ class LightningAttentionBackend(MambaAttnBackendBase): f"linear_backend for linear attention in hybrid_linear_backend: {self.linear_backend}" ) + def init_forward_metadata_out_graph( + self, + forward_batch: ForwardBatch, + in_capture: bool = False, + ): + # seq_lens_cpu is unused by the underlying _replay_metadata for + # non-target-verify modes; pass it through for compatibility. + bs = forward_batch.batch_size + metadata = self._replay_metadata( + bs, + forward_batch.req_pool_indices, + forward_batch.forward_mode, + forward_batch.spec_info, + forward_batch.seq_lens_cpu if not in_capture else None, + ) + self.forward_metadata = BailingLinearMetadata.prepare_decode( + metadata.query_start_loc, + metadata.mamba_cache_indices, + bs, + forward_batch.seq_lens, + ) + def init_forward_metadata(self, forward_batch: ForwardBatch): metadata = self._forward_metadata(forward_batch) self.forward_metadata = BailingLinearMetadata.prepare_mixed( @@ -79,45 +101,6 @@ class LightningAttentionBackend(MambaAttnBackendBase): forward_batch, ) - def init_forward_metadata_capture_cuda_graph( - self, - bs: int, - num_tokens: int, - req_pool_indices: torch.Tensor, - seq_lens: torch.Tensor, - encoder_lens: Optional[torch.Tensor], - forward_mode: ForwardMode, - spec_info: Optional[Union[EagleDraftInput, EagleVerifyInput]], - ): - self.init_forward_metadata_replay_cuda_graph( - bs=bs, - req_pool_indices=req_pool_indices, - seq_lens=seq_lens, - seq_lens_sum=None, - encoder_lens=encoder_lens, - forward_mode=forward_mode, - spec_info=spec_info, - seq_lens_cpu=None, - ) - - def init_forward_metadata_replay_cuda_graph( - self, - bs: int, - req_pool_indices: torch.Tensor, - seq_lens: torch.Tensor, - seq_lens_sum: int, - encoder_lens: Optional[torch.Tensor], - forward_mode: ForwardMode, - spec_info: Optional[Union[EagleDraftInput, EagleVerifyInput]], - seq_lens_cpu: Optional[torch.Tensor], - ): - metadata = self._replay_metadata( - bs, req_pool_indices, forward_mode, spec_info, seq_lens_cpu - ) - self.forward_metadata = BailingLinearMetadata.prepare_decode( - metadata.query_start_loc, metadata.mamba_cache_indices, bs, seq_lens - ) - @staticmethod def _build_slope_tensor( n_attention_heads: int, num_hidden_layers: int, device="cuda" diff --git a/python/sglang/srt/layers/attention/tbo_backend.py b/python/sglang/srt/layers/attention/tbo_backend.py index 76d83b7b7..335d8cc4e 100644 --- a/python/sglang/srt/layers/attention/tbo_backend.py +++ b/python/sglang/srt/layers/attention/tbo_backend.py @@ -1,13 +1,11 @@ -from typing import TYPE_CHECKING, Callable, List, Optional - -import torch +from types import SimpleNamespace +from typing import TYPE_CHECKING, Callable, List from sglang.srt.batch_overlap import two_batch_overlap from sglang.srt.layers.attention.base_attn_backend import AttentionBackend -from sglang.srt.speculative.spec_info import SpecInput if TYPE_CHECKING: - from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode + from sglang.srt.model_executor.forward_batch_info import ForwardBatch class TboAttnBackend(AttentionBackend): @@ -27,6 +25,88 @@ class TboAttnBackend(AttentionBackend): children=[creator() for _ in range(2)], ) + def init_forward_metadata_out_graph( + self, + forward_batch: "ForwardBatch", + in_capture: bool = False, + ): + self.primary.init_forward_metadata_out_graph( + forward_batch=forward_batch, in_capture=in_capture + ) + tbo_children = getattr(forward_batch, "tbo_children", None) + if tbo_children is not None: + for child, forward_batch_child in zip( + self.children, tbo_children, strict=True + ): + if forward_batch_child.batch_size > 0: + child.init_forward_metadata_out_graph( + forward_batch=forward_batch_child, in_capture=in_capture + ) + return + if in_capture: + return + # Replay path: build_replay_fb_view returns a SimpleNamespace and + # tbo_plugin.replay_prepare does not call prepare_raw, so split the + # padded buffers here using the same indices the eager path would. + self._dispatch_children_from_replay_view(forward_batch) + + def _dispatch_children_from_replay_view(self, fb_view) -> None: + bs = fb_view.batch_size + forward_mode = fb_view.forward_mode + spec_info = fb_view.spec_info + token_num_per_seq = two_batch_overlap.get_token_num_per_seq( + forward_mode=forward_mode, spec_info=spec_info + ) + num_tokens = bs * token_num_per_seq + ( + tbo_split_seq_index, + tbo_split_token_index, + ) = two_batch_overlap.compute_split_indices_for_cuda_graph_replay( + forward_mode=forward_mode, + cuda_graph_num_tokens=num_tokens, + spec_info=spec_info, + ) + bs_left = tbo_split_seq_index + bs_right = bs - bs_left + for child, child_bs, seq_slice, tok_slice in ( + ( + self.children[0], + bs_left, + slice(None, tbo_split_seq_index), + slice(None, tbo_split_token_index), + ), + ( + self.children[1], + bs_right, + slice(tbo_split_seq_index, None), + slice(tbo_split_token_index, None), + ), + ): + if child_bs == 0: + continue + child_fb_view = _build_tbo_child_replay_fb_view( + fb_view, + child_bs=child_bs, + seq_slice=seq_slice, + tok_slice=tok_slice, + token_num_per_seq=token_num_per_seq, + ) + child.init_forward_metadata_out_graph( + forward_batch=child_fb_view, in_capture=False + ) + + def init_forward_metadata_in_graph(self, forward_batch: "ForwardBatch"): + self.primary.init_forward_metadata_in_graph(forward_batch=forward_batch) + tbo_children = getattr(forward_batch, "tbo_children", None) + if tbo_children is not None: + for child, forward_batch_child in zip( + self.children, tbo_children, strict=True + ): + if forward_batch_child.batch_size > 0: + child.init_forward_metadata_in_graph( + forward_batch=forward_batch_child + ) + def init_forward_metadata(self, forward_batch: "ForwardBatch"): self.primary.init_forward_metadata(forward_batch=forward_batch) if forward_batch.tbo_children is not None: @@ -42,140 +122,10 @@ class TboAttnBackend(AttentionBackend): # TODO for children, maybe can provide *smaller* max_bs to optimize item.init_cuda_graph_state(max_bs=max_bs, max_num_tokens=max_num_tokens) - def init_forward_metadata_capture_cuda_graph( - self, - bs: int, - num_tokens: int, - req_pool_indices: torch.Tensor, - seq_lens: torch.Tensor, - encoder_lens: Optional[torch.Tensor], - forward_mode: "ForwardMode", - spec_info: Optional[SpecInput], - ): - self.primary.init_forward_metadata_capture_cuda_graph( - bs=bs, - num_tokens=num_tokens, - req_pool_indices=req_pool_indices, - seq_lens=seq_lens, - encoder_lens=encoder_lens, - forward_mode=forward_mode, - spec_info=spec_info, - ) - - self._init_forward_metadata_cuda_graph_children( - fn_name="init_forward_metadata_capture_cuda_graph", - bs=bs, - req_pool_indices=req_pool_indices, - seq_lens=seq_lens, - encoder_lens=encoder_lens, - forward_mode=forward_mode, - spec_info=spec_info, - capture_num_tokens=num_tokens, - ) - - def init_forward_metadata_replay_cuda_graph( - self, - bs: int, - req_pool_indices: torch.Tensor, - seq_lens: torch.Tensor, - seq_lens_sum: int, - encoder_lens: Optional[torch.Tensor], - forward_mode: "ForwardMode", - spec_info: Optional[SpecInput], - seq_lens_cpu: Optional[torch.Tensor], - ): - self.primary.init_forward_metadata_replay_cuda_graph( - bs=bs, - req_pool_indices=req_pool_indices, - seq_lens=seq_lens, - seq_lens_sum=seq_lens_sum, - encoder_lens=encoder_lens, - forward_mode=forward_mode, - spec_info=spec_info, - seq_lens_cpu=seq_lens_cpu, - ) - - self._init_forward_metadata_cuda_graph_children( - fn_name="init_forward_metadata_replay_cuda_graph", - bs=bs, - req_pool_indices=req_pool_indices, - seq_lens=seq_lens, - encoder_lens=encoder_lens, - forward_mode=forward_mode, - spec_info=spec_info, - replay_seq_lens_sum=seq_lens_sum, - replay_seq_lens_cpu=seq_lens_cpu, - ) - - def _init_forward_metadata_cuda_graph_children( - self, - fn_name: str, - # common args - bs: int, - req_pool_indices: torch.Tensor, - seq_lens: torch.Tensor, - encoder_lens: Optional[torch.Tensor], - forward_mode: "ForwardMode", - spec_info: Optional[SpecInput], - # capture args - capture_num_tokens: int = None, - # replay args - replay_seq_lens_sum: int = None, - replay_seq_lens_cpu: Optional[torch.Tensor] = None, - ): - token_num_per_seq = two_batch_overlap.get_token_num_per_seq( - forward_mode=forward_mode, spec_info=spec_info - ) - if fn_name == "init_forward_metadata_capture_cuda_graph": - assert ( - capture_num_tokens == bs * token_num_per_seq - ), "For target-verify or decode mode, num_tokens should be equal to token_num_per_seq * bs" - num_tokens = bs * token_num_per_seq - - tbo_split_seq_index, tbo_split_token_index = ( - two_batch_overlap.compute_split_indices_for_cuda_graph_replay( - forward_mode=forward_mode, - cuda_graph_num_tokens=num_tokens, - spec_info=spec_info, - ) - ) - - num_tokens_child_left = tbo_split_token_index - num_tokens_child_right = num_tokens - tbo_split_token_index - bs_child_left = tbo_split_seq_index - bs_child_right = bs - bs_child_left - - assert ( - num_tokens_child_left > 0 and num_tokens_child_right > 0 - ), f"{num_tokens_child_left=} {num_tokens_child_right=} {forward_mode=} {num_tokens=}" - - common_pre_split_args = dict( - fn_name=fn_name, - bs=bs, - req_pool_indices=req_pool_indices, - seq_lens=seq_lens, - encoder_lens=encoder_lens, - forward_mode=forward_mode, - spec_info=spec_info, - capture_num_tokens=capture_num_tokens, - replay_seq_lens_sum=replay_seq_lens_sum, - replay_seq_lens_cpu=replay_seq_lens_cpu, - ) - - args_left = _init_forward_metadata_cuda_graph_split( - output_bs=bs_child_left, - seq_slice=slice(None, tbo_split_seq_index), - **common_pre_split_args, - ) - args_right = _init_forward_metadata_cuda_graph_split( - output_bs=bs_child_right, - seq_slice=slice(tbo_split_seq_index, None), - **common_pre_split_args, - ) - - child_left, child_right = self.children - getattr(child_left, fn_name)(**args_left) - getattr(child_right, fn_name)(**args_right) + def on_after_cuda_graph_warmup(self): + self.primary.on_after_cuda_graph_warmup() + for child in self.children: + child.on_after_cuda_graph_warmup() def get_cuda_graph_seq_len_fill_value(self): ans = self.primary.get_cuda_graph_seq_len_fill_value() @@ -196,75 +146,59 @@ class TboAttnBackend(AttentionBackend): return self.primary.get_indexer_metadata(layer_id, forward_batch) -def _init_forward_metadata_cuda_graph_split( - fn_name: str, +def _build_tbo_child_replay_fb_view( + fb_view, + *, + child_bs: int, seq_slice: slice, - output_bs: int, - # common args - bs: int, - req_pool_indices: torch.Tensor, - seq_lens: torch.Tensor, - encoder_lens: Optional[torch.Tensor], - forward_mode: "ForwardMode", - spec_info: Optional[SpecInput], - # capture args - capture_num_tokens: int = None, - # replay args - replay_seq_lens_sum: int = None, - replay_seq_lens_cpu: Optional[torch.Tensor] = None, -): - token_num_per_seq = two_batch_overlap.get_token_num_per_seq( - forward_mode=forward_mode, spec_info=spec_info - ) - assert encoder_lens is None, "encoder_lens is not supported yet" + tok_slice: slice, + token_num_per_seq: int, +) -> SimpleNamespace: + """Slice a parent replay fb_view into a per-child view. + + Mirrors the legacy ``_init_forward_metadata_cuda_graph_split`` (deleted + along with the cuda_graph variants) for the new + ``init_forward_metadata_out_graph(fb_view)`` contract: padded + capture-time buffers are sliced per child, spec_info is split, and + seq_lens_sum is recomputed from the sliced ``seq_lens_cpu``. + """ + assert ( + getattr(fb_view, "encoder_lens", None) is None + ), "TBO replay split does not support encoder_lens yet" + spec_info = getattr(fb_view, "spec_info", None) if spec_info is not None: - output_spec_info = two_batch_overlap.split_spec_info( + start_seq = seq_slice.start or 0 + end_seq = seq_slice.stop if seq_slice.stop is not None else start_seq + child_bs + child_spec_info = two_batch_overlap.split_spec_info( spec_info=spec_info, - start_seq_index=seq_slice.start if seq_slice.start is not None else 0, - end_seq_index=seq_slice.stop if seq_slice.stop is not None else bs, - start_token_index=( - seq_slice.start * token_num_per_seq - if seq_slice.start is not None - else 0 - ), - end_token_index=( - seq_slice.stop * token_num_per_seq - if seq_slice.stop is not None - else bs * token_num_per_seq - ), + start_seq_index=start_seq, + end_seq_index=end_seq, + start_token_index=start_seq * token_num_per_seq, + end_token_index=end_seq * token_num_per_seq, ) - else: - output_spec_info = None - ans = dict( - bs=output_bs, - req_pool_indices=req_pool_indices[seq_slice], - seq_lens=seq_lens[seq_slice], - # directly forward - forward_mode=forward_mode, - # ignore + child_spec_info = None + child_seq_lens_cpu = fb_view.seq_lens_cpu[seq_slice] + parent_input_ids = getattr(fb_view, "input_ids", None) + parent_out_cache_loc = getattr(fb_view, "out_cache_loc", None) + return SimpleNamespace( + batch_size=child_bs, + forward_mode=fb_view.forward_mode, + actual_forward_mode=getattr( + fb_view, "actual_forward_mode", fb_view.forward_mode + ), + input_ids=( + parent_input_ids[tok_slice] if parent_input_ids is not None else None + ), + req_pool_indices=fb_view.req_pool_indices[seq_slice], + seq_lens=fb_view.seq_lens[seq_slice], + seq_lens_sum=int(child_seq_lens_cpu.sum()), + seq_lens_cpu=child_seq_lens_cpu, encoder_lens=None, - spec_info=output_spec_info, + out_cache_loc=( + parent_out_cache_loc[tok_slice] + if parent_out_cache_loc is not None + else None + ), + spec_info=child_spec_info, ) - - if fn_name == "init_forward_metadata_capture_cuda_graph": - assert ( - capture_num_tokens == bs * token_num_per_seq - ), "Only support num_tokens==bs * token_num_per_seq for target-verify or decode mode" - ans.update( - dict( - num_tokens=output_bs * token_num_per_seq, - ) - ) - elif fn_name == "init_forward_metadata_replay_cuda_graph": - output_seq_lens_cpu = replay_seq_lens_cpu[seq_slice] - ans.update( - dict( - seq_lens_sum=output_seq_lens_cpu.sum().item(), - seq_lens_cpu=output_seq_lens_cpu, - ) - ) - else: - raise NotImplementedError - - return ans diff --git a/python/sglang/srt/layers/attention/triton_backend.py b/python/sglang/srt/layers/attention/triton_backend.py index 1a0e7afc0..01c27f3e8 100644 --- a/python/sglang/srt/layers/attention/triton_backend.py +++ b/python/sglang/srt/layers/attention/triton_backend.py @@ -458,6 +458,59 @@ class TritonAttnBackend(AttentionBackend): ) return qo_indptr, kv_indptr, num_tokens_per_bs + def init_forward_metadata_out_graph( + self, + forward_batch: ForwardBatch, + in_capture: bool = False, + ): + bs = forward_batch.batch_size + req_pool_indices = forward_batch.req_pool_indices + seq_lens = forward_batch.seq_lens + forward_mode = forward_batch.forward_mode + spec_info = forward_batch.spec_info + + if in_capture: + assert forward_batch.encoder_lens is None, "Not supported" + # Multi-step speculative decode: kv buffers come from spec_info + # rather than the cuda-graph pool, so replay is not involved. + if forward_mode.is_decode_or_idle() and spec_info is not None: + self.forward_metadata = ForwardMetadata( + attn_logits=self.cuda_graph_attn_logits, + attn_lse=self.cuda_graph_attn_lse, + max_extend_len=None, + num_kv_splits=self.cuda_graph_num_kv_splits, + kv_indptr=spec_info.kv_indptr, + kv_indices=spec_info.kv_indices, + qo_indptr=None, + custom_mask=None, + mask_indptr=None, + window_kv_indptr=self.window_kv_indptr, + window_kv_indices=None, + window_num_kv_splits=None, + window_kv_offsets=None, + swa_attn_logits=self.cuda_graph_swa_attn_logits, + ) + return + + self._apply_cuda_graph_metadata( + bs=bs, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + forward_mode=forward_mode, + spec_info=spec_info, + ) + self.forward_metadata = self._build_cuda_graph_forward_metadata( + bs, forward_mode, spec_info + ) + else: + self._apply_cuda_graph_metadata( + bs=bs, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + forward_mode=forward_mode, + spec_info=spec_info, + ) + def init_forward_metadata(self, forward_batch: ForwardBatch): """Init auxiliary variables for triton attention backend.""" @@ -835,66 +888,18 @@ class TritonAttnBackend(AttentionBackend): else: raise ValueError(f"Invalid forward mode: {forward_mode=} for CUDA Graph.") - def init_forward_metadata_capture_cuda_graph( - self, - bs: int, - num_tokens: int, - req_pool_indices: torch.Tensor, - seq_lens: torch.Tensor, - encoder_lens: Optional[torch.Tensor], - forward_mode: ForwardMode, - spec_info: Optional[SpecInput], - ): - assert encoder_lens is None, "Not supported" - - # Multi-step speculative decode: kv buffers come from spec_info rather - # than the cuda-graph pool, so replay is not involved for this path. - if forward_mode.is_decode_or_idle() and spec_info is not None: - self.forward_metadata = ForwardMetadata( - attn_logits=self.cuda_graph_attn_logits, - attn_lse=self.cuda_graph_attn_lse, - max_extend_len=None, - num_kv_splits=self.cuda_graph_num_kv_splits, - kv_indptr=spec_info.kv_indptr, - kv_indices=spec_info.kv_indices, - qo_indptr=None, - custom_mask=None, - mask_indptr=None, - window_kv_indptr=self.window_kv_indptr, - window_kv_indices=None, - window_num_kv_splits=None, - window_kv_offsets=None, - swa_attn_logits=self.cuda_graph_swa_attn_logits, - ) - return - - # Run the same buffer update as replay, then freeze the result into - # a ForwardMetadata whose tensor fields point into the cuda-graph buffers. - self.init_forward_metadata_replay_cuda_graph( - bs=bs, - req_pool_indices=req_pool_indices, - seq_lens=seq_lens, - seq_lens_sum=None, - encoder_lens=encoder_lens, - forward_mode=forward_mode, - spec_info=spec_info, - seq_lens_cpu=None, - ) - self.forward_metadata = self._build_cuda_graph_forward_metadata( - bs, forward_mode, spec_info - ) - - def init_forward_metadata_replay_cuda_graph( + def _apply_cuda_graph_metadata( self, bs: int, req_pool_indices: torch.Tensor, seq_lens: torch.Tensor, - seq_lens_sum: int, - encoder_lens: Optional[torch.Tensor], forward_mode: ForwardMode, spec_info: Optional[SpecInput], - seq_lens_cpu: Optional[torch.Tensor], ): + """Shared capture+replay body for the cuda-graph init path. + + Public entry: :py:meth:`init_forward_metadata_out_graph`. + """ # NOTE: encoder_lens expected to be zeros or None if forward_mode.is_decode_or_idle(): assert spec_info is None, "Multi-step cuda graph init is not done here." @@ -1417,36 +1422,45 @@ class TritonMultiStepDraftBackend: cuda_graph_num_kv_splits_buf=self.cuda_graph_num_kv_splits, ) - def init_forward_metadata_capture_cuda_graph(self, forward_batch: ForwardBatch): - def call_fn(i, forward_batch): - self.attn_backends[i].init_forward_metadata_capture_cuda_graph( - forward_batch.batch_size, - forward_batch.batch_size * self.topk, - forward_batch.req_pool_indices, - forward_batch.seq_lens, - encoder_lens=None, + def init_forward_metadata_out_graph( + self, + forward_batch: ForwardBatch, + in_capture: bool = False, + ): + from sglang.srt.model_executor.forward_batch_info import build_inner_fb_view + + if in_capture: + inner_fb = build_inner_fb_view( + forward_batch, + bs=forward_batch.batch_size, forward_mode=ForwardMode.DECODE, - spec_info=forward_batch.spec_info, ) - self.common_template(forward_batch, None, call_fn) + def call_fn(i, _forward_batch): + self.attn_backends[i].init_forward_metadata_out_graph( + inner_fb, in_capture=True + ) - def init_forward_metadata_replay_cuda_graph( - self, forward_batch: ForwardBatch, bs: int - ): - self.common_template(forward_batch, None, None) + self.common_template(forward_batch, None, call_fn) + else: + bs = forward_batch.batch_size + self.common_template(forward_batch, None, None) - # NOTE: Multi-step's attention backends use the slice of - # - kv_indptr buffer (cuda graph and non-cuda graph) - # - kv_indices buffer (cuda graph only) - # So we don't need to assign the KV indices inside the attention backend. + # NOTE: Multi-step's attention backends use the slice of + # - kv_indptr buffer (cuda graph and non-cuda graph) + # - kv_indices buffer (cuda graph only) + # So we don't need to assign the KV indices inside the attention backend. - # Compute num_kv_splits only once - num_token = forward_batch.batch_size * self.topk - self.attn_backends[-1].get_num_kv_splits( - self.attn_backends[-1].cuda_graph_num_kv_splits[:num_token], - forward_batch.seq_lens[:bs], - ) + # Compute num_kv_splits only once + num_token = bs * self.topk + self.attn_backends[-1].get_num_kv_splits( + self.attn_backends[-1].cuda_graph_num_kv_splits[:num_token], + forward_batch.seq_lens[:bs], + ) + + def init_forward_metadata_in_graph(self, forward_batch: ForwardBatch) -> None: + for attn_backend in self.attn_backends: + attn_backend.init_forward_metadata_in_graph(forward_batch) def update_sliding_window_buffer( diff --git a/python/sglang/srt/layers/attention/trtllm_mha_backend.py b/python/sglang/srt/layers/attention/trtllm_mha_backend.py index 69c1409d2..ec3f64c36 100644 --- a/python/sglang/srt/layers/attention/trtllm_mha_backend.py +++ b/python/sglang/srt/layers/attention/trtllm_mha_backend.py @@ -397,50 +397,19 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend): return metadata - def init_forward_metadata_capture_cuda_graph( - self, - bs: int, - num_tokens: int, - req_pool_indices: torch.Tensor, - seq_lens: torch.Tensor, - encoder_lens: Optional[torch.Tensor], - forward_mode: ForwardMode, - spec_info: Optional[SpecInput], - ): - """Initialize metadata for CUDA graph capture.""" - seq_lens_cpu = seq_lens.cpu() - self._build_cuda_graph_metadata( - bs, num_tokens, forward_mode, spec_info, seq_lens.device - ) - self.init_forward_metadata_replay_cuda_graph( - bs=bs, - req_pool_indices=req_pool_indices, - seq_lens=seq_lens, - seq_lens_sum=None, - encoder_lens=encoder_lens, - forward_mode=forward_mode, - spec_info=spec_info, - seq_lens_cpu=seq_lens_cpu, - ) - if forward_mode.is_draft_extend(): - # CUDA graph bakes max_seq_len_q as a constant. replay() sets it to - # max(num_accept_tokens_cpu) which is None/empty at capture time, - # falling back to 1. Restore the correct upper bound so the kernel - # sees num_tokens_per_bs (not 1) for all replays of this graph. - self.forward_metadata.max_seq_len_q = num_tokens // bs - - def init_forward_metadata_replay_cuda_graph( + def _apply_cuda_graph_metadata( self, bs: int, req_pool_indices: torch.Tensor, seq_lens: torch.Tensor, - seq_lens_sum: int, - encoder_lens: Optional[torch.Tensor], forward_mode: ForwardMode, spec_info: Optional[SpecInput], seq_lens_cpu: Optional[torch.Tensor], ): - """Replay CUDA graph with new inputs.""" + """Shared capture+replay body for the cuda-graph init path. + + Public entry: :py:meth:`init_forward_metadata_out_graph`. + """ seq_lens = seq_lens[:bs] seq_lens_cpu = seq_lens_cpu[:bs] req_pool_indices = req_pool_indices[:bs] @@ -571,6 +540,48 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend): page_size=self.page_size, ) + def init_forward_metadata_out_graph( + self, + forward_batch: ForwardBatch, + in_capture: bool = False, + ): + bs = forward_batch.batch_size + req_pool_indices = forward_batch.req_pool_indices + seq_lens = forward_batch.seq_lens + encoder_lens = forward_batch.encoder_lens + forward_mode = forward_batch.forward_mode + spec_info = forward_batch.spec_info + + if in_capture: + num_tokens = forward_batch.positions.numel() + seq_lens_cpu = seq_lens.cpu() + self._build_cuda_graph_metadata( + bs, num_tokens, forward_mode, spec_info, seq_lens.device + ) + self._apply_cuda_graph_metadata( + bs=bs, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + forward_mode=forward_mode, + spec_info=spec_info, + seq_lens_cpu=seq_lens_cpu, + ) + if forward_mode.is_draft_extend(): + # CUDA graph bakes max_seq_len_q as a constant. replay() sets it + # to max(num_accept_tokens_cpu) which is None/empty at capture + # time, falling back to 1. Restore the correct upper bound so + # the kernel sees num_tokens_per_bs (not 1) for all replays. + self.forward_metadata.max_seq_len_q = num_tokens // bs + else: + self._apply_cuda_graph_metadata( + bs=bs, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + forward_mode=forward_mode, + spec_info=spec_info, + seq_lens_cpu=forward_batch.seq_lens_cpu, + ) + def init_forward_metadata(self, forward_batch: ForwardBatch): """Initialize the metadata for a forward pass.""" @@ -899,39 +910,25 @@ class TRTLLMHAAttnMultiStepDraftBackend(FlashInferMultiStepDraftBackend): for i in range(self.speculative_num_steps - 1): self.attn_backends[i].init_cuda_graph_state(max_bs, max_num_tokens) - def init_forward_metadata_capture_cuda_graph( + def init_forward_metadata_out_graph( self, forward_batch: ForwardBatch, + in_capture: bool = False, ): + from sglang.srt.model_executor.forward_batch_info import build_inner_fb_view + assert forward_batch.spec_info is not None assert forward_batch.spec_info.is_draft_input() + # TRTLLM-MHA uses encoder_lens from the original fb for inner dispatch + # (FlashInfer parent forces encoder_lens=None instead). + inner_fb = build_inner_fb_view( + forward_batch, + bs=forward_batch.batch_size, + forward_mode=ForwardMode.DECODE, + encoder_lens=forward_batch.encoder_lens, + ) for i in range(self.speculative_num_steps - 1): - self.attn_backends[i].init_forward_metadata_capture_cuda_graph( - forward_batch.batch_size, - forward_batch.batch_size * self.topk, - forward_batch.req_pool_indices, - forward_batch.seq_lens, - encoder_lens=forward_batch.encoder_lens, - forward_mode=ForwardMode.DECODE, - spec_info=forward_batch.spec_info, - ) - - def init_forward_metadata_replay_cuda_graph( - self, forward_batch: ForwardBatch, bs: int - ): - assert forward_batch.spec_info is not None - assert forward_batch.spec_info.is_draft_input() - - for i in range(self.speculative_num_steps - 1): - - self.attn_backends[i].init_forward_metadata_replay_cuda_graph( - bs, - forward_batch.req_pool_indices, - forward_batch.seq_lens, - forward_batch.seq_lens_sum, - encoder_lens=forward_batch.encoder_lens, - forward_mode=ForwardMode.DECODE, - spec_info=forward_batch.spec_info, - seq_lens_cpu=forward_batch.seq_lens_cpu, + self.attn_backends[i].init_forward_metadata_out_graph( + inner_fb, in_capture=in_capture ) diff --git a/python/sglang/srt/layers/attention/trtllm_mla_backend.py b/python/sglang/srt/layers/attention/trtllm_mla_backend.py index b9d55bfe9..bef1f6791 100755 --- a/python/sglang/srt/layers/attention/trtllm_mla_backend.py +++ b/python/sglang/srt/layers/attention/trtllm_mla_backend.py @@ -38,7 +38,6 @@ if is_flashinfer_available(): if TYPE_CHECKING: from sglang.srt.layers.radix_attention import RadixAttention from sglang.srt.model_executor.model_runner import ModelRunner - from sglang.srt.speculative.spec_info import SpecInput logger = logging.getLogger(__name__) @@ -499,83 +498,23 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend): self.decode_cuda_graph_metadata[bs] = metadata self.forward_decode_metadata = metadata - def init_forward_metadata_capture_cuda_graph( - self, - bs: int, - num_tokens: int, - req_pool_indices: torch.Tensor, - seq_lens: torch.Tensor, - encoder_lens: Optional[torch.Tensor], - forward_mode: ForwardMode, - spec_info: Optional[SpecInput], - ): - """Initialize metadata for CUDA graph capture.""" - - # Delegate to parent for non-decode modes. - if ( - not forward_mode.is_decode_or_idle() - and not forward_mode.is_target_verify() - and not forward_mode.is_draft_extend(include_v2=True) - ): - return super().init_forward_metadata_capture_cuda_graph( - bs, - num_tokens, - req_pool_indices, - seq_lens, - encoder_lens, - forward_mode, - spec_info, - ) - - self._init_cuda_graph_metadata( - bs, num_tokens, forward_mode, seq_lens, seq_lens.device - ) - self.init_forward_metadata_replay_cuda_graph( - bs=bs, - req_pool_indices=req_pool_indices, - seq_lens=seq_lens, - seq_lens_sum=None, - encoder_lens=encoder_lens, - forward_mode=forward_mode, - spec_info=spec_info, - seq_lens_cpu=seq_lens.cpu(), - ) - - def init_forward_metadata_replay_cuda_graph( + def _apply_cuda_graph_metadata( self, bs: int, req_pool_indices: torch.Tensor, seq_lens: torch.Tensor, - seq_lens_sum: int, - encoder_lens: Optional[torch.Tensor], forward_mode: ForwardMode, - spec_info: Optional[SpecInput], - seq_lens_cpu: Optional[torch.Tensor], ): - """Replay CUDA graph with new inputs.""" - # Delegate to parent for non-decode modes. - if ( - not forward_mode.is_decode_or_idle() - and not forward_mode.is_target_verify() - and not forward_mode.is_draft_extend(include_v2=True) - ): - return super().init_forward_metadata_replay_cuda_graph( - bs, - req_pool_indices, - seq_lens, - seq_lens_sum, - encoder_lens, - forward_mode, - spec_info, - seq_lens_cpu, - ) + """Shared decode / target-verify / draft-extend capture+replay body. + Public entry: :py:meth:`init_forward_metadata_out_graph` (which routes + the non-decode-family modes to the FlashInferMLA parent). + """ metadata = self.decode_cuda_graph_metadata[bs] if forward_mode.is_target_verify(): seq_lens = seq_lens[:bs] + self.num_draft_tokens metadata.seq_lens_k.copy_(seq_lens.to(dtype=torch.int32)) - del seq_lens_sum # not handle "num_draft_tokens" but we do not need it elif forward_mode.is_draft_extend(include_v2=True): num_tokens_per_bs = self.num_draft_tokens metadata.max_seq_len_q = num_tokens_per_bs @@ -618,6 +557,47 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend): if fallback_to_flashinfer_impl: super().init_mha_chunk_metadata(forward_batch) + def init_forward_metadata_out_graph( + self, + forward_batch: ForwardBatch, + in_capture: bool = False, + ): + forward_mode = forward_batch.forward_mode + + if ( + not forward_mode.is_decode_or_idle() + and not forward_mode.is_target_verify() + and not forward_mode.is_draft_extend(include_v2=True) + ): + return super().init_forward_metadata_out_graph( + forward_batch, in_capture=in_capture + ) + + bs = forward_batch.batch_size + if in_capture: + num_tokens = forward_batch.positions.numel() + seq_lens_cpu = forward_batch.seq_lens.cpu() + self._init_cuda_graph_metadata( + bs, + num_tokens, + forward_mode, + forward_batch.seq_lens, + forward_batch.seq_lens.device, + ) + self._apply_cuda_graph_metadata( + bs=bs, + req_pool_indices=forward_batch.req_pool_indices, + seq_lens=forward_batch.seq_lens, + forward_mode=forward_mode, + ) + else: + self._apply_cuda_graph_metadata( + bs=bs, + req_pool_indices=forward_batch.req_pool_indices, + seq_lens=forward_batch.seq_lens, + forward_mode=forward_mode, + ) + def init_forward_metadata(self, forward_batch: ForwardBatch): """Initialize the metadata for a forward pass.""" # Delegate to parent for non-decode modes. @@ -1284,17 +1264,21 @@ class TRTLLMMLAMultiStepDraftBackend(FlashInferMLAMultiStepDraftBackend): for i in range(self.speculative_num_steps - 1): self.attn_backends[i].init_forward_metadata(forward_batch) - def init_forward_metadata_replay_cuda_graph( - self, forward_batch: ForwardBatch, bs: int + def init_forward_metadata_out_graph( + self, + forward_batch: ForwardBatch, + in_capture: bool = False, ): - for i in range(self.speculative_num_steps - 1): - self.attn_backends[i].init_forward_metadata_replay_cuda_graph( - bs, - forward_batch.req_pool_indices, - forward_batch.seq_lens, - seq_lens_sum=None, - encoder_lens=None, - forward_mode=ForwardMode.DECODE, - spec_info=forward_batch.spec_info, - seq_lens_cpu=forward_batch.seq_lens_cpu, + from sglang.srt.model_executor.forward_batch_info import build_inner_fb_view + + if in_capture: + return super().init_forward_metadata_out_graph( + forward_batch, in_capture=in_capture ) + inner_fb = build_inner_fb_view( + forward_batch, + bs=forward_batch.batch_size, + forward_mode=ForwardMode.DECODE, + ) + for i in range(self.speculative_num_steps - 1): + self.attn_backends[i].init_forward_metadata_out_graph(inner_fb) diff --git a/python/sglang/srt/layers/attention/wave_backend.py b/python/sglang/srt/layers/attention/wave_backend.py index c7975c96c..0922eca37 100644 --- a/python/sglang/srt/layers/attention/wave_backend.py +++ b/python/sglang/srt/layers/attention/wave_backend.py @@ -147,6 +147,53 @@ class WaveAttnBackend(AttentionBackend): MAX_NUM_SEQ=SCHEDULE_SEQ, ) + def init_forward_metadata_out_graph( + self, + forward_batch: ForwardBatch, + in_capture: bool = False, + ): + bs = forward_batch.batch_size + req_pool_indices = forward_batch.req_pool_indices + seq_lens = forward_batch.seq_lens + forward_mode = forward_batch.forward_mode + spec_info = forward_batch.spec_info + + if in_capture: + assert forward_batch.encoder_lens is None, "Not supported" + # kv buffers come from spec_info rather than the cuda-graph pool. + if forward_mode.is_decode_or_idle() and spec_info is not None: + self.forward_metadata = ForwardMetadata( + attn_logits=self.cuda_graph_attn_logits, + attn_lse=self.cuda_graph_attn_lse, + max_extend_len=None, + num_kv_splits=self.cuda_graph_num_kv_splits, + kv_indptr=spec_info.kv_indptr, + kv_indices=spec_info.kv_indices, + qo_indptr=None, + custom_mask=None, + mask_indptr=None, + ) + return + + self._apply_cuda_graph_metadata( + bs=bs, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + forward_mode=forward_mode, + spec_info=spec_info, + ) + self.forward_metadata = self._build_cuda_graph_forward_metadata( + bs, forward_mode, spec_info + ) + else: + self._apply_cuda_graph_metadata( + bs=bs, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + forward_mode=forward_mode, + spec_info=spec_info, + ) + def init_forward_metadata(self, forward_batch: ForwardBatch): """Init auxiliary variables for wave attention backend.""" @@ -373,59 +420,18 @@ class WaveAttnBackend(AttentionBackend): else: raise ValueError(f"Invalid forward mode: {forward_mode=} for CUDA Graph.") - def init_forward_metadata_capture_cuda_graph( - self, - bs: int, - num_tokens: int, - req_pool_indices: torch.Tensor, - seq_lens: torch.Tensor, - encoder_lens: Optional[torch.Tensor], - forward_mode: ForwardMode, - spec_info: Optional[SpecInput], - ): - assert encoder_lens is None, "Not supported" - - # Multi-step speculative decode: kv buffers come from spec_info rather than - # the cuda-graph pool, so replay is not involved for this path. - if forward_mode.is_decode_or_idle() and spec_info is not None: - self.forward_metadata = ForwardMetadata( - attn_logits=self.cuda_graph_attn_logits, - attn_lse=self.cuda_graph_attn_lse, - max_extend_len=None, - num_kv_splits=self.cuda_graph_num_kv_splits, - kv_indptr=spec_info.kv_indptr, - kv_indices=spec_info.kv_indices, - qo_indptr=None, - custom_mask=None, - mask_indptr=None, - ) - return - - self.init_forward_metadata_replay_cuda_graph( - bs=bs, - req_pool_indices=req_pool_indices, - seq_lens=seq_lens, - seq_lens_sum=None, - encoder_lens=encoder_lens, - forward_mode=forward_mode, - spec_info=spec_info, - seq_lens_cpu=None, - ) - self.forward_metadata = self._build_cuda_graph_forward_metadata( - bs, forward_mode, spec_info - ) - - def init_forward_metadata_replay_cuda_graph( + def _apply_cuda_graph_metadata( self, bs: int, req_pool_indices: torch.Tensor, seq_lens: torch.Tensor, - seq_lens_sum: int, - encoder_lens: Optional[torch.Tensor], forward_mode: ForwardMode, spec_info: Optional[SpecInput], - seq_lens_cpu: Optional[torch.Tensor], ): + """Shared capture+replay body for the cuda-graph init path. + + Public entry: :py:meth:`init_forward_metadata_out_graph`. + """ if forward_mode.is_decode_or_idle(): kv_indptr = self.kv_indptr kv_indices = self.cuda_graph_kv_indices diff --git a/python/sglang/srt/layers/attention/xpu_backend.py b/python/sglang/srt/layers/attention/xpu_backend.py index 3b5743799..41293a20b 100644 --- a/python/sglang/srt/layers/attention/xpu_backend.py +++ b/python/sglang/srt/layers/attention/xpu_backend.py @@ -906,7 +906,10 @@ class XPUAttentionBackend(AttentionBackend): return 1 def _init_local_attn_metadata( - self, forwardbatch: ForwardBatch, metadata: FlashAttentionMetadata, device + self, + forwardbatch: ForwardBatch, + metadata: FlashAttentionMetadata, + device, ): """Centralized utility to initialize local_attn_metadata if chunked attention is enabled.""" if self.attention_chunk_size is None: diff --git a/python/sglang/srt/model_executor/cuda_graph_runner.py b/python/sglang/srt/model_executor/cuda_graph_runner.py index e610da7fa..e58d43b7e 100644 --- a/python/sglang/srt/model_executor/cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/cuda_graph_runner.py @@ -24,6 +24,7 @@ import os from contextlib import contextmanager from dataclasses import dataclass from functools import partial +from types import SimpleNamespace from typing import TYPE_CHECKING, Callable, Dict, List, Optional, Tuple, Union import torch @@ -107,6 +108,57 @@ if TYPE_CHECKING: _has_foreach_copy = hasattr(torch, "_foreach_copy_") +def build_replay_fb_view( + forward_batch: "ForwardBatch", + buffers: "DecodeInputBuffers", + bs: int, + raw_bs: int, + num_tokens: int, + seq_len_fill_value: int, + capture_forward_mode: "ForwardMode", + is_encoder_decoder: bool, +) -> SimpleNamespace: + """Construct a ForwardBatch-like view for backend replay-side init. + + Combines the original ``forward_batch`` (for unpadded / per-iter + fields like ``spec_info``, ``out_cache_loc``, and the runtime + ``actual_forward_mode``) with the padded capture-time buffers from + ``buffers`` (for ``req_pool_indices``, ``seq_lens``, ``seq_lens_cpu``, + ``encoder_lens``). + + Field semantics: + + - ``forward_mode``: the capture-time mode (``capture_forward_mode``), + used by backends for bucket / dispatch decisions (e.g. choosing + between decode / target-verify / draft-extend code paths). + - ``actual_forward_mode``: the original runtime ``forward_batch + .forward_mode``, which may be ``IDLE`` even when the captured + graph corresponds to ``DECODE``. DSV4's replay metadata prep + uses this for IDLE-batch substitution; other backends ignore it. + + This view subsumes the ``_replay_forward_batch`` side channel DSV4 + previously read out-of-band — step 04 swaps that mechanism for this + explicit fb_view field. + """ + return SimpleNamespace( + batch_size=bs, + forward_mode=capture_forward_mode, + actual_forward_mode=forward_batch.forward_mode, + input_ids=buffers.input_ids[:num_tokens], + req_pool_indices=buffers.req_pool_indices[:bs], + seq_lens=buffers.seq_lens[:bs], + seq_lens_sum=( + None + if forward_batch.seq_lens_sum is None + else forward_batch.seq_lens_sum + (bs - raw_bs) * seq_len_fill_value + ), + seq_lens_cpu=buffers.seq_lens_cpu[:bs], + encoder_lens=buffers.encoder_lens[:bs] if is_encoder_decoder else None, + out_cache_loc=getattr(forward_batch, "out_cache_loc", None), + spec_info=forward_batch.spec_info, + ) + + def _grouped_foreach_copy_(dsts: List[torch.Tensor], srcs: List[torch.Tensor]) -> None: """Call torch._foreach_copy_ grouped by (dst_dtype, src_dtype) pairs.""" @@ -1099,15 +1151,7 @@ class CudaGraphRunner: if lora_ids is not None: self.model_runner.lora_manager.prepare_lora_batch(forward_batch) - attn_backend.init_forward_metadata_capture_cuda_graph( - bs, - num_tokens, - req_pool_indices, - seq_lens, - encoder_lens, - forward_batch.forward_mode, - forward_batch.spec_info, - ) + attn_backend.init_forward_metadata_out_graph(forward_batch, in_capture=True) def run_once(): # Without this, warmup-1 caches the translation; the capture @@ -1116,6 +1160,10 @@ class CudaGraphRunner: if self.model_runner.is_hybrid_swa: self.model_runner.token_to_kv_pool.invalidate_loc_cache() + # Must run inside the capture block: warmup mutations here are + # undone by on_after_cuda_graph_warmup so capture starts clean. + attn_backend.init_forward_metadata_in_graph(forward_batch) + forward_batch.dp_local_start_pos = forward_batch.dp_local_num_tokens = ( None ) @@ -1269,25 +1317,17 @@ class CudaGraphRunner: attn_backend = self.model_runner.decode_attn_backend_group[stream_idx] else: attn_backend = self.attn_backend - # FIXME: implicit channel for backends (dsv4) that need forward_batch - # in replay metadata prep. Should become a real param on the interface. - attn_backend._replay_forward_batch = forward_batch - seq_lens_sum_arg = ( - None - if forward_batch.seq_lens_sum is None - else forward_batch.seq_lens_sum + (bs - raw_bs) * self.seq_len_fill_value + fb_view = build_replay_fb_view( + forward_batch=forward_batch, + buffers=buffers, + bs=bs, + raw_bs=raw_bs, + num_tokens=bs * self.num_tokens_per_bs, + seq_len_fill_value=self.seq_len_fill_value, + capture_forward_mode=self.capture_forward_mode, + is_encoder_decoder=self.is_encoder_decoder, ) - attn_backend.init_forward_metadata_replay_cuda_graph( - bs, - buffers.req_pool_indices[:bs], - buffers.seq_lens[:bs], - seq_lens_sum_arg, - buffers.encoder_lens[:bs] if self.is_encoder_decoder else None, - self.capture_forward_mode, - forward_batch.spec_info, - seq_lens_cpu=buffers.seq_lens_cpu[:bs], - ) - attn_backend._replay_forward_batch = None + attn_backend.init_forward_metadata_out_graph(fb_view) # Store fields self.raw_bs = raw_bs diff --git a/python/sglang/srt/model_executor/forward_batch_info.py b/python/sglang/srt/model_executor/forward_batch_info.py index 5f0fdf433..79759ac86 100644 --- a/python/sglang/srt/model_executor/forward_batch_info.py +++ b/python/sglang/srt/model_executor/forward_batch_info.py @@ -1233,6 +1233,46 @@ def enable_num_token_non_padded(): return get_moe_expert_parallel_world_size() > 1 +def build_inner_fb_view( + forward_batch: ForwardBatch, + *, + bs: int, + forward_mode: ForwardMode, + encoder_lens: Optional[torch.Tensor] = None, +): + """Build a ForwardBatch-like view for MultiStep draft wrapper dispatch. + + MultiStep draft wrappers (FlashInferMultiStepDraftBackend, + AiterMultiStepDraftBackend, TritonMultiStepDraftBackend, etc.) need + to dispatch to per-step inner backends' + :py:meth:`AttentionBackend.init_forward_metadata_out_graph` with an + overridden ``forward_mode`` (typically pinned to ``DECODE``) and + sometimes overridden ``encoder_lens``. The result is a thin + namespace mirroring just the fields backend init reads, avoiding + the cost of allocating a real ``ForwardBatch``. + + ``actual_forward_mode`` carries the original runtime + ``forward_batch.forward_mode`` (e.g., spec-decode draft) so backends + that check it for IDLE substitution (DSV4) see the unaltered value. + """ + from types import SimpleNamespace + + return SimpleNamespace( + batch_size=bs, + forward_mode=forward_mode, + actual_forward_mode=forward_batch.forward_mode, + input_ids=getattr(forward_batch, "input_ids", None), + positions=getattr(forward_batch, "positions", None), + req_pool_indices=forward_batch.req_pool_indices, + seq_lens=forward_batch.seq_lens, + seq_lens_sum=forward_batch.seq_lens_sum, + seq_lens_cpu=forward_batch.seq_lens_cpu, + encoder_lens=encoder_lens, + out_cache_loc=getattr(forward_batch, "out_cache_loc", None), + spec_info=forward_batch.spec_info, + ) + + class PPProxyTensors: # adapted from https://github.com/vllm-project/vllm/blob/d14e98d924724b284dc5eaf8070d935e214e50c0/vllm/sequence.py#L1103 tensors: Dict[str, torch.Tensor] 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 8af733255..bd15e2148 100644 --- a/python/sglang/srt/model_executor/piecewise_cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/piecewise_cuda_graph_runner.py @@ -419,7 +419,6 @@ class PiecewiseCudaGraphRunner: return_pooled_hidden_states=self.capture_return_pooled_hidden_states, ) - # Attention backend self.model_runner.attn_backend.init_forward_metadata(forward_batch) forward_batch.dp_local_start_pos = forward_batch.dp_local_num_tokens = None set_dp_buffer_len(None, num_tokens, forward_batch.dp_padding_mode.is_max_len()) @@ -798,7 +797,6 @@ class PiecewiseCudaGraphRunner: self.moe_fusions, dsa_indexers=self.dsa_indexers, ): - # Due to the dispatch kernel for MLA model, we init the metadata with original forward_batch self.model_runner.attn_backend.init_forward_metadata(forward_batch) output = self.model_runner.model.forward( static_forward_batch.input_ids, diff --git a/python/sglang/srt/models/deepseek_v4.py b/python/sglang/srt/models/deepseek_v4.py index 953a36b23..cbbfde35c 100644 --- a/python/sglang/srt/models/deepseek_v4.py +++ b/python/sglang/srt/models/deepseek_v4.py @@ -1582,9 +1582,6 @@ class DeepseekV4Model(nn.Module): for _attr in ("freqs_cis_c4", "freqs_cis_c128"): if hasattr(forward_batch, _attr): delattr(forward_batch, _attr) - # Upgrade lazy raw metadata on the main stream once before any layer - # forks alt-streams; later per-layer calls become no-ops. - get_attn_backend()._maybe_upgrade_forward_metadata() use_fused = self.use_fused_mhc_post_pre prev_residual, prev_post, prev_comb = None, None, None diff --git a/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py b/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py index 3c7482839..e7dec20da 100644 --- a/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py @@ -373,6 +373,8 @@ class EAGLEDraftCudaGraphRunner: if self.model_runner.is_hybrid_swa: self.model_runner.token_to_kv_pool.invalidate_loc_cache() + self.draft_attn_backend.init_forward_metadata_in_graph(forward_batch) + forward_batch.dp_local_start_pos = forward_batch.dp_local_num_tokens = None set_dp_buffer_len( global_dp_buffer_len, @@ -392,8 +394,8 @@ class EAGLEDraftCudaGraphRunner: return ret with forward_context(ForwardContext(attn_backend=self.draft_attn_backend)): - self.draft_attn_backend.init_forward_metadata_capture_cuda_graph( - forward_batch + self.draft_attn_backend.init_forward_metadata_out_graph( + forward_batch, in_capture=True ) self.deepep_adapter.capture(is_extend_in_batch=False) self._capture_init(run_once) @@ -509,9 +511,8 @@ class EAGLEDraftCudaGraphRunner: buffers.seq_lens_cpu[:raw_bs].copy_(forward_batch.seq_lens_cpu) forward_batch.seq_lens_cpu = buffers.seq_lens_cpu[:bs] - self.draft_attn_backend.init_forward_metadata_replay_cuda_graph( - forward_batch, bs - ) + # forward_batch.batch_size was overwritten to bs above when padding. + self.draft_attn_backend.init_forward_metadata_out_graph(forward_batch) self.raw_bs = raw_bs self.bs = bs # TODO: The forward_batch.seq_len_sum might need to be updated to reflect the padding in the cuda graph diff --git a/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py b/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py index a231ad738..aea21c332 100644 --- a/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py @@ -423,14 +423,8 @@ class EAGLEDraftExtendCudaGraphRunner: with forward_context( ForwardContext(attn_backend=self.draft_extend_attn_backend) ): - self.draft_extend_attn_backend.init_forward_metadata_capture_cuda_graph( - bs=bs, - num_tokens=num_tokens, - req_pool_indices=req_pool_indices, - seq_lens=seq_lens, - encoder_lens=None, - forward_mode=self.forward_mode, - spec_info=spec_info, + self.draft_extend_attn_backend.init_forward_metadata_out_graph( + forward_batch, in_capture=True ) self.deepep_adapter.capture(is_extend_in_batch=True) @@ -542,19 +536,24 @@ class EAGLEDraftExtendCudaGraphRunner: forward_batch.spec_info.num_correct_drafts = buffers.num_correct_drafts[:bs] forward_batch.spec_info.num_accept_tokens = buffers.num_accept_tokens[:bs] + from types import SimpleNamespace + seq_lens_sum = forward_batch.seq_lens_sum if seq_lens_sum is not None: seq_lens_sum = seq_lens_sum + (bs - raw_bs) * self.seq_len_fill_value - self.draft_extend_attn_backend.init_forward_metadata_replay_cuda_graph( - bs=bs, + fb_view = SimpleNamespace( + batch_size=bs, + forward_mode=self.forward_mode, + input_ids=getattr(forward_batch, "input_ids", None), req_pool_indices=buffers.req_pool_indices, seq_lens=buffers.seq_lens, seq_lens_sum=seq_lens_sum, - encoder_lens=None, - forward_mode=self.forward_mode, - spec_info=forward_batch.spec_info, seq_lens_cpu=buffers.seq_lens_cpu, + encoder_lens=None, + out_cache_loc=forward_batch.out_cache_loc, + spec_info=forward_batch.spec_info, ) + self.draft_extend_attn_backend.init_forward_metadata_out_graph(fb_view) # Replay self.raw_bs = raw_bs diff --git a/python/sglang/srt/speculative/frozen_kv_mtp_worker.py b/python/sglang/srt/speculative/frozen_kv_mtp_worker.py index 267d9f797..cf1a2da56 100644 --- a/python/sglang/srt/speculative/frozen_kv_mtp_worker.py +++ b/python/sglang/srt/speculative/frozen_kv_mtp_worker.py @@ -288,34 +288,33 @@ class FrozenKVMTPWorker(TpModelWorker): self, forward_batch: ForwardBatch ) -> None: with self._frozen_kv_target_view(forward_batch): - self.draft_attn_backend.init_forward_metadata_capture_cuda_graph( - forward_batch.batch_size, - forward_batch.positions.numel(), - forward_batch.req_pool_indices, - forward_batch.seq_lens, - encoder_lens=None, - forward_mode=ForwardMode.DECODE, - spec_info=None, + self.draft_attn_backend.init_forward_metadata_out_graph( + forward_batch, in_capture=True ) def _init_frozen_kv_metadata_replay_cuda_graph( self, forward_batch: ForwardBatch, bs: int, seq_lens_sum: int ) -> None: + from types import SimpleNamespace + + fb_view = SimpleNamespace( + batch_size=bs, + forward_mode=ForwardMode.DECODE, + input_ids=getattr(forward_batch, "input_ids", None), + req_pool_indices=forward_batch.req_pool_indices[:bs], + seq_lens=forward_batch.seq_lens[:bs], + seq_lens_sum=seq_lens_sum, + seq_lens_cpu=( + forward_batch.seq_lens_cpu[:bs] + if forward_batch.seq_lens_cpu is not None + else None + ), + encoder_lens=None, + out_cache_loc=getattr(forward_batch, "out_cache_loc", None), + spec_info=None, + ) with self._frozen_kv_target_view(forward_batch): - self.draft_attn_backend.init_forward_metadata_replay_cuda_graph( - bs, - forward_batch.req_pool_indices[:bs], - forward_batch.seq_lens[:bs], - seq_lens_sum, - encoder_lens=None, - forward_mode=ForwardMode.DECODE, - spec_info=None, - seq_lens_cpu=( - forward_batch.seq_lens_cpu[:bs] - if forward_batch.seq_lens_cpu is not None - else None - ), - ) + self.draft_attn_backend.init_forward_metadata_out_graph(fb_view) def init_cuda_graphs(self) -> None: if self.server_args.disable_cuda_graph or self.speculative_num_steps <= 1: diff --git a/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py b/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py index 30beb43b3..9e188afb8 100644 --- a/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py @@ -485,15 +485,7 @@ class MultiLayerEagleDraftExtendCudaGraphRunner: return ret with forward_context(ForwardContext(attn_backend=attn_backend)): - attn_backend.init_forward_metadata_capture_cuda_graph( - bs=bs, - num_tokens=num_tokens, - req_pool_indices=forward_batch.req_pool_indices, - seq_lens=forward_batch.seq_lens, - encoder_lens=None, - forward_mode=self.forward_mode, - spec_info=forward_batch.spec_info, - ) + attn_backend.init_forward_metadata_out_graph(forward_batch, in_capture=True) self.deepep_adapter.capture(is_extend_in_batch=True) self._capture_init(run_once) out = self._capture_graph( @@ -572,19 +564,24 @@ class MultiLayerEagleDraftExtendCudaGraphRunner: forward_batch.spec_info.positions = buffers.positions[:num_tokens] forward_batch.spec_info.extend_seq_lens_tensor = buffers.extend_seq_lens[:bs] - self.eagle_worker.draft_extend_attn_backend_list[ - self.step - ].init_forward_metadata_replay_cuda_graph( - bs=bs, + from types import SimpleNamespace + + fb_view = SimpleNamespace( + batch_size=bs, + forward_mode=self.forward_mode, + input_ids=getattr(forward_batch, "input_ids", None), req_pool_indices=buffers.req_pool_indices, seq_lens=buffers.seq_lens, seq_lens_sum=forward_batch.seq_lens_sum + (bs - raw_bs) * self.seq_len_fill_value, - encoder_lens=None, - forward_mode=self.forward_mode, - spec_info=forward_batch.spec_info, seq_lens_cpu=buffers.seq_lens_cpu, + encoder_lens=None, + out_cache_loc=forward_batch.out_cache_loc, + spec_info=forward_batch.spec_info, ) + self.eagle_worker.draft_extend_attn_backend_list[ + self.step + ].init_forward_metadata_out_graph(fb_view) # Replay self.raw_bs = raw_bs diff --git a/python/sglang/test/kits/attention_unittest/attention_methods/dsv4_attention.py b/python/sglang/test/kits/attention_unittest/attention_methods/dsv4_attention.py index 75cc359cb..830756db1 100644 --- a/python/sglang/test/kits/attention_unittest/attention_methods/dsv4_attention.py +++ b/python/sglang/test/kits/attention_unittest/attention_methods/dsv4_attention.py @@ -1079,7 +1079,6 @@ def _seed_c4_if_needed(fixture: DSV4AttentionFixture) -> None: compress_ratios. """ if fixture.case.compress_ratio == 4: - fixture.backend._maybe_upgrade_forward_metadata() _seed_c4_sparse_indices(fixture, num_entries=_DSV4_EXTRA_ENTRIES) @@ -1273,7 +1272,6 @@ def _pure_torch_dsv4_combined_reference( # `c4_sparse_page_indices` back to all -1 on the next upgrade) — the # reference must observe the same seeded indices the backend forward saw. _seed_c4_if_needed(fixture) - fixture.backend._maybe_upgrade_forward_metadata() md = fixture.backend.forward_metadata.core_metadata runner = fixture.runner max_context_len = runner.req_to_token_pool.req_to_token.shape[1] @@ -1554,9 +1552,6 @@ def run_dsv4_compress_attention_case( q_input, _ = fixture.actual_module.project(fixture.input_hidden) with torch.no_grad(), forward_context(ForwardContext(attn_backend=fixture.backend)): fixture.backend.init_forward_metadata(fixture.forward_batch) - # Trigger lazy upgrade so we can patch the metadata that the smoke - # case relies on (specifically c4_sparse_page_indices). - fixture.backend._maybe_upgrade_forward_metadata() if case.compress_ratio == 4: _seed_c4_sparse_indices(fixture, num_entries=extra_entries) actual = fixture.backend.forward( diff --git a/python/sglang/test/kits/attention_unittest/runner_modes/cuda_graph_decode_runner.py b/python/sglang/test/kits/attention_unittest/runner_modes/cuda_graph_decode_runner.py index 8cd74bda3..73ced7afd 100644 --- a/python/sglang/test/kits/attention_unittest/runner_modes/cuda_graph_decode_runner.py +++ b/python/sglang/test/kits/attention_unittest/runner_modes/cuda_graph_decode_runner.py @@ -291,36 +291,31 @@ def _init_cuda_graph_capture_metadata(backend, capture_batch_size: int, batch): max_bs=capture_batch_size, max_num_tokens=batch.input_ids.numel(), ) - backend.init_forward_metadata_capture_cuda_graph( - bs=capture_batch_size, - num_tokens=batch.input_ids.numel(), - req_pool_indices=batch.req_pool_indices, - seq_lens=batch.seq_lens, - encoder_lens=batch.encoder_lens, - forward_mode=batch.forward_mode, - spec_info=batch.spec_info, - ) + backend.init_forward_metadata_out_graph(batch, in_capture=True) + backend.init_forward_metadata_in_graph(batch) def _init_cuda_graph_replay_metadata(backend, capture_batch_size: int, batch): - # Some backends (e.g., `DeepseekV4AttnBackend`) read out-of-band attributes - # off the backend during replay metadata init — production wires this in - # `sglang/srt/model_executor/cuda_graph_runner.py:1234`. Mirror that - # contract so backends that don't use it just store-and-clear the field. - backend._replay_forward_batch = batch - try: - backend.init_forward_metadata_replay_cuda_graph( - bs=capture_batch_size, - req_pool_indices=batch.req_pool_indices, - seq_lens=batch.seq_lens, - seq_lens_sum=batch.seq_lens_sum, - encoder_lens=batch.encoder_lens, - forward_mode=batch.forward_mode, - spec_info=batch.spec_info, - seq_lens_cpu=batch.seq_lens_cpu, - ) - finally: - backend._replay_forward_batch = None + from types import SimpleNamespace + + fb_view = SimpleNamespace( + batch_size=capture_batch_size, + forward_mode=batch.forward_mode, + actual_forward_mode=batch.forward_mode, + input_ids=batch.input_ids, + positions=getattr(batch, "positions", None), + req_pool_indices=batch.req_pool_indices, + seq_lens=batch.seq_lens, + seq_lens_sum=batch.seq_lens_sum, + seq_lens_cpu=batch.seq_lens_cpu, + encoder_lens=batch.encoder_lens, + out_cache_loc=getattr(batch, "out_cache_loc", None), + spec_info=batch.spec_info, + ) + backend.init_forward_metadata_out_graph(fb_view) + # No real cuda graph here, so run the in-graph step explicitly to produce + # the Full metadata the forward path expects (no-op for non-DSV4). + backend.init_forward_metadata_in_graph(fb_view) # Best-effort metadata-shape sanity check — catches negative kv_lens and # non-monotonic indptr that would otherwise leave real-row output correct # but corrupt padded-row scratch state. See `metadata_invariants.py`. diff --git a/python/sglang/test/kits/attention_unittest/runner_modes/speculative_draft_extend_runner.py b/python/sglang/test/kits/attention_unittest/runner_modes/speculative_draft_extend_runner.py index 819282d36..dfcaaa5c7 100644 --- a/python/sglang/test/kits/attention_unittest/runner_modes/speculative_draft_extend_runner.py +++ b/python/sglang/test/kits/attention_unittest/runner_modes/speculative_draft_extend_runner.py @@ -1022,13 +1022,6 @@ class EagleDraftExtendCudaGraphRunnerAdapter: make_forward_batch: Callable[ [Any, Any, Any, EagleDraftRunnerSettings], ForwardBatch ] - # Optional hook invoked with `(draft_extend_attn_backend, batch)` right - # before `graph_runner.replay(batch)`. DSV4 needs this to set the - # out-of-band `_replay_forward_batch` attribute that - # `DeepseekV4AttnBackend.init_forward_metadata_replay_cuda_graph` reads - # (the multi-step DECODE wrapper sets it internally, but the single- - # backend DRAFT_EXTEND path does not). - pre_replay: Callable[[Any, ForwardBatch], None] = None check_case: Callable[[Any, EagleDraftRunnerSettings], None] = ( lambda _case, _settings: None ) @@ -1268,12 +1261,7 @@ def run_eagle_draft_extend_cuda_graph_runner_case( adapter.prepare_replay_state(graph_fixture, case, draft_inputs, settings) testcase.assertTrue(graph_runner.can_run(graph_batch)) - if adapter.pre_replay is not None: - adapter.pre_replay(graph_backend, graph_batch) actual = graph_runner.replay(graph_batch) - if adapter.pre_replay is not None: - # Best-effort cleanup of any out-of-band state pre_replay set. - adapter.pre_replay(graph_backend, None) adapter.assert_outputs_close(actual, expected, settings) finally: _reset_cuda_graph_test_buffers() @@ -1934,21 +1922,6 @@ def _dsv4_assert_draft_extend_outputs_close(actual, expected, settings) -> None: ) -def _dsv4_draft_extend_pre_replay( - draft_extend_attn_backend, - batch: ForwardBatch | None, -) -> None: - """Set/clear the out-of-band `_replay_forward_batch` attribute that - `DeepseekV4AttnBackend.init_forward_metadata_replay_cuda_graph` reads. - - The DSV4 multi-step DECODE wrapper sets this internally - (`deepseek_v4_backend.py:1231,1242`), but the single-backend DRAFT_EXTEND - path used by `_create_dsv4_prefill_backend` does not. Set before - `replay()` and clear afterwards to mimic the multi-step pattern. - """ - draft_extend_attn_backend._replay_forward_batch = batch - - def run_dsv4_eagle_draft_extend_cuda_graph_runner_case( testcase, case: DSV4AttentionCase, @@ -2001,7 +1974,6 @@ def run_dsv4_eagle_draft_extend_cuda_graph_runner_case( make_draft_inputs=_make_dsv4_draft_extend_inputs, prepare_replay_state=_prepare_dsv4_draft_extend_replay_state, make_forward_batch=_make_dsv4_eagle_draft_extend_forward_batch, - pre_replay=_dsv4_draft_extend_pre_replay, assert_outputs_close=_dsv4_assert_draft_extend_outputs_close, ) run_eagle_draft_extend_cuda_graph_runner_case( diff --git a/test/manual/attention/test_trtllm_mla_backend.py b/test/manual/attention/test_trtllm_mla_backend.py index 25cfe2a6d..5470b9053 100755 --- a/test/manual/attention/test_trtllm_mla_backend.py +++ b/test/manual/attention/test_trtllm_mla_backend.py @@ -1,5 +1,6 @@ import math import unittest +from types import SimpleNamespace import numpy as np import torch @@ -1030,15 +1031,15 @@ class TestTRTLLMMLA(CustomTestCase): ) req_pool_indices = torch.arange(batch_size, device=config["device"]) - backend.init_forward_metadata_capture_cuda_graph( - bs=batch_size, - num_tokens=batch_size, + capture_fb = SimpleNamespace( + batch_size=batch_size, + forward_mode=ForwardMode.DECODE, req_pool_indices=req_pool_indices, seq_lens=seq_lens, - encoder_lens=None, - forward_mode=ForwardMode.DECODE, + positions=torch.arange(batch_size, device=config["device"]), spec_info=None, ) + backend.init_forward_metadata_out_graph(capture_fb, in_capture=True) # Verify capture metadata self.assertIn(batch_size, backend.decode_cuda_graph_metadata) @@ -1054,16 +1055,15 @@ class TestTRTLLMMLA(CustomTestCase): device=config["device"], ) - backend.init_forward_metadata_replay_cuda_graph( - bs=batch_size, + replay_fb = SimpleNamespace( + batch_size=batch_size, + forward_mode=ForwardMode.DECODE, req_pool_indices=req_pool_indices, seq_lens=new_seq_lens, - seq_lens_sum=new_seq_lens.sum().item(), - encoder_lens=None, - forward_mode=ForwardMode.DECODE, - spec_info=None, seq_lens_cpu=new_seq_lens.cpu(), + spec_info=None, ) + backend.init_forward_metadata_out_graph(replay_fb) # Verify replay updated the metadata replay_metadata = backend.forward_decode_metadata diff --git a/test/registered/attention/unittests/dense/test_extend_init_contract.py b/test/registered/attention/unittests/dense/test_extend_init_contract.py new file mode 100644 index 000000000..06eb9f123 --- /dev/null +++ b/test/registered/attention/unittests/dense/test_extend_init_contract.py @@ -0,0 +1,147 @@ +"""Unit tests for the init contract that piecewise + breakable cuda graph +capture rely on with ``forward_mode=EXTEND``. + +The piecewise + breakable runners (unlike the full ``cuda_graph_runner``) +capture prefill chunks, which means their capture path passes +``forward_mode=ForwardMode.EXTEND`` to ``attn_backend.init_forward_metadata*``. + +Backends like FlashInfer and FA3 implement two separate init bodies: + +- ``init_forward_metadata(fb)`` — the eager entry. Handles all modes + including plain ``EXTEND`` (full prefill / chunked prefill). +- ``init_forward_metadata_out_graph(fb, in_capture=True)`` — the + bucket-keyed wrapper prep used by the full cuda graph runner. Only + handles modes the full runner captures: ``DECODE`` / ``IDLE`` / + ``TARGET_VERIFY`` / ``DRAFT_EXTEND`` / ``DLLM_EXTEND``. Plain + ``EXTEND`` is not in scope here and the body raises on it. + +These tests pin both halves of that contract so a future refactor that +incorrectly routes piecewise/breakable capture through ``_out_graph(in_capture=True)`` +fails at unit-test time instead of at GPU CI / e2e time. This is the +specific regression #26735 introduced and then fixed +(see ``piecewise_cuda_graph_runner.py`` / +``breakable_cuda_graph_runner.py`` capture sites). +""" + +import sys +import unittest +from pathlib import Path + +import torch + +from sglang.srt.model_executor.forward_batch_info import ForwardMode +from sglang.srt.utils import get_device_sm +from sglang.test.test_utils import CustomTestCase + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) + +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.kits.attention_unittest.attention_methods.dense_attention import ( + DenseAttentionCase, + build_dense_attention_fixture, +) + +register_cuda_ci(est_time=10, stage="base-a", runner_config="1-gpu-small") + +_EXTEND_CASE = DenseAttentionCase( + name="extend_no_prefix_smoke", + backend="fa3", # overridden per backend below + forward_mode=ForwardMode.EXTEND, + num_heads=4, + num_kv_heads=4, + page_size=16, + prefix_lens=(0,), + extend_lens=(16,), +) + + +@unittest.skipIf(not torch.cuda.is_available(), "CUDA is required") +class TestExtendInitContract(CustomTestCase): + + def _make_case(self, backend: str) -> DenseAttentionCase: + return DenseAttentionCase( + name=f"extend_no_prefix_{backend}", + backend=backend, + forward_mode=_EXTEND_CASE.forward_mode, + num_heads=_EXTEND_CASE.num_heads, + num_kv_heads=_EXTEND_CASE.num_kv_heads, + page_size=_EXTEND_CASE.page_size, + prefix_lens=_EXTEND_CASE.prefix_lens, + extend_lens=_EXTEND_CASE.extend_lens, + ) + + def _build_fixture(self, backend: str, *, head_dim: int = 16): + case = self._make_case(backend) + try: + return build_dense_attention_fixture(self, case, head_dim=head_dim) + except (AssertionError, ImportError, ModuleNotFoundError) as exc: + self.skipTest(f"backend {backend} unavailable: {exc}") + + def _assert_extend_eager_init_well_formed( + self, backend: str, *, head_dim: int = 16 + ): + fixture = self._build_fixture(backend, head_dim=head_dim) + fixture.backend.init_forward_metadata(fixture.forward_batch) + meta = fixture.backend.forward_metadata + self.assertIsNotNone( + meta, + f"{backend}: init_forward_metadata(EXTEND) left forward_metadata as None — " + "piecewise/breakable capture would crash on the first forward_extend call.", + ) + page_table = getattr(meta, "page_table", None) + if page_table is not None: + self.assertIsInstance(page_table, torch.Tensor) + + @unittest.skipIf( + get_device_sm() >= 100 or get_device_sm() < 80, + "FA3 backend requires SM 80-90", + ) + def test_fa3_extend_eager_init(self): + self._assert_extend_eager_init_well_formed("fa3") + + def test_flashinfer_extend_eager_init(self): + # FlashInfer's JIT prefill kernel needs head_dim ≥ 64. + self._assert_extend_eager_init_well_formed("flashinfer", head_dim=128) + + def test_triton_extend_eager_init(self): + self._assert_extend_eager_init_well_formed("triton") + + def _assert_out_graph_in_capture_rejects_extend( + self, backend: str, *, head_dim: int = 16 + ): + """``init_forward_metadata_out_graph(fb, in_capture=True)`` is the + bucket-prep path. Plain EXTEND mode isn't in its supported set — + piecewise/breakable capture used to route through here and crash. + This pins the constraint: if the path starts handling EXTEND + (raises become passes), revisit the piecewise/breakable capture + wiring and consider unifying the API. + """ + fixture = self._build_fixture(backend, head_dim=head_dim) + try: + fixture.backend.init_forward_metadata_out_graph( + fixture.forward_batch, in_capture=True + ) + except (ValueError, AttributeError, KeyError, AssertionError): + return + meta = fixture.backend.forward_metadata + self.assertIsNotNone( + meta, + f"{backend}: _out_graph(in_capture=True) accepted EXTEND but left " + "forward_metadata as None — this is the specific bug pattern that " + "broke piecewise capture (FA3._apply_cuda_graph_metadata falls " + "through for unsupported modes and sets self.forward_metadata = None).", + ) + + @unittest.skipIf( + get_device_sm() >= 100 or get_device_sm() < 80, + "FA3 backend requires SM 80-90", + ) + def test_fa3_out_graph_capture_rejects_extend(self): + self._assert_out_graph_in_capture_rejects_extend("fa3") + + def test_flashinfer_out_graph_capture_rejects_extend(self): + self._assert_out_graph_in_capture_rejects_extend("flashinfer", head_dim=128) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/attention/unittests/dense/test_tbo.py b/test/registered/attention/unittests/dense/test_tbo.py index a5867451f..603b6c5c5 100644 --- a/test/registered/attention/unittests/dense/test_tbo.py +++ b/test/registered/attention/unittests/dense/test_tbo.py @@ -94,14 +94,15 @@ class TestTboAttnDenseAttentionBackendCorrectness(CustomTestCase): "FA3 backend requires SM 80-90", ) def test_tbo_target_verify_cuda_graph_capture_delegates_to_primary_capture(self): - """TBO capture must invoke ``primary.init_forward_metadata_capture_cuda_graph``, - not ``primary.init_forward_metadata_replay_cuda_graph``. + """TBO capture must dispatch primary's + ``init_forward_metadata_out_graph(fb, in_capture=True)`` (the capture + path), not the replay path. Backends like FlashAttention store per-bs metadata in dicts populated - only by their capture path (via ``_bind_metadata_buffers``). If TBO - short-circuits its capture to its own replay (which delegates to - ``primary.replay``), those dicts are empty and replay raises - ``KeyError: bs``. Reproduces the deepep-4-gpu-h100 failure where + only by the in_capture=True branch (via ``_bind_metadata_buffers``). + If TBO short-circuits its capture to its own replay path, those dicts + are empty and replay raises ``KeyError: bs``. Reproduces the + deepep-4-gpu-h100 failure where ``flashattention_backend.target_verify_metadata[bs]`` lookup blew up during ``init_device_graphs``. @@ -126,19 +127,100 @@ class TestTboAttnDenseAttentionBackendCorrectness(CustomTestCase): capture_bs = case.batch_size num_tokens = sum(case.extend_lens) wrapper.init_cuda_graph_state(max_bs=capture_bs, max_num_tokens=num_tokens) - # This is the failing call before the fix: TBO.capture delegating to - # primary.replay (instead of primary.capture) reads an unpopulated - # ``target_verify_metadata[bs]`` dict and raises KeyError. - wrapper.init_forward_metadata_capture_cuda_graph( - bs=capture_bs, - num_tokens=num_tokens, + wrapper.init_forward_metadata_out_graph(batch, in_capture=True) + + @unittest.skipIf( + get_device_sm() >= 100 or get_device_sm() < 80, + "FA3 backend requires SM 80-90", + ) + def test_tbo_target_verify_cuda_graph_replay_splits_children(self): + """TBO replay must split the padded capture-time buffers into per-child + views before dispatching to ``init_forward_metadata_out_graph(fb_view, + in_capture=False)``. + + ``cuda_graph_runner.replay_prepare`` constructs a ``SimpleNamespace`` + fb_view via ``build_replay_fb_view`` — it has no ``tbo_children`` + attribute because ``tbo_plugin.replay_prepare`` does not call + ``prepare_raw``. Without the in-line split, TBO either crashes + (AttributeError on ``fb_view.tbo_children``) or silently leaves child + metadata stale. + + Asserts both children's ``init_forward_metadata_out_graph`` is invoked + with sliced ``req_pool_indices`` / ``seq_lens`` / ``seq_lens_cpu`` + whose lengths match the split derived from + ``compute_split_indices_for_cuda_graph_replay``. + """ + from types import SimpleNamespace + from unittest.mock import MagicMock + + from sglang.srt.batch_overlap.two_batch_overlap import ( + compute_split_indices_for_cuda_graph_replay, + ) + + case = self.TARGET_VERIFY_CAPTURE_CASE + fixture = self._build_and_wrap(case) + wrapper = fixture.backend + batch = fixture.forward_batch + _prepare_spec_verify_batch( + case, + batch, + topk=1, + spec_kind="eagle", + device=str(batch.seq_lens.device), + ) + + capture_bs = case.batch_size + num_tokens_per_bs = sum(case.extend_lens) // capture_bs + num_tokens = capture_bs * num_tokens_per_bs + split_seq_index, split_token_index = ( + compute_split_indices_for_cuda_graph_replay( + forward_mode=batch.forward_mode, + cuda_graph_num_tokens=num_tokens, + spec_info=batch.spec_info, + ) + ) + self.assertGreater(split_seq_index, 0) + self.assertLess(split_seq_index, capture_bs) + + # fb_view shaped like build_replay_fb_view's output: a SimpleNamespace + # with no tbo_children attribute. + fb_view = SimpleNamespace( + batch_size=capture_bs, + forward_mode=batch.forward_mode, + actual_forward_mode=batch.forward_mode, + input_ids=batch.input_ids, req_pool_indices=batch.req_pool_indices, seq_lens=batch.seq_lens, - encoder_lens=batch.encoder_lens, - forward_mode=batch.forward_mode, + seq_lens_sum=int(batch.seq_lens_cpu.sum()), + seq_lens_cpu=batch.seq_lens_cpu, + encoder_lens=None, + out_cache_loc=batch.out_cache_loc, spec_info=batch.spec_info, ) + # Pure mocks (no `wraps=...`) so the dispatcher's slicing/contract is + # observed without invoking real backend bodies. + primary_mock = MagicMock() + child_mocks = [MagicMock(), MagicMock()] + wrapper.primary = primary_mock + wrapper.children = child_mocks + + wrapper.init_forward_metadata_out_graph(fb_view, in_capture=False) + + primary_mock.init_forward_metadata_out_graph.assert_called_once() + for child_mock in child_mocks: + child_mock.init_forward_metadata_out_graph.assert_called_once() + child_fbs = [ + m.init_forward_metadata_out_graph.call_args.kwargs["forward_batch"] + for m in child_mocks + ] + self.assertEqual(child_fbs[0].batch_size, split_seq_index) + self.assertEqual(child_fbs[1].batch_size, capture_bs - split_seq_index) + self.assertEqual(child_fbs[0].req_pool_indices.shape[0], split_seq_index) + self.assertEqual( + child_fbs[1].req_pool_indices.shape[0], capture_bs - split_seq_index + ) + if __name__ == "__main__": unittest.main() diff --git a/test/registered/attention/unittests/gdn/test_triton.py b/test/registered/attention/unittests/gdn/test_triton.py index ffddcc91a..d20a62864 100644 --- a/test/registered/attention/unittests/gdn/test_triton.py +++ b/test/registered/attention/unittests/gdn/test_triton.py @@ -399,38 +399,35 @@ class TestTritonGDNBackendCorrectness(CustomTestCase): linear_attn_backend.init_forward_metadata, sentinel_forward_batch ) + def _make_sentinel_fb(self): + return SimpleNamespace( + batch_size=3, + forward_mode=ForwardMode.DECODE, + req_pool_indices=object(), + seq_lens=object(), + seq_lens_cpu=object(), + seq_lens_sum=42, + spec_info=object(), + encoder_lens=None, + positions=object(), + input_ids=object(), + out_cache_loc=None, + ) + def test_hybrid_dispatch_replay_init_forward_metadata_fan_out(self): backend, full_attn_backend, linear_attn_backend = ( self._make_dispatch_spy_backend() ) - sentinel_req_pool = object() - sentinel_seq_lens = object() - sentinel_seq_lens_cpu = object() - sentinel_spec_info = object() - - backend.init_forward_metadata_replay_cuda_graph( - bs=3, - req_pool_indices=sentinel_req_pool, - seq_lens=sentinel_seq_lens, - seq_lens_sum=42, - encoder_lens=None, - forward_mode=ForwardMode.DECODE, - spec_info=sentinel_spec_info, - seq_lens_cpu=sentinel_seq_lens_cpu, - ) + fb = self._make_sentinel_fb() + backend.init_forward_metadata_out_graph(fb) # We assert sentinel identity rather than exact (args, kwargs) shape # so a positional↔keyword refactor inside `HybridLinearAttnBackend` - # doesn't trip the test as long as the values still flow through. + # doesn't trip the test as long as the fb still flows through. for sub_backend in (full_attn_backend, linear_attn_backend): self._assert_fanout_forwarded( - sub_backend.init_forward_metadata_replay_cuda_graph, - sentinel_req_pool, - sentinel_seq_lens, - sentinel_seq_lens_cpu, - sentinel_spec_info, - ForwardMode.DECODE, + sub_backend.init_forward_metadata_out_graph, fb ) def test_hybrid_dispatch_capture_init_forward_metadata_fan_out(self): @@ -439,27 +436,12 @@ class TestTritonGDNBackendCorrectness(CustomTestCase): backend, full_attn_backend, linear_attn_backend = ( self._make_dispatch_spy_backend() ) - sentinel_req_pool = object() - sentinel_seq_lens = object() - sentinel_spec_info = object() - - backend.init_forward_metadata_capture_cuda_graph( - bs=3, - num_tokens=3, - req_pool_indices=sentinel_req_pool, - seq_lens=sentinel_seq_lens, - encoder_lens=None, - forward_mode=ForwardMode.DECODE, - spec_info=sentinel_spec_info, - ) + fb = self._make_sentinel_fb() + backend.init_forward_metadata_out_graph(fb, in_capture=True) for sub_backend in (full_attn_backend, linear_attn_backend): self._assert_fanout_forwarded( - sub_backend.init_forward_metadata_capture_cuda_graph, - sentinel_req_pool, - sentinel_seq_lens, - sentinel_spec_info, - ForwardMode.DECODE, + sub_backend.init_forward_metadata_out_graph, fb ) diff --git a/test/registered/attention/unittests/mamba/test_mamba2.py b/test/registered/attention/unittests/mamba/test_mamba2.py index 8fa154842..4b62df96a 100644 --- a/test/registered/attention/unittests/mamba/test_mamba2.py +++ b/test/registered/attention/unittests/mamba/test_mamba2.py @@ -216,7 +216,7 @@ class TestTritonMamba2BackendCorrectness(CustomTestCase): # padding the `forward_batch.input_ids` / `out_cache_loc`. def test_mamba2_replay_metadata_padding_indices(self): - # Drive `init_forward_metadata_replay_cuda_graph` directly with + # Drive `init_forward_metadata_out_graph` (replay path) directly with # `seq_lens_cpu=[5, 1, 1]` (two trailing rows at the cuda-graph # fill value 1) so the padding-row count is observable in # `state_indices_list[bs - 1]`. @@ -245,16 +245,17 @@ class TestTritonMamba2BackendCorrectness(CustomTestCase): req_pool_indices ] = torch.tensor([7, 0, 0], dtype=torch.int32, device=device) - backend.init_forward_metadata_replay_cuda_graph( - bs=bs, + fb = SimpleNamespace( + batch_size=bs, + forward_mode=ForwardMode.DECODE, req_pool_indices=req_pool_indices, seq_lens=seq_lens, - seq_lens_sum=int(seq_lens_cpu.sum().item()), - encoder_lens=None, - forward_mode=ForwardMode.DECODE, - spec_info=None, seq_lens_cpu=seq_lens_cpu, + seq_lens_sum=int(seq_lens_cpu.sum().item()), + spec_info=None, + encoder_lens=None, ) + backend.init_forward_metadata_out_graph(fb) state_indices = backend.state_indices_list[bs - 1].cpu().tolist() self.assertEqual( @@ -326,62 +327,44 @@ class TestTritonMamba2BackendCorrectness(CustomTestCase): linear_attn_backend.init_forward_metadata, sentinel_forward_batch ) + def _make_sentinel_fb(self): + return SimpleNamespace( + batch_size=3, + forward_mode=ForwardMode.DECODE, + req_pool_indices=object(), + seq_lens=object(), + seq_lens_cpu=object(), + seq_lens_sum=42, + spec_info=object(), + encoder_lens=None, + positions=object(), + input_ids=object(), + out_cache_loc=None, + ) + def test_hybrid_dispatch_replay_init_forward_metadata_fan_out(self): backend, full_attn_backend, linear_attn_backend = ( self._make_dispatch_spy_backend() ) - sentinel_req_pool = object() - sentinel_seq_lens = object() - sentinel_seq_lens_cpu = object() - sentinel_spec_info = object() - - backend.init_forward_metadata_replay_cuda_graph( - bs=3, - req_pool_indices=sentinel_req_pool, - seq_lens=sentinel_seq_lens, - seq_lens_sum=42, - encoder_lens=None, - forward_mode=ForwardMode.DECODE, - spec_info=sentinel_spec_info, - seq_lens_cpu=sentinel_seq_lens_cpu, - ) + fb = self._make_sentinel_fb() + backend.init_forward_metadata_out_graph(fb) for sub_backend in (full_attn_backend, linear_attn_backend): self._assert_fanout_forwarded( - sub_backend.init_forward_metadata_replay_cuda_graph, - sentinel_req_pool, - sentinel_seq_lens, - sentinel_seq_lens_cpu, - sentinel_spec_info, - ForwardMode.DECODE, + sub_backend.init_forward_metadata_out_graph, fb ) def test_hybrid_dispatch_capture_init_forward_metadata_fan_out(self): backend, full_attn_backend, linear_attn_backend = ( self._make_dispatch_spy_backend() ) - sentinel_req_pool = object() - sentinel_seq_lens = object() - sentinel_spec_info = object() - - backend.init_forward_metadata_capture_cuda_graph( - bs=3, - num_tokens=3, - req_pool_indices=sentinel_req_pool, - seq_lens=sentinel_seq_lens, - encoder_lens=None, - forward_mode=ForwardMode.DECODE, - spec_info=sentinel_spec_info, - ) + fb = self._make_sentinel_fb() + backend.init_forward_metadata_out_graph(fb, in_capture=True) for sub_backend in (full_attn_backend, linear_attn_backend): self._assert_fanout_forwarded( - sub_backend.init_forward_metadata_capture_cuda_graph, - sentinel_req_pool, - sentinel_seq_lens, - sentinel_spec_info, - ForwardMode.DECODE, + sub_backend.init_forward_metadata_out_graph, fb )