[refactor] init_forward_metadata 3-method ABC + side-channel removal + ForwardMetadata type rename (#26735)

Co-authored-by: Claude Opus 4.7 <noreply@anthropic.com>
This commit is contained in:
Cheng Wan
2026-06-02 10:33:33 -07:00
committed by GitHub
co-authored by Claude Opus 4.7
parent 6c69756fa8
commit 99da43b900
42 changed files with 1999 additions and 2081 deletions
@@ -362,6 +362,29 @@ class AscendAttnBackend(AttentionBackend):
): ):
pass 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): def init_forward_metadata(self, forward_batch: ForwardBatch):
"""Init the metadata for a forward pass.""" """Init the metadata for a forward pass."""
self.forward_metadata = ForwardMetadata() self.forward_metadata = ForwardMetadata()
@@ -538,39 +561,19 @@ class AscendAttnBackend(AttentionBackend):
self.graph_metadata[bs] = metadata self.graph_metadata[bs] = metadata
return metadata return metadata
def init_forward_metadata_capture_cuda_graph( def _apply_cuda_graph_metadata(
self, self,
bs: int, bs: int,
num_tokens: int,
req_pool_indices: torch.Tensor, req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor, seq_lens: torch.Tensor,
encoder_lens: Optional[torch.Tensor], seq_lens_cpu: torch.Tensor,
forward_mode: ForwardMode, forward_mode: ForwardMode,
spec_info: Optional[SpecInput], spec_info: Optional[SpecInput],
): ):
self._init_cuda_graph_metadata(bs, forward_mode, seq_lens) """Shared capture+replay body for the cuda-graph init path.
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( Public entry: :py:meth:`init_forward_metadata_out_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],
):
metadata = self.graph_metadata[bs] metadata = self.graph_metadata[bs]
max_len = seq_lens_cpu[:bs].max().item() max_len = seq_lens_cpu[:bs].max().item()
if forward_mode.is_target_verify(): if forward_mode.is_target_verify():
@@ -2395,6 +2398,32 @@ class AscendAttnMultiStepDraftBackend:
for i in range(self.speculative_num_steps - 1): for i in range(self.speculative_num_steps - 1):
call_fn(i, forward_batch) 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 init_forward_metadata(self, forward_batch: ForwardBatch):
def call_fn(i, forward_batch): def call_fn(i, forward_batch):
assert forward_batch.spec_info is not None 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): def init_cuda_graph_state(self, max_bs, max_num_tokens):
for i in range(self.speculative_num_steps): for i in range(self.speculative_num_steps):
self.attn_backends[i].init_cuda_graph_state(max_bs, max_num_tokens) 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)
@@ -72,6 +72,21 @@ class AscendGDNAttnBackend(AscendMambaAttnBackendBase):
else: else:
self.ssm_state_indices = cache_indices 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): def init_forward_metadata(self, forward_batch: ForwardBatch):
if forward_batch.forward_mode.is_draft_extend(True): if forward_batch.forward_mode.is_draft_extend(True):
return return
@@ -83,53 +98,6 @@ class AscendGDNAttnBackend(AscendMambaAttnBackendBase):
) )
self.graph_mode = False 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( def forward_decode(
self, self,
layer: RadixLinearAttention, layer: RadixLinearAttention,
@@ -131,6 +131,10 @@ class AscendMambaAttnBackendBase(MambaAttnBackendBase):
spec_info: Optional[SpecInput], spec_info: Optional[SpecInput],
seq_lens_cpu: Optional[torch.Tensor], seq_lens_cpu: Optional[torch.Tensor],
): ):
# 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( num_padding = torch.count_nonzero(
seq_lens_cpu == self.get_cuda_graph_seq_len_fill_value() seq_lens_cpu == self.get_cuda_graph_seq_len_fill_value()
) )
@@ -820,6 +820,25 @@ class AiterAttnBackend(AttentionBackend):
) )
return output 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): def init_forward_metadata(self, forward_batch: ForwardBatch):
"""Init auxiliary variables for aiter attention backend.""" """Init auxiliary variables for aiter attention backend."""
@@ -1482,28 +1501,7 @@ class AiterAttnBackend(AttentionBackend):
device=self.device, device=self.device,
) )
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[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(
self, self,
bs: int, bs: int,
req_pool_indices: torch.Tensor, 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] 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 init_forward_metadata_out_graph(
def call_fn(i, forward_batch): self,
self.attn_backends[i].init_forward_metadata_capture_cuda_graph( forward_batch: ForwardBatch,
forward_batch.batch_size, in_capture: bool = False,
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 call_fn(i, forward_batch): from sglang.srt.model_executor.forward_batch_info import build_inner_fb_view
self.attn_backends[i].init_forward_metadata_replay_cuda_graph(
bs, inner_fb = build_inner_fb_view(
forward_batch.req_pool_indices, forward_batch,
forward_batch.seq_lens, bs=forward_batch.batch_size,
seq_lens_sum=-1,
encoder_lens=None,
forward_mode=ForwardMode.DECODE, forward_mode=ForwardMode.DECODE,
spec_info=forward_batch.spec_info, )
seq_lens_cpu=forward_batch.seq_lens_cpu,
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) 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)
@@ -1,6 +1,6 @@
from __future__ import annotations from __future__ import annotations
from abc import ABC, abstractmethod from abc import ABC
from typing import TYPE_CHECKING, Optional from typing import TYPE_CHECKING, Optional
import torch import torch
@@ -11,52 +11,84 @@ from sglang.srt.utils.common import is_npu
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.layers.attention.dsa.dsa_indexer import BaseIndexerMetadata from sglang.srt.layers.attention.dsa.dsa_indexer import BaseIndexerMetadata
from sglang.srt.layers.radix_attention import RadixAttention 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 from sglang.srt.speculative.spec_info import SpecInput
class AttentionBackend(ABC): 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. # Opt out only when this backend never reads seq_lens_cpu / seq_lens_sum.
needs_cpu_seq_lens: bool = True 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): def init_cuda_graph_state(self, max_bs: int, max_num_tokens: int):
"""Init the global shared states for cuda graph.""" """Init the global shared states for cuda graph."""
raise NotImplementedError() 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): def get_cuda_graph_seq_len_fill_value(self):
"""Get the fill value for padded seq lens. Typically, it is 0 or 1.""" """Get the fill value for padded seq lens. Typically, it is 0 or 1."""
raise NotImplementedError() raise NotImplementedError()
@@ -14,13 +14,12 @@ import triton
from sglang.srt.layers.attention.flashinfer_mla_backend import FlashInferMLAAttnBackend 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.attention.utils import create_flashmla_kv_indices_triton
from sglang.srt.layers.dp_attention import get_attention_tp_size 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 from sglang.srt.utils import is_cuda
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.layers.radix_attention import RadixAttention from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.model_executor.model_runner import ModelRunner from sglang.srt.model_executor.model_runner import ModelRunner
from sglang.srt.speculative.spec_info import SpecInput
_is_cuda = is_cuda() _is_cuda = is_cuda()
if _is_cuda: if _is_cuda:
@@ -79,6 +78,37 @@ class CutlassMLABackend(FlashInferMLAAttnBackend):
self.q_data_type = model_runner.dtype self.q_data_type = model_runner.dtype
self.kv_cache_dim = self.kv_lora_rank + self.qk_rope_head_dim 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): def init_forward_metadata(self, forward_batch: ForwardBatch):
bs = forward_batch.batch_size bs = forward_batch.batch_size
@@ -143,77 +173,6 @@ class CutlassMLABackend(FlashInferMLAAttnBackend):
) )
self.cuda_graph_kv_indices = cuda_graph_kv_indices 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): def get_cuda_graph_seq_len_fill_value(self):
return 1 return 1
@@ -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.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode 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.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 import ceil_align
from sglang.srt.utils.common import is_sm120_supported from sglang.srt.utils.common import is_sm120_supported
@@ -387,7 +386,6 @@ class DeepseekV4AttnBackend(
DSV4RawVerifyMetadata, DSV4RawVerifyMetadata,
DSV4RawDecodeMetadata, DSV4RawDecodeMetadata,
] = None ] = None
self._replay_forward_batch: Optional[ForwardBatch] = None # FIXME: out-of-band
def _move_to_device(self, x: List[int]) -> torch.Tensor: def _move_to_device(self, x: List[int]) -> torch.Tensor:
pin_tensor = torch.tensor(x, dtype=torch.int32, pin_memory=True) 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, 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: def init_forward_metadata(self, forward_batch: ForwardBatch) -> None:
if self.mtp_enabled and forward_batch.forward_mode.is_idle(): if self.mtp_enabled and forward_batch.forward_mode.is_idle():
return return
@@ -686,8 +814,7 @@ class DeepseekV4AttnBackend(
if forward_batch.forward_mode.is_decode_or_idle(): if forward_batch.forward_mode.is_decode_or_idle():
# DSv4 bakes this step's KV write target (c4/c128) into metadata, # 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 # so slice the shared multi-step out_cache_loc now, not at forward time.
# forward time.
out_cache_loc = forward_batch.out_cache_loc out_cache_loc = forward_batch.out_cache_loc
if self.topk > 0 and self.speculative_num_steps > 1: if self.topk > 0 and self.speculative_num_steps > 1:
out_cache_loc = per_step_draft_out_cache_loc( out_cache_loc = per_step_draft_out_cache_loc(
@@ -734,154 +861,24 @@ class DeepseekV4AttnBackend(
raise NotImplementedError(f"unsupported mode {forward_batch.forward_mode=}") raise NotImplementedError(f"unsupported mode {forward_batch.forward_mode=}")
self.forward_metadata = metadata 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: def init_cuda_graph_state(self, max_bs: int, max_num_tokens: int) -> None:
self.cuda_graph_metadata_of_bucket_and_bs: Dict[ self.cuda_graph_metadata_of_bucket_and_bs: Dict[
_GraphBucket, _GraphBucket,
Dict[ Dict[
int, int,
Union[DSV4Metadata, DSV4RawDecodeMetadata, DSV4RawVerifyMetadata], Union[
DSV4Metadata,
DSV4RawDecodeMetadata,
DSV4RawVerifyMetadata,
],
], ],
] = {bucket: {} for bucket in _GraphBucket} ] = {bucket: {} for bucket in _GraphBucket}
self.draft_extend_num_tokens_per_bs = ( self.draft_extend_num_tokens_per_bs = (
max_num_tokens // max_bs if max_bs > 0 else 1 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( def replay_cuda_graph_metadata_from(
self, self,
bs: int, bs: int,
@@ -938,24 +935,6 @@ class DeepseekV4AttnBackend(
cache_nope_fp8_rope_bf16_pack=swa_k_pack, 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( def forward(
self, self,
q: torch.Tensor, q: torch.Tensor,
@@ -968,8 +947,6 @@ class DeepseekV4AttnBackend(
attn_sink: Optional[torch.Tensor] = None, attn_sink: Optional[torch.Tensor] = None,
**_, **_,
) -> torch.Tensor: ) -> torch.Tensor:
self._maybe_upgrade_forward_metadata()
if self.mtp_enabled and forward_batch.forward_mode.is_idle(): 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) 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): def init_forward_metadata(self, forward_batch: ForwardBatch):
for i in range(self.speculative_num_steps - 1): for i in range(self.speculative_num_steps - 1):
self.attn_backends[i].init_forward_metadata(forward_batch) self.attn_backends[i].init_forward_metadata(forward_batch)
@@ -1241,49 +1264,10 @@ class DeepseekV4MultiStepBackend(DeepseekV4AttnBackend):
for i in range(self.speculative_num_steps): for i in range(self.speculative_num_steps):
self.attn_backends[i].init_cuda_graph_state(max_bs, max_num_tokens) 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): def on_after_cuda_graph_warmup(self):
for backend in self.attn_backends: for backend in self.attn_backends:
backend.on_after_cuda_graph_warmup() 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): def _pad_tensor_to_size(tensor: torch.Tensor, size: int, *, value: int = 0):
if value == 0: if value == 0:
@@ -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.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode 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.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 import ceil_align
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -381,7 +380,6 @@ class DeepseekV4HipRadixBackend(
DSV4RawVerifyMetadata, DSV4RawVerifyMetadata,
DSV4RawDecodeMetadata, DSV4RawDecodeMetadata,
] = None ] = None
self._replay_forward_batch: Optional[ForwardBatch] = None # FIXME: out-of-band
def _move_to_device(self, x: List[int]) -> torch.Tensor: def _move_to_device(self, x: List[int]) -> torch.Tensor:
pin_tensor = torch.tensor(x, dtype=torch.int32, pin_memory=True) 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, 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: def init_forward_metadata(self, forward_batch: ForwardBatch) -> None:
if self.mtp_enabled and forward_batch.forward_mode.is_idle(): if self.mtp_enabled and forward_batch.forward_mode.is_idle():
return return
@@ -676,8 +801,7 @@ class DeepseekV4HipRadixBackend(
if forward_batch.forward_mode.is_decode_or_idle(): if forward_batch.forward_mode.is_decode_or_idle():
# DSv4 bakes this step's KV write target (c4/c128) into metadata, # 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 # so slice the shared multi-step out_cache_loc now, not at forward time.
# forward time.
out_cache_loc = forward_batch.out_cache_loc out_cache_loc = forward_batch.out_cache_loc
if self.topk > 0 and self.speculative_num_steps > 1: if self.topk > 0 and self.speculative_num_steps > 1:
out_cache_loc = per_step_draft_out_cache_loc( out_cache_loc = per_step_draft_out_cache_loc(
@@ -725,154 +849,24 @@ class DeepseekV4HipRadixBackend(
raise NotImplementedError(f"unsupported mode {forward_batch.forward_mode=}") raise NotImplementedError(f"unsupported mode {forward_batch.forward_mode=}")
self.forward_metadata = metadata 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: def init_cuda_graph_state(self, max_bs: int, max_num_tokens: int) -> None:
self.cuda_graph_metadata_of_bucket_and_bs: Dict[ self.cuda_graph_metadata_of_bucket_and_bs: Dict[
_GraphBucket, _GraphBucket,
Dict[ Dict[
int, int,
Union[DSV4Metadata, DSV4RawDecodeMetadata, DSV4RawVerifyMetadata], Union[
DSV4Metadata,
DSV4RawDecodeMetadata,
DSV4RawVerifyMetadata,
],
], ],
] = {bucket: {} for bucket in _GraphBucket} ] = {bucket: {} for bucket in _GraphBucket}
self.draft_extend_num_tokens_per_bs = ( self.draft_extend_num_tokens_per_bs = (
max_num_tokens // max_bs if max_bs > 0 else 1 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( def replay_cuda_graph_metadata_from(
self, self,
bs: int, bs: int,
@@ -929,24 +923,6 @@ class DeepseekV4HipRadixBackend(
cache_nope_fp8_rope_bf16_pack=swa_k_pack, 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( def forward(
self, self,
q: torch.Tensor, q: torch.Tensor,
@@ -959,8 +935,6 @@ class DeepseekV4HipRadixBackend(
attn_sink: Optional[torch.Tensor] = None, attn_sink: Optional[torch.Tensor] = None,
**_, **_,
) -> torch.Tensor: ) -> torch.Tensor:
self._maybe_upgrade_forward_metadata()
if self.mtp_enabled and forward_batch.forward_mode.is_idle(): 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) 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): def init_forward_metadata(self, forward_batch: ForwardBatch):
for i in range(self.speculative_num_steps - 1): for i in range(self.speculative_num_steps - 1):
self.attn_backends[i].init_forward_metadata(forward_batch) self.attn_backends[i].init_forward_metadata(forward_batch)
@@ -1220,49 +1240,10 @@ class DeepseekV4MultiStepBackend(DeepseekV4HipRadixBackend):
for i in range(self.speculative_num_steps): for i in range(self.speculative_num_steps):
self.attn_backends[i].init_cuda_graph_state(max_bs, max_num_tokens) 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): def on_after_cuda_graph_warmup(self):
for backend in self.attn_backends: for backend in self.attn_backends:
backend.on_after_cuda_graph_warmup() 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): def _pad_tensor_to_size(tensor: torch.Tensor, size: int, *, value: int = 0):
if value == 0: if value == 0:
@@ -408,6 +408,25 @@ class DeepseekSparseAttnBackend(
) )
return page_table[:, strided_indices] // page_size 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): def init_forward_metadata(self, forward_batch: ForwardBatch):
"""Init the metadata for a forward pass.""" """Init the metadata for a forward pass."""
batch_size = forward_batch.batch_size batch_size = forward_batch.batch_size
@@ -971,42 +990,23 @@ class DeepseekSparseAttnBackend(
self.decode_cuda_graph_metadata[bs] = metadata self.decode_cuda_graph_metadata[bs] = metadata
self.forward_metadata = metadata 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[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(
self, self,
bs: int, bs: int,
req_pool_indices: torch.Tensor, req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor, seq_lens: torch.Tensor,
seq_lens_sum: int, seq_lens_cpu: torch.Tensor,
encoder_lens: Optional[torch.Tensor],
forward_mode: ForwardMode, forward_mode: ForwardMode,
spec_info: Optional[SpecInput], spec_info: Optional[SpecInput],
seq_lens_cpu: Optional[torch.Tensor],
out_cache_loc: Optional[torch.Tensor] = None, out_cache_loc: Optional[torch.Tensor] = None,
actual_forward_mode: Optional[ForwardMode] = 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 assert seq_lens_cpu is not None
if bs not in self.decode_cuda_graph_metadata: if bs not in self.decode_cuda_graph_metadata:
@@ -2370,21 +2370,26 @@ class DeepseekSparseAttnMultiStepBackend:
for i in range(self.speculative_num_steps - 1): for i in range(self.speculative_num_steps - 1):
self.attn_backends[i].init_cuda_graph_state(max_bs, max_num_tokens) 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 init_forward_metadata_out_graph(
for i in range(self.speculative_num_steps - 1): self,
self.attn_backends[i].init_forward_metadata_capture_cuda_graph( forward_batch: ForwardBatch,
forward_batch.batch_size, in_capture: bool = False,
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
): ):
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(): if envs.SGLANG_DSA_ENABLE_MTP_PRECOMPUTE_METADATA.get():
# Precompute metadata once (shared across all backends) # Precompute metadata once (shared across all backends)
precomputed = self.attn_backends[0]._precompute_replay_metadata( precomputed = self.attn_backends[0]._precompute_replay_metadata(
@@ -2542,20 +2547,21 @@ class DeepseekSparseAttnMultiStepBackend:
forward_mode=ForwardMode.DECODE, forward_mode=ForwardMode.DECODE,
) )
else: else:
# Fallback: compute metadata separately for each backend
for i in range(self.speculative_num_steps - 1): 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, bs=bs,
req_pool_indices=forward_batch.req_pool_indices, req_pool_indices=forward_batch.req_pool_indices,
seq_lens=forward_batch.seq_lens, 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,
forward_mode=ForwardMode.DECODE, forward_mode=ForwardMode.DECODE,
spec_info=forward_batch.spec_info, spec_info=forward_batch.spec_info,
seq_lens_cpu=forward_batch.seq_lens_cpu,
out_cache_loc=None, 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) # Backward-compat aliases (deprecated: use DSA class names)
DeepseekSparseAttnBackend = DeepseekSparseAttnBackend DeepseekSparseAttnBackend = DeepseekSparseAttnBackend
@@ -57,9 +57,6 @@ class CompressorBackendMixin:
assert isinstance(metadata, FusedCompressMetadata) assert isinstance(metadata, FusedCompressMetadata)
return metadata return metadata
def _maybe_upgrade_forward_metadata(self) -> None:
pass
def forward_compress( def forward_compress(
self, self,
*, *,
@@ -153,11 +150,6 @@ class CompressorBackendMixin:
) -> None: ) -> None:
if forward_batch.forward_mode.is_idle(): if forward_batch.forward_mode.is_idle():
return 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 token_to_kv_pool = self.token_to_kv_pool
if TYPE_CHECKING: if TYPE_CHECKING:
assert isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool) assert isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool)
@@ -187,8 +179,6 @@ class CompressorBackendMixin:
compressor: Compressor, compressor: Compressor,
) -> None: ) -> None:
assert is_overlap_compress(compressor.ratio) 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 token_to_kv_pool = self.token_to_kv_pool
if TYPE_CHECKING: if TYPE_CHECKING:
assert isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool) assert isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool)
@@ -404,9 +404,6 @@ class CompressorBackendMixin:
super().__init__() super().__init__()
self.forward_metadata: DSV4Metadata self.forward_metadata: DSV4Metadata
# NOTE: Will be overridden
def _maybe_upgrade_forward_metadata(self): ...
def _get_paged_compress_metadata(self, compress_ratio: int) -> CompressMetadata: def _get_paged_compress_metadata(self, compress_ratio: int) -> CompressMetadata:
attr_name = f"c{compress_ratio}_compress_metadata" attr_name = f"c{compress_ratio}_compress_metadata"
return getattr(self.forward_metadata, attr_name) return getattr(self.forward_metadata, attr_name)
@@ -483,7 +480,6 @@ class CompressorBackendMixin:
if forward_batch.forward_mode.is_idle(): if forward_batch.forward_mode.is_idle():
return return
self._maybe_upgrade_forward_metadata()
token_to_kv_pool = self.token_to_kv_pool token_to_kv_pool = self.token_to_kv_pool
token_to_kv_pool = cast("DeepSeekV4TokenToKVPool", token_to_kv_pool) token_to_kv_pool = cast("DeepSeekV4TokenToKVPool", token_to_kv_pool)
kv_score_input = compressor.compute_kv_score(x, forward_batch) kv_score_input = compressor.compute_kv_score(x, forward_batch)
@@ -443,9 +443,6 @@ class C4IndexerBackendMixin:
) -> None: ) -> None:
if forward_batch.forward_mode.is_idle(): if forward_batch.forward_mode.is_idle():
return 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 token_to_kv_pool = self.token_to_kv_pool
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -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.distributed.parallel_state import get_tensor_model_parallel_rank
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend 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 from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -172,6 +174,35 @@ class DualChunkFlashAttentionBackend(AttentionBackend):
end_head = start_head + self.num_heads end_head = start_head + self.num_heads
return [layer_sparse_attention_config[i] for i in range(start_head, end_head)] 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): def init_forward_metadata(self, forward_batch: ForwardBatch):
"""Initialize forward metadata hence all layers in the forward pass can reuse it.""" """Initialize forward metadata hence all layers in the forward pass can reuse it."""
@@ -577,49 +608,17 @@ class DualChunkFlashAttentionBackend(AttentionBackend):
self.forward_metadata = metadata self.forward_metadata = metadata
def init_forward_metadata_capture_cuda_graph( def _apply_cuda_graph_metadata(
self, self,
bs: int, bs: int,
num_tokens: int,
req_pool_indices: torch.Tensor, req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor, seq_lens: torch.Tensor,
encoder_lens: Optional[torch.Tensor],
forward_mode: ForwardMode, forward_mode: ForwardMode,
spec_info: Optional[None],
): ):
self._bind_metadata_buffers(bs, req_pool_indices, forward_mode) """Shared capture+replay body for the cuda-graph init path.
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
def init_forward_metadata_replay_cuda_graph( Public entry: :py:meth:`init_forward_metadata_out_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."""
assert forward_mode.is_decode() assert forward_mode.is_decode()
seq_lens = seq_lens[:bs] seq_lens = seq_lens[:bs]
req_pool_indices = req_pool_indices[:bs] req_pool_indices = req_pool_indices[:bs]
@@ -273,6 +273,102 @@ class FlashAttentionBackend(AttentionBackend):
num_splits=self.num_splits, 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): def init_forward_metadata(self, forward_batch: ForwardBatch):
"""Initialize forward metadata hence all layers in the forward pass can reuse it.""" """Initialize forward metadata hence all layers in the forward pass can reuse it."""
metadata = FlashAttentionMetadata() metadata = FlashAttentionMetadata()
@@ -1901,77 +1997,7 @@ class FlashAttentionBackend(AttentionBackend):
return metadata, metadata_expand return metadata, metadata_expand
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[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(
self, self,
bs: int, bs: int,
req_pool_indices: torch.Tensor, req_pool_indices: torch.Tensor,
@@ -1983,7 +2009,13 @@ class FlashAttentionBackend(AttentionBackend):
seq_lens_cpu: Optional[torch.Tensor], seq_lens_cpu: Optional[torch.Tensor],
out_cache_loc: Optional[torch.Tensor] = None, 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 = seq_lens[:bs]
seq_lens_cpu = seq_lens_cpu[:bs] seq_lens_cpu = seq_lens_cpu[:bs]
req_pool_indices = req_pool_indices[: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) 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: if encoder_lens is not None:
# Per-request varlen encoder support (e.g. MossVL different images). # Per-request varlen encoder support (e.g. MossVL different images).
metadata.encoder_max_seq_len_k = int(encoder_lens.max().item()) metadata.encoder_max_seq_len_k = int(encoder_lens.max().item())
@@ -2353,7 +2395,10 @@ class FlashAttentionBackend(AttentionBackend):
return 1 return 1
def _maybe_init_local_attn_metadata( 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.""" """Centralized utility to initialize local_attn_metadata if chunked attention is enabled."""
if not self.has_local_attention: if not self.has_local_attention:
@@ -2628,45 +2673,33 @@ class FlashAttentionMultiStepBackend:
for i in range(self.speculative_num_steps - 1): for i in range(self.speculative_num_steps - 1):
self.attn_backends[i].init_cuda_graph_state(max_bs, max_num_tokens) 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, self,
forward_batch: ForwardBatch, 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 not None
assert forward_batch.spec_info.is_draft_input() assert forward_batch.spec_info.is_draft_input()
for i in range(self.speculative_num_steps - 1): inner_fb = build_inner_fb_view(
self.attn_backends[i].init_forward_metadata_capture_cuda_graph( forward_batch,
forward_batch.batch_size, bs=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, forward_mode=ForwardMode.DECODE,
spec_info=forward_batch.spec_info, encoder_lens=forward_batch.encoder_lens,
) )
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): for i in range(self.speculative_num_steps - 1):
# TODO: incrementally update the metadata for the later steps, # TODO: incrementally update the metadata for the later steps,
# so that they do not need to recompute everything from scratch. # so that they do not need to recompute everything from scratch.
self.attn_backends[i].init_forward_metadata_replay_cuda_graph( self.attn_backends[i].init_forward_metadata_out_graph(
bs, inner_fb, in_capture=in_capture
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,
) )
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()) @torch.compile(dynamic=True, backend=get_compiler_backend())
def draft_decode_set_expand_metadata( def draft_decode_set_expand_metadata(
@@ -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): def init_forward_metadata(self, forward_batch: ForwardBatch):
if forward_batch.forward_mode.is_decode_or_idle(): if forward_batch.forward_mode.is_decode_or_idle():
self.indices_updater_decode.update( self.indices_updater_decode.update(
@@ -662,85 +725,6 @@ class FlashInferAttnBackend(AttentionBackend):
else: else:
raise ValueError(f"Invalid mode: {forward_mode=}") 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): def get_cuda_graph_seq_len_fill_value(self):
return 1 return 1
@@ -1668,37 +1652,27 @@ class FlashInferMultiStepDraftBackend:
max_bs, max_num_tokens, kv_indices_buf=self.cuda_graph_kv_indices[i] 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 init_forward_metadata_out_graph(
def call_fn(i, forward_batch): self,
self.attn_backends[i].init_forward_metadata_capture_cuda_graph( forward_batch: ForwardBatch,
forward_batch.batch_size, in_capture: bool = False,
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 call_fn(i, forward_batch): from sglang.srt.model_executor.forward_batch_info import build_inner_fb_view
self.attn_backends[i].init_forward_metadata_replay_cuda_graph(
bs, bs = forward_batch.batch_size
forward_batch.req_pool_indices,
forward_batch.seq_lens, def call_fn(i, fb):
seq_lens_sum=-1, inner_fb = build_inner_fb_view(fb, bs=bs, forward_mode=ForwardMode.DECODE)
encoder_lens=None, self.attn_backends[i].init_forward_metadata_out_graph(
forward_mode=ForwardMode.DECODE, inner_fb, in_capture=in_capture
spec_info=forward_batch.spec_info,
seq_lens_cpu=forward_batch.seq_lens_cpu,
) )
self.common_template(forward_batch, self.cuda_graph_kv_indices, call_fn) 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( def should_use_tensor_core(
kv_cache_dtype: torch.dtype, kv_cache_dtype: torch.dtype,
@@ -289,6 +289,80 @@ class FlashInferMLAAttnBackend(AttentionBackend):
self.decode_cuda_graph_metadata = {} self.decode_cuda_graph_metadata = {}
self.prefill_cuda_graph_metadata = {} # For verify 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): def init_forward_metadata(self, forward_batch: ForwardBatch):
if forward_batch.forward_mode.is_decode_or_idle(): if forward_batch.forward_mode.is_decode_or_idle():
self.indices_updater_decode.update( self.indices_updater_decode.update(
@@ -374,83 +448,20 @@ class FlashInferMLAAttnBackend(AttentionBackend):
"kv_indices": self.cuda_graph_kv_indices, "kv_indices": self.cuda_graph_kv_indices,
} }
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[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(
self, self,
bs: int, bs: int,
req_pool_indices: torch.Tensor, req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor, seq_lens: torch.Tensor,
seq_lens_sum: int, seq_lens_sum: int,
encoder_lens: Optional[torch.Tensor],
forward_mode: ForwardMode, forward_mode: ForwardMode,
spec_info: Optional[SpecInput], spec_info: Optional[SpecInput],
seq_lens_cpu: Optional[torch.Tensor], 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(): if forward_mode.is_decode_or_idle():
assert seq_lens_cpu is not None assert seq_lens_cpu is not None
kv_len_arr_cpu = seq_lens_cpu[:bs] 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] 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 init_forward_metadata_out_graph(
def call_fn(i, forward_batch): self,
self.attn_backends[i].init_forward_metadata_capture_cuda_graph( forward_batch: ForwardBatch,
forward_batch.batch_size, in_capture: bool = False,
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 call_fn(i, forward_batch): from sglang.srt.model_executor.forward_batch_info import build_inner_fb_view
self.attn_backends[i].init_forward_metadata_replay_cuda_graph(
bs, inner_fb = build_inner_fb_view(
forward_batch.req_pool_indices, forward_batch,
forward_batch.seq_lens, bs=forward_batch.batch_size,
seq_lens_sum=-1,
encoder_lens=None,
forward_mode=ForwardMode.DECODE, forward_mode=ForwardMode.DECODE,
spec_info=forward_batch.spec_info, )
seq_lens_cpu=forward_batch.seq_lens_cpu,
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) 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( def fast_mla_decode_plan(
self, self,
@@ -21,7 +21,6 @@ from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMo
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.layers.radix_attention import RadixAttention from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.model_executor.model_runner import ModelRunner from sglang.srt.model_executor.model_runner import ModelRunner
from sglang.srt.speculative.spec_info import SpecInput
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -86,6 +85,25 @@ class FlashMLABackend(FlashInferMLAAttnBackend):
self.cuda_graph_mla_metadata_view = None self.cuda_graph_mla_metadata_view = None
self.cuda_graph_num_splits_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): def init_forward_metadata(self, forward_batch: ForwardBatch):
bs = forward_batch.batch_size bs = forward_batch.batch_size
if forward_batch.forward_mode.is_decode_or_idle(): 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_mla_metadata_view = None
self.cuda_graph_num_splits_view = None self.cuda_graph_num_splits_view = None
def init_forward_metadata_capture_cuda_graph( def _apply_decode_target_verify_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[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(
self, self,
bs: int, bs: int,
req_pool_indices: torch.Tensor, req_pool_indices: torch.Tensor,
seq_lens: 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], 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 = seq_lens[:bs]
seq_lens_cpu = seq_lens_cpu[:bs] if seq_lens_cpu is not None else None 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_num_splits_view,
self.cuda_graph_kv_indices[:bs, :max_seqlen_pad], 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): def get_cuda_graph_seq_len_fill_value(self):
return 1 return 1
@@ -516,39 +494,29 @@ class FlashMLAMultiStepDraftBackend:
max_bs, max_num_tokens, block_kv_indices=None max_bs, max_num_tokens, block_kv_indices=None
) )
def init_forward_metadata_capture_cuda_graph(self, forward_batch: ForwardBatch): def init_forward_metadata_out_graph(
def call_fn(i, forward_batch): self,
# EAGLE draft worker uses DECODE mode for draft steps forward_batch: ForwardBatch,
from sglang.srt.model_executor.forward_batch_info import ForwardMode in_capture: bool = False,
# 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 call_fn(i, forward_batch): from sglang.srt.model_executor.forward_batch_info import (
from sglang.srt.model_executor.forward_batch_info import ForwardMode ForwardMode,
build_inner_fb_view,
)
self.attn_backends[i].init_forward_metadata_replay_cuda_graph( inner_fb = build_inner_fb_view(
bs, forward_batch,
forward_batch.req_pool_indices, bs=forward_batch.batch_size,
forward_batch.seq_lens,
seq_lens_sum=-1,
encoder_lens=None,
forward_mode=ForwardMode.DECODE, forward_mode=ForwardMode.DECODE,
spec_info=forward_batch.spec_info, )
seq_lens_cpu=forward_batch.seq_lens_cpu,
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) 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)
@@ -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.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, ForwardMode
from sglang.srt.model_executor.model_runner import ModelRunner from sglang.srt.model_executor.model_runner import ModelRunner
from sglang.srt.speculative.spec_info import SpecInput
class HybridAttnBackend(AttentionBackend): class HybridAttnBackend(AttentionBackend):
@@ -52,6 +51,14 @@ class HybridAttnBackend(AttentionBackend):
else: else:
return self.prefill_backend 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): def init_forward_metadata(self, forward_batch: ForwardBatch):
backend = self._select_backend(forward_batch.forward_mode) backend = self._select_backend(forward_batch.forward_mode)
backend.init_forward_metadata(forward_batch) backend.init_forward_metadata(forward_batch)
@@ -66,50 +73,6 @@ class HybridAttnBackend(AttentionBackend):
# that will be used for target_verify. # that will be used for target_verify.
self.prefill_backend.init_cuda_graph_state(max_bs, max_num_tokens) 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): def get_cuda_graph_seq_len_fill_value(self):
return self.decode_backend.get_cuda_graph_seq_len_fill_value() return self.decode_backend.get_cuda_graph_seq_len_fill_value()
@@ -258,6 +258,21 @@ class MambaAttnBackendBase(AttentionBackend):
has_mamba_track_mask=has_mamba_track_mask, 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): def init_forward_metadata(self, forward_batch: ForwardBatch):
self._execute_deferred_mamba_cow_and_clear(forward_batch) self._execute_deferred_mamba_cow_and_clear(forward_batch)
self.forward_metadata = self._forward_metadata(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), 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( def init_forward_metadata_capture_cpu_graph(
self, self,
bs: int, bs: int,
@@ -698,6 +677,27 @@ class Mamba2AttnBackend(MambaAttnBackendBase):
model_runner.server_args.mamba_track_interval >= self.mamba_chunk_size 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})" ), 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): def init_forward_metadata(self, forward_batch: ForwardBatch):
self._execute_deferred_mamba_cow_and_clear(forward_batch) self._execute_deferred_mamba_cow_and_clear(forward_batch)
metadata = self._forward_metadata(forward_batch) metadata = self._forward_metadata(forward_batch)
@@ -707,49 +707,6 @@ class Mamba2AttnBackend(MambaAttnBackendBase):
forward_batch, 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( def forward(
self, self,
mixer: MambaMixer2, mixer: MambaMixer2,
@@ -847,6 +804,16 @@ class HybridLinearAttnBackend(AttentionBackend):
assert layer_id is not None, "either layer or layer_id must be provided" assert layer_id is not None, "either layer or layer_id must be provided"
return layer_id in self.full_attn_layers 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): def init_forward_metadata(self, forward_batch: ForwardBatch):
if forward_batch.forward_mode.is_draft_extend_v2(): if forward_batch.forward_mode.is_draft_extend_v2():
# DRAFT_EXTEND_V2 only runs full-attn layers in the draft model, # 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: for attn_backend in self.attn_backend_list:
attn_backend.init_cpu_graph_state(max_bs, max_num_tokens) 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( def init_forward_metadata_capture_cpu_graph(
self, self,
bs: int, bs: int,
@@ -906,29 +852,6 @@ class HybridLinearAttnBackend(AttentionBackend):
spec_info, 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): def get_cuda_graph_seq_len_fill_value(self):
return self.full_attn_backend.get_cuda_graph_seq_len_fill_value() return self.full_attn_backend.get_cuda_graph_seq_len_fill_value()
@@ -1,6 +1,5 @@
import logging import logging
import math import math
from typing import Optional, Union
import torch import torch
@@ -9,12 +8,13 @@ from sglang.srt.layers.attention.linear.lightning_attn import (
BailingLinearKernel, BailingLinearKernel,
linear_decode_forward_triton, 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.attention.linear.seg_la import SegLaMeta, seg_la_fwd
from sglang.srt.layers.radix_attention import RadixAttention 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.model_executor.model_runner import ModelRunner
from sglang.srt.speculative.eagle_info import EagleDraftInput, EagleVerifyInput
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -71,6 +71,28 @@ class LightningAttentionBackend(MambaAttnBackendBase):
f"linear_backend for linear attention in hybrid_linear_backend: {self.linear_backend}" 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): def init_forward_metadata(self, forward_batch: ForwardBatch):
metadata = self._forward_metadata(forward_batch) metadata = self._forward_metadata(forward_batch)
self.forward_metadata = BailingLinearMetadata.prepare_mixed( self.forward_metadata = BailingLinearMetadata.prepare_mixed(
@@ -79,45 +101,6 @@ class LightningAttentionBackend(MambaAttnBackendBase):
forward_batch, 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 @staticmethod
def _build_slope_tensor( def _build_slope_tensor(
n_attention_heads: int, num_hidden_layers: int, device="cuda" n_attention_heads: int, num_hidden_layers: int, device="cuda"
+140 -206
View File
@@ -1,13 +1,11 @@
from typing import TYPE_CHECKING, Callable, List, Optional from types import SimpleNamespace
from typing import TYPE_CHECKING, Callable, List
import torch
from sglang.srt.batch_overlap import two_batch_overlap from sglang.srt.batch_overlap import two_batch_overlap
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
from sglang.srt.speculative.spec_info import SpecInput
if TYPE_CHECKING: 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): class TboAttnBackend(AttentionBackend):
@@ -27,6 +25,88 @@ class TboAttnBackend(AttentionBackend):
children=[creator() for _ in range(2)], 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"): def init_forward_metadata(self, forward_batch: "ForwardBatch"):
self.primary.init_forward_metadata(forward_batch=forward_batch) self.primary.init_forward_metadata(forward_batch=forward_batch)
if forward_batch.tbo_children is not None: 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 # 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) item.init_cuda_graph_state(max_bs=max_bs, max_num_tokens=max_num_tokens)
def init_forward_metadata_capture_cuda_graph( def on_after_cuda_graph_warmup(self):
self, self.primary.on_after_cuda_graph_warmup()
bs: int, for child in self.children:
num_tokens: int, child.on_after_cuda_graph_warmup()
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 get_cuda_graph_seq_len_fill_value(self): def get_cuda_graph_seq_len_fill_value(self):
ans = self.primary.get_cuda_graph_seq_len_fill_value() 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) return self.primary.get_indexer_metadata(layer_id, forward_batch)
def _init_forward_metadata_cuda_graph_split( def _build_tbo_child_replay_fb_view(
fn_name: str, fb_view,
*,
child_bs: int,
seq_slice: slice, seq_slice: slice,
output_bs: int, tok_slice: slice,
# common args token_num_per_seq: int,
bs: int, ) -> SimpleNamespace:
req_pool_indices: torch.Tensor, """Slice a parent replay fb_view into a per-child view.
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"
if spec_info is not None:
output_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
),
)
else: Mirrors the legacy ``_init_forward_metadata_cuda_graph_split`` (deleted
output_spec_info = None along with the cuda_graph variants) for the new
ans = dict( ``init_forward_metadata_out_graph(fb_view)`` contract: padded
bs=output_bs, capture-time buffers are sliced per child, spec_info is split, and
req_pool_indices=req_pool_indices[seq_slice], seq_lens_sum is recomputed from the sliced ``seq_lens_cpu``.
seq_lens=seq_lens[seq_slice], """
# directly forward
forward_mode=forward_mode,
# ignore
encoder_lens=None,
spec_info=output_spec_info,
)
if fn_name == "init_forward_metadata_capture_cuda_graph":
assert ( assert (
capture_num_tokens == bs * token_num_per_seq getattr(fb_view, "encoder_lens", None) is None
), "Only support num_tokens==bs * token_num_per_seq for target-verify or decode mode" ), "TBO replay split does not support encoder_lens yet"
ans.update( spec_info = getattr(fb_view, "spec_info", None)
dict( if spec_info is not None:
num_tokens=output_bs * token_num_per_seq, 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(
elif fn_name == "init_forward_metadata_replay_cuda_graph": spec_info=spec_info,
output_seq_lens_cpu = replay_seq_lens_cpu[seq_slice] start_seq_index=start_seq,
ans.update( end_seq_index=end_seq,
dict( start_token_index=start_seq * token_num_per_seq,
seq_lens_sum=output_seq_lens_cpu.sum().item(), end_token_index=end_seq * token_num_per_seq,
seq_lens_cpu=output_seq_lens_cpu,
)
) )
else: else:
raise NotImplementedError child_spec_info = None
child_seq_lens_cpu = fb_view.seq_lens_cpu[seq_slice]
return ans 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,
out_cache_loc=(
parent_out_cache_loc[tok_slice]
if parent_out_cache_loc is not None
else None
),
spec_info=child_spec_info,
)
@@ -458,6 +458,59 @@ class TritonAttnBackend(AttentionBackend):
) )
return qo_indptr, kv_indptr, num_tokens_per_bs 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): def init_forward_metadata(self, forward_batch: ForwardBatch):
"""Init auxiliary variables for triton attention backend.""" """Init auxiliary variables for triton attention backend."""
@@ -835,66 +888,18 @@ class TritonAttnBackend(AttentionBackend):
else: else:
raise ValueError(f"Invalid forward mode: {forward_mode=} for CUDA Graph.") raise ValueError(f"Invalid forward mode: {forward_mode=} for CUDA Graph.")
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[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(
self, self,
bs: int, bs: int,
req_pool_indices: torch.Tensor, req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor, seq_lens: torch.Tensor,
seq_lens_sum: int,
encoder_lens: Optional[torch.Tensor],
forward_mode: ForwardMode, forward_mode: ForwardMode,
spec_info: Optional[SpecInput], 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 # NOTE: encoder_lens expected to be zeros or None
if forward_mode.is_decode_or_idle(): if forward_mode.is_decode_or_idle():
assert spec_info is None, "Multi-step cuda graph init is not done here." assert spec_info is None, "Multi-step cuda graph init is not done here."
@@ -1417,23 +1422,28 @@ class TritonMultiStepDraftBackend:
cuda_graph_num_kv_splits_buf=self.cuda_graph_num_kv_splits, cuda_graph_num_kv_splits_buf=self.cuda_graph_num_kv_splits,
) )
def init_forward_metadata_capture_cuda_graph(self, forward_batch: ForwardBatch): def init_forward_metadata_out_graph(
def call_fn(i, forward_batch): self,
self.attn_backends[i].init_forward_metadata_capture_cuda_graph( forward_batch: ForwardBatch,
forward_batch.batch_size, in_capture: bool = False,
forward_batch.batch_size * self.topk, ):
forward_batch.req_pool_indices, from sglang.srt.model_executor.forward_batch_info import build_inner_fb_view
forward_batch.seq_lens,
encoder_lens=None, if in_capture:
inner_fb = build_inner_fb_view(
forward_batch,
bs=forward_batch.batch_size,
forward_mode=ForwardMode.DECODE, forward_mode=ForwardMode.DECODE,
spec_info=forward_batch.spec_info, )
def call_fn(i, _forward_batch):
self.attn_backends[i].init_forward_metadata_out_graph(
inner_fb, in_capture=True
) )
self.common_template(forward_batch, None, call_fn) self.common_template(forward_batch, None, call_fn)
else:
def init_forward_metadata_replay_cuda_graph( bs = forward_batch.batch_size
self, forward_batch: ForwardBatch, bs: int
):
self.common_template(forward_batch, None, None) self.common_template(forward_batch, None, None)
# NOTE: Multi-step's attention backends use the slice of # NOTE: Multi-step's attention backends use the slice of
@@ -1442,12 +1452,16 @@ class TritonMultiStepDraftBackend:
# So we don't need to assign the KV indices inside the attention backend. # So we don't need to assign the KV indices inside the attention backend.
# Compute num_kv_splits only once # Compute num_kv_splits only once
num_token = forward_batch.batch_size * self.topk num_token = bs * self.topk
self.attn_backends[-1].get_num_kv_splits( self.attn_backends[-1].get_num_kv_splits(
self.attn_backends[-1].cuda_graph_num_kv_splits[:num_token], self.attn_backends[-1].cuda_graph_num_kv_splits[:num_token],
forward_batch.seq_lens[:bs], 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( def update_sliding_window_buffer(
window_kv_indptr, window_kv_indptr,
@@ -397,50 +397,19 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
return 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],
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(
self, self,
bs: int, bs: int,
req_pool_indices: torch.Tensor, req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor, seq_lens: torch.Tensor,
seq_lens_sum: int,
encoder_lens: Optional[torch.Tensor],
forward_mode: ForwardMode, forward_mode: ForwardMode,
spec_info: Optional[SpecInput], spec_info: Optional[SpecInput],
seq_lens_cpu: Optional[torch.Tensor], 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 = seq_lens[:bs]
seq_lens_cpu = seq_lens_cpu[:bs] seq_lens_cpu = seq_lens_cpu[:bs]
req_pool_indices = req_pool_indices[:bs] req_pool_indices = req_pool_indices[:bs]
@@ -571,6 +540,48 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
page_size=self.page_size, 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): def init_forward_metadata(self, forward_batch: ForwardBatch):
"""Initialize the metadata for a forward pass.""" """Initialize the metadata for a forward pass."""
@@ -899,39 +910,25 @@ class TRTLLMHAAttnMultiStepDraftBackend(FlashInferMultiStepDraftBackend):
for i in range(self.speculative_num_steps - 1): for i in range(self.speculative_num_steps - 1):
self.attn_backends[i].init_cuda_graph_state(max_bs, max_num_tokens) 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, self,
forward_batch: ForwardBatch, 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 not None
assert forward_batch.spec_info.is_draft_input() assert forward_batch.spec_info.is_draft_input()
for i in range(self.speculative_num_steps - 1): # TRTLLM-MHA uses encoder_lens from the original fb for inner dispatch
self.attn_backends[i].init_forward_metadata_capture_cuda_graph( # (FlashInfer parent forces encoder_lens=None instead).
forward_batch.batch_size, inner_fb = build_inner_fb_view(
forward_batch.batch_size * self.topk, forward_batch,
forward_batch.req_pool_indices, bs=forward_batch.batch_size,
forward_batch.seq_lens,
encoder_lens=forward_batch.encoder_lens,
forward_mode=ForwardMode.DECODE, 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, encoder_lens=forward_batch.encoder_lens,
forward_mode=ForwardMode.DECODE, )
spec_info=forward_batch.spec_info, for i in range(self.speculative_num_steps - 1):
seq_lens_cpu=forward_batch.seq_lens_cpu, self.attn_backends[i].init_forward_metadata_out_graph(
inner_fb, in_capture=in_capture
) )
@@ -38,7 +38,6 @@ if is_flashinfer_available():
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.layers.radix_attention import RadixAttention from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.model_executor.model_runner import ModelRunner from sglang.srt.model_executor.model_runner import ModelRunner
from sglang.srt.speculative.spec_info import SpecInput
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -499,83 +498,23 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
self.decode_cuda_graph_metadata[bs] = metadata self.decode_cuda_graph_metadata[bs] = metadata
self.forward_decode_metadata = metadata self.forward_decode_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[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(
self, self,
bs: int, bs: int,
req_pool_indices: torch.Tensor, req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor, seq_lens: torch.Tensor,
seq_lens_sum: int,
encoder_lens: Optional[torch.Tensor],
forward_mode: ForwardMode, forward_mode: ForwardMode,
spec_info: Optional[SpecInput],
seq_lens_cpu: Optional[torch.Tensor],
): ):
"""Replay CUDA graph with new inputs.""" """Shared decode / target-verify / draft-extend capture+replay body.
# 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,
)
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] metadata = self.decode_cuda_graph_metadata[bs]
if forward_mode.is_target_verify(): if forward_mode.is_target_verify():
seq_lens = seq_lens[:bs] + self.num_draft_tokens seq_lens = seq_lens[:bs] + self.num_draft_tokens
metadata.seq_lens_k.copy_(seq_lens.to(dtype=torch.int32)) 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): elif forward_mode.is_draft_extend(include_v2=True):
num_tokens_per_bs = self.num_draft_tokens num_tokens_per_bs = self.num_draft_tokens
metadata.max_seq_len_q = num_tokens_per_bs metadata.max_seq_len_q = num_tokens_per_bs
@@ -618,6 +557,47 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
if fallback_to_flashinfer_impl: if fallback_to_flashinfer_impl:
super().init_mha_chunk_metadata(forward_batch) 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): def init_forward_metadata(self, forward_batch: ForwardBatch):
"""Initialize the metadata for a forward pass.""" """Initialize the metadata for a forward pass."""
# Delegate to parent for non-decode modes. # Delegate to parent for non-decode modes.
@@ -1284,17 +1264,21 @@ class TRTLLMMLAMultiStepDraftBackend(FlashInferMLAMultiStepDraftBackend):
for i in range(self.speculative_num_steps - 1): for i in range(self.speculative_num_steps - 1):
self.attn_backends[i].init_forward_metadata(forward_batch) self.attn_backends[i].init_forward_metadata(forward_batch)
def init_forward_metadata_replay_cuda_graph( def init_forward_metadata_out_graph(
self, forward_batch: ForwardBatch, bs: int self,
forward_batch: ForwardBatch,
in_capture: bool = False,
): ):
for i in range(self.speculative_num_steps - 1): from sglang.srt.model_executor.forward_batch_info import build_inner_fb_view
self.attn_backends[i].init_forward_metadata_replay_cuda_graph(
bs, if in_capture:
forward_batch.req_pool_indices, return super().init_forward_metadata_out_graph(
forward_batch.seq_lens, forward_batch, in_capture=in_capture
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,
) )
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)
@@ -147,6 +147,53 @@ class WaveAttnBackend(AttentionBackend):
MAX_NUM_SEQ=SCHEDULE_SEQ, 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): def init_forward_metadata(self, forward_batch: ForwardBatch):
"""Init auxiliary variables for wave attention backend.""" """Init auxiliary variables for wave attention backend."""
@@ -373,59 +420,18 @@ class WaveAttnBackend(AttentionBackend):
else: else:
raise ValueError(f"Invalid forward mode: {forward_mode=} for CUDA Graph.") raise ValueError(f"Invalid forward mode: {forward_mode=} for CUDA Graph.")
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[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(
self, self,
bs: int, bs: int,
req_pool_indices: torch.Tensor, req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor, seq_lens: torch.Tensor,
seq_lens_sum: int,
encoder_lens: Optional[torch.Tensor],
forward_mode: ForwardMode, forward_mode: ForwardMode,
spec_info: Optional[SpecInput], 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(): if forward_mode.is_decode_or_idle():
kv_indptr = self.kv_indptr kv_indptr = self.kv_indptr
kv_indices = self.cuda_graph_kv_indices kv_indices = self.cuda_graph_kv_indices
@@ -906,7 +906,10 @@ class XPUAttentionBackend(AttentionBackend):
return 1 return 1
def _init_local_attn_metadata( 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.""" """Centralized utility to initialize local_attn_metadata if chunked attention is enabled."""
if self.attention_chunk_size is None: if self.attention_chunk_size is None:
@@ -24,6 +24,7 @@ import os
from contextlib import contextmanager from contextlib import contextmanager
from dataclasses import dataclass from dataclasses import dataclass
from functools import partial from functools import partial
from types import SimpleNamespace
from typing import TYPE_CHECKING, Callable, Dict, List, Optional, Tuple, Union from typing import TYPE_CHECKING, Callable, Dict, List, Optional, Tuple, Union
import torch import torch
@@ -107,6 +108,57 @@ if TYPE_CHECKING:
_has_foreach_copy = hasattr(torch, "_foreach_copy_") _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: def _grouped_foreach_copy_(dsts: List[torch.Tensor], srcs: List[torch.Tensor]) -> None:
"""Call torch._foreach_copy_ grouped by (dst_dtype, src_dtype) pairs.""" """Call torch._foreach_copy_ grouped by (dst_dtype, src_dtype) pairs."""
@@ -1099,15 +1151,7 @@ class CudaGraphRunner:
if lora_ids is not None: if lora_ids is not None:
self.model_runner.lora_manager.prepare_lora_batch(forward_batch) self.model_runner.lora_manager.prepare_lora_batch(forward_batch)
attn_backend.init_forward_metadata_capture_cuda_graph( attn_backend.init_forward_metadata_out_graph(forward_batch, in_capture=True)
bs,
num_tokens,
req_pool_indices,
seq_lens,
encoder_lens,
forward_batch.forward_mode,
forward_batch.spec_info,
)
def run_once(): def run_once():
# Without this, warmup-1 caches the translation; the capture # Without this, warmup-1 caches the translation; the capture
@@ -1116,6 +1160,10 @@ class CudaGraphRunner:
if self.model_runner.is_hybrid_swa: if self.model_runner.is_hybrid_swa:
self.model_runner.token_to_kv_pool.invalidate_loc_cache() 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 = ( forward_batch.dp_local_start_pos = forward_batch.dp_local_num_tokens = (
None None
) )
@@ -1269,25 +1317,17 @@ class CudaGraphRunner:
attn_backend = self.model_runner.decode_attn_backend_group[stream_idx] attn_backend = self.model_runner.decode_attn_backend_group[stream_idx]
else: else:
attn_backend = self.attn_backend attn_backend = self.attn_backend
# FIXME: implicit channel for backends (dsv4) that need forward_batch fb_view = build_replay_fb_view(
# in replay metadata prep. Should become a real param on the interface. forward_batch=forward_batch,
attn_backend._replay_forward_batch = forward_batch buffers=buffers,
seq_lens_sum_arg = ( bs=bs,
None raw_bs=raw_bs,
if forward_batch.seq_lens_sum is None num_tokens=bs * self.num_tokens_per_bs,
else forward_batch.seq_lens_sum + (bs - raw_bs) * self.seq_len_fill_value 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( attn_backend.init_forward_metadata_out_graph(fb_view)
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
# Store fields # Store fields
self.raw_bs = raw_bs self.raw_bs = raw_bs
@@ -1233,6 +1233,46 @@ def enable_num_token_non_padded():
return get_moe_expert_parallel_world_size() > 1 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: class PPProxyTensors:
# adapted from https://github.com/vllm-project/vllm/blob/d14e98d924724b284dc5eaf8070d935e214e50c0/vllm/sequence.py#L1103 # adapted from https://github.com/vllm-project/vllm/blob/d14e98d924724b284dc5eaf8070d935e214e50c0/vllm/sequence.py#L1103
tensors: Dict[str, torch.Tensor] tensors: Dict[str, torch.Tensor]
@@ -419,7 +419,6 @@ class PiecewiseCudaGraphRunner:
return_pooled_hidden_states=self.capture_return_pooled_hidden_states, return_pooled_hidden_states=self.capture_return_pooled_hidden_states,
) )
# Attention backend
self.model_runner.attn_backend.init_forward_metadata(forward_batch) self.model_runner.attn_backend.init_forward_metadata(forward_batch)
forward_batch.dp_local_start_pos = forward_batch.dp_local_num_tokens = None 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()) set_dp_buffer_len(None, num_tokens, forward_batch.dp_padding_mode.is_max_len())
@@ -798,7 +797,6 @@ class PiecewiseCudaGraphRunner:
self.moe_fusions, self.moe_fusions,
dsa_indexers=self.dsa_indexers, 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) self.model_runner.attn_backend.init_forward_metadata(forward_batch)
output = self.model_runner.model.forward( output = self.model_runner.model.forward(
static_forward_batch.input_ids, static_forward_batch.input_ids,
-3
View File
@@ -1582,9 +1582,6 @@ class DeepseekV4Model(nn.Module):
for _attr in ("freqs_cis_c4", "freqs_cis_c128"): for _attr in ("freqs_cis_c4", "freqs_cis_c128"):
if hasattr(forward_batch, _attr): if hasattr(forward_batch, _attr):
delattr(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 use_fused = self.use_fused_mhc_post_pre
prev_residual, prev_post, prev_comb = None, None, None prev_residual, prev_post, prev_comb = None, None, None
@@ -373,6 +373,8 @@ class EAGLEDraftCudaGraphRunner:
if self.model_runner.is_hybrid_swa: if self.model_runner.is_hybrid_swa:
self.model_runner.token_to_kv_pool.invalidate_loc_cache() 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 forward_batch.dp_local_start_pos = forward_batch.dp_local_num_tokens = None
set_dp_buffer_len( set_dp_buffer_len(
global_dp_buffer_len, global_dp_buffer_len,
@@ -392,8 +394,8 @@ class EAGLEDraftCudaGraphRunner:
return ret return ret
with forward_context(ForwardContext(attn_backend=self.draft_attn_backend)): with forward_context(ForwardContext(attn_backend=self.draft_attn_backend)):
self.draft_attn_backend.init_forward_metadata_capture_cuda_graph( self.draft_attn_backend.init_forward_metadata_out_graph(
forward_batch forward_batch, in_capture=True
) )
self.deepep_adapter.capture(is_extend_in_batch=False) self.deepep_adapter.capture(is_extend_in_batch=False)
self._capture_init(run_once) self._capture_init(run_once)
@@ -509,9 +511,8 @@ class EAGLEDraftCudaGraphRunner:
buffers.seq_lens_cpu[:raw_bs].copy_(forward_batch.seq_lens_cpu) buffers.seq_lens_cpu[:raw_bs].copy_(forward_batch.seq_lens_cpu)
forward_batch.seq_lens_cpu = buffers.seq_lens_cpu[:bs] forward_batch.seq_lens_cpu = buffers.seq_lens_cpu[:bs]
self.draft_attn_backend.init_forward_metadata_replay_cuda_graph( # forward_batch.batch_size was overwritten to bs above when padding.
forward_batch, bs self.draft_attn_backend.init_forward_metadata_out_graph(forward_batch)
)
self.raw_bs = raw_bs self.raw_bs = raw_bs
self.bs = bs self.bs = bs
# TODO: The forward_batch.seq_len_sum might need to be updated to reflect the padding in the cuda graph # TODO: The forward_batch.seq_len_sum might need to be updated to reflect the padding in the cuda graph
@@ -423,14 +423,8 @@ class EAGLEDraftExtendCudaGraphRunner:
with forward_context( with forward_context(
ForwardContext(attn_backend=self.draft_extend_attn_backend) ForwardContext(attn_backend=self.draft_extend_attn_backend)
): ):
self.draft_extend_attn_backend.init_forward_metadata_capture_cuda_graph( self.draft_extend_attn_backend.init_forward_metadata_out_graph(
bs=bs, forward_batch, in_capture=True
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.deepep_adapter.capture(is_extend_in_batch=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_correct_drafts = buffers.num_correct_drafts[:bs]
forward_batch.spec_info.num_accept_tokens = buffers.num_accept_tokens[: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 seq_lens_sum = forward_batch.seq_lens_sum
if seq_lens_sum is not None: if seq_lens_sum is not None:
seq_lens_sum = seq_lens_sum + (bs - raw_bs) * self.seq_len_fill_value 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( fb_view = SimpleNamespace(
bs=bs, batch_size=bs,
forward_mode=self.forward_mode,
input_ids=getattr(forward_batch, "input_ids", None),
req_pool_indices=buffers.req_pool_indices, req_pool_indices=buffers.req_pool_indices,
seq_lens=buffers.seq_lens, seq_lens=buffers.seq_lens,
seq_lens_sum=seq_lens_sum, 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, 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 # Replay
self.raw_bs = raw_bs self.raw_bs = raw_bs
@@ -288,34 +288,33 @@ class FrozenKVMTPWorker(TpModelWorker):
self, forward_batch: ForwardBatch self, forward_batch: ForwardBatch
) -> None: ) -> None:
with self._frozen_kv_target_view(forward_batch): with self._frozen_kv_target_view(forward_batch):
self.draft_attn_backend.init_forward_metadata_capture_cuda_graph( self.draft_attn_backend.init_forward_metadata_out_graph(
forward_batch.batch_size, forward_batch, in_capture=True
forward_batch.positions.numel(),
forward_batch.req_pool_indices,
forward_batch.seq_lens,
encoder_lens=None,
forward_mode=ForwardMode.DECODE,
spec_info=None,
) )
def _init_frozen_kv_metadata_replay_cuda_graph( def _init_frozen_kv_metadata_replay_cuda_graph(
self, forward_batch: ForwardBatch, bs: int, seq_lens_sum: int self, forward_batch: ForwardBatch, bs: int, seq_lens_sum: int
) -> None: ) -> None:
with self._frozen_kv_target_view(forward_batch): from types import SimpleNamespace
self.draft_attn_backend.init_forward_metadata_replay_cuda_graph(
bs, fb_view = SimpleNamespace(
forward_batch.req_pool_indices[:bs], batch_size=bs,
forward_batch.seq_lens[:bs],
seq_lens_sum,
encoder_lens=None,
forward_mode=ForwardMode.DECODE, forward_mode=ForwardMode.DECODE,
spec_info=None, 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=( seq_lens_cpu=(
forward_batch.seq_lens_cpu[:bs] forward_batch.seq_lens_cpu[:bs]
if forward_batch.seq_lens_cpu is not None if forward_batch.seq_lens_cpu is not None
else 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_out_graph(fb_view)
def init_cuda_graphs(self) -> None: def init_cuda_graphs(self) -> None:
if self.server_args.disable_cuda_graph or self.speculative_num_steps <= 1: if self.server_args.disable_cuda_graph or self.speculative_num_steps <= 1:
@@ -485,15 +485,7 @@ class MultiLayerEagleDraftExtendCudaGraphRunner:
return ret return ret
with forward_context(ForwardContext(attn_backend=attn_backend)): with forward_context(ForwardContext(attn_backend=attn_backend)):
attn_backend.init_forward_metadata_capture_cuda_graph( attn_backend.init_forward_metadata_out_graph(forward_batch, in_capture=True)
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,
)
self.deepep_adapter.capture(is_extend_in_batch=True) self.deepep_adapter.capture(is_extend_in_batch=True)
self._capture_init(run_once) self._capture_init(run_once)
out = self._capture_graph( out = self._capture_graph(
@@ -572,19 +564,24 @@ class MultiLayerEagleDraftExtendCudaGraphRunner:
forward_batch.spec_info.positions = buffers.positions[:num_tokens] forward_batch.spec_info.positions = buffers.positions[:num_tokens]
forward_batch.spec_info.extend_seq_lens_tensor = buffers.extend_seq_lens[:bs] forward_batch.spec_info.extend_seq_lens_tensor = buffers.extend_seq_lens[:bs]
self.eagle_worker.draft_extend_attn_backend_list[ from types import SimpleNamespace
self.step
].init_forward_metadata_replay_cuda_graph( fb_view = SimpleNamespace(
bs=bs, batch_size=bs,
forward_mode=self.forward_mode,
input_ids=getattr(forward_batch, "input_ids", None),
req_pool_indices=buffers.req_pool_indices, req_pool_indices=buffers.req_pool_indices,
seq_lens=buffers.seq_lens, seq_lens=buffers.seq_lens,
seq_lens_sum=forward_batch.seq_lens_sum seq_lens_sum=forward_batch.seq_lens_sum
+ (bs - raw_bs) * self.seq_len_fill_value, + (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, 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 # Replay
self.raw_bs = raw_bs self.raw_bs = raw_bs
@@ -1079,7 +1079,6 @@ def _seed_c4_if_needed(fixture: DSV4AttentionFixture) -> None:
compress_ratios. compress_ratios.
""" """
if fixture.case.compress_ratio == 4: if fixture.case.compress_ratio == 4:
fixture.backend._maybe_upgrade_forward_metadata()
_seed_c4_sparse_indices(fixture, num_entries=_DSV4_EXTRA_ENTRIES) _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 # `c4_sparse_page_indices` back to all -1 on the next upgrade) — the
# reference must observe the same seeded indices the backend forward saw. # reference must observe the same seeded indices the backend forward saw.
_seed_c4_if_needed(fixture) _seed_c4_if_needed(fixture)
fixture.backend._maybe_upgrade_forward_metadata()
md = fixture.backend.forward_metadata.core_metadata md = fixture.backend.forward_metadata.core_metadata
runner = fixture.runner runner = fixture.runner
max_context_len = runner.req_to_token_pool.req_to_token.shape[1] 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) q_input, _ = fixture.actual_module.project(fixture.input_hidden)
with torch.no_grad(), forward_context(ForwardContext(attn_backend=fixture.backend)): with torch.no_grad(), forward_context(ForwardContext(attn_backend=fixture.backend)):
fixture.backend.init_forward_metadata(fixture.forward_batch) 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: if case.compress_ratio == 4:
_seed_c4_sparse_indices(fixture, num_entries=extra_entries) _seed_c4_sparse_indices(fixture, num_entries=extra_entries)
actual = fixture.backend.forward( actual = fixture.backend.forward(
@@ -291,36 +291,31 @@ def _init_cuda_graph_capture_metadata(backend, capture_batch_size: int, batch):
max_bs=capture_batch_size, max_bs=capture_batch_size,
max_num_tokens=batch.input_ids.numel(), max_num_tokens=batch.input_ids.numel(),
) )
backend.init_forward_metadata_capture_cuda_graph( backend.init_forward_metadata_out_graph(batch, in_capture=True)
bs=capture_batch_size, backend.init_forward_metadata_in_graph(batch)
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,
)
def _init_cuda_graph_replay_metadata(backend, capture_batch_size: int, batch): def _init_cuda_graph_replay_metadata(backend, capture_batch_size: int, batch):
# Some backends (e.g., `DeepseekV4AttnBackend`) read out-of-band attributes from types import SimpleNamespace
# off the backend during replay metadata init — production wires this in
# `sglang/srt/model_executor/cuda_graph_runner.py:1234`. Mirror that fb_view = SimpleNamespace(
# contract so backends that don't use it just store-and-clear the field. batch_size=capture_batch_size,
backend._replay_forward_batch = batch forward_mode=batch.forward_mode,
try: actual_forward_mode=batch.forward_mode,
backend.init_forward_metadata_replay_cuda_graph( input_ids=batch.input_ids,
bs=capture_batch_size, positions=getattr(batch, "positions", None),
req_pool_indices=batch.req_pool_indices, req_pool_indices=batch.req_pool_indices,
seq_lens=batch.seq_lens, seq_lens=batch.seq_lens,
seq_lens_sum=batch.seq_lens_sum, 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, 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,
) )
finally: backend.init_forward_metadata_out_graph(fb_view)
backend._replay_forward_batch = None # 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 # Best-effort metadata-shape sanity check — catches negative kv_lens and
# non-monotonic indptr that would otherwise leave real-row output correct # non-monotonic indptr that would otherwise leave real-row output correct
# but corrupt padded-row scratch state. See `metadata_invariants.py`. # but corrupt padded-row scratch state. See `metadata_invariants.py`.
@@ -1022,13 +1022,6 @@ class EagleDraftExtendCudaGraphRunnerAdapter:
make_forward_batch: Callable[ make_forward_batch: Callable[
[Any, Any, Any, EagleDraftRunnerSettings], ForwardBatch [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] = ( check_case: Callable[[Any, EagleDraftRunnerSettings], None] = (
lambda _case, _settings: 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) adapter.prepare_replay_state(graph_fixture, case, draft_inputs, settings)
testcase.assertTrue(graph_runner.can_run(graph_batch)) 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) 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) adapter.assert_outputs_close(actual, expected, settings)
finally: finally:
_reset_cuda_graph_test_buffers() _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( def run_dsv4_eagle_draft_extend_cuda_graph_runner_case(
testcase, testcase,
case: DSV4AttentionCase, case: DSV4AttentionCase,
@@ -2001,7 +1974,6 @@ def run_dsv4_eagle_draft_extend_cuda_graph_runner_case(
make_draft_inputs=_make_dsv4_draft_extend_inputs, make_draft_inputs=_make_dsv4_draft_extend_inputs,
prepare_replay_state=_prepare_dsv4_draft_extend_replay_state, prepare_replay_state=_prepare_dsv4_draft_extend_replay_state,
make_forward_batch=_make_dsv4_eagle_draft_extend_forward_batch, 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, assert_outputs_close=_dsv4_assert_draft_extend_outputs_close,
) )
run_eagle_draft_extend_cuda_graph_runner_case( run_eagle_draft_extend_cuda_graph_runner_case(
@@ -1,5 +1,6 @@
import math import math
import unittest import unittest
from types import SimpleNamespace
import numpy as np import numpy as np
import torch import torch
@@ -1030,15 +1031,15 @@ class TestTRTLLMMLA(CustomTestCase):
) )
req_pool_indices = torch.arange(batch_size, device=config["device"]) req_pool_indices = torch.arange(batch_size, device=config["device"])
backend.init_forward_metadata_capture_cuda_graph( capture_fb = SimpleNamespace(
bs=batch_size, batch_size=batch_size,
num_tokens=batch_size, forward_mode=ForwardMode.DECODE,
req_pool_indices=req_pool_indices, req_pool_indices=req_pool_indices,
seq_lens=seq_lens, seq_lens=seq_lens,
encoder_lens=None, positions=torch.arange(batch_size, device=config["device"]),
forward_mode=ForwardMode.DECODE,
spec_info=None, spec_info=None,
) )
backend.init_forward_metadata_out_graph(capture_fb, in_capture=True)
# Verify capture metadata # Verify capture metadata
self.assertIn(batch_size, backend.decode_cuda_graph_metadata) self.assertIn(batch_size, backend.decode_cuda_graph_metadata)
@@ -1054,16 +1055,15 @@ class TestTRTLLMMLA(CustomTestCase):
device=config["device"], device=config["device"],
) )
backend.init_forward_metadata_replay_cuda_graph( replay_fb = SimpleNamespace(
bs=batch_size, batch_size=batch_size,
forward_mode=ForwardMode.DECODE,
req_pool_indices=req_pool_indices, req_pool_indices=req_pool_indices,
seq_lens=new_seq_lens, 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(), seq_lens_cpu=new_seq_lens.cpu(),
spec_info=None,
) )
backend.init_forward_metadata_out_graph(replay_fb)
# Verify replay updated the metadata # Verify replay updated the metadata
replay_metadata = backend.forward_decode_metadata replay_metadata = backend.forward_decode_metadata
@@ -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()
@@ -94,14 +94,15 @@ class TestTboAttnDenseAttentionBackendCorrectness(CustomTestCase):
"FA3 backend requires SM 80-90", "FA3 backend requires SM 80-90",
) )
def test_tbo_target_verify_cuda_graph_capture_delegates_to_primary_capture(self): def test_tbo_target_verify_cuda_graph_capture_delegates_to_primary_capture(self):
"""TBO capture must invoke ``primary.init_forward_metadata_capture_cuda_graph``, """TBO capture must dispatch primary's
not ``primary.init_forward_metadata_replay_cuda_graph``. ``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 Backends like FlashAttention store per-bs metadata in dicts populated
only by their capture path (via ``_bind_metadata_buffers``). If TBO only by the in_capture=True branch (via ``_bind_metadata_buffers``).
short-circuits its capture to its own replay (which delegates to If TBO short-circuits its capture to its own replay path, those dicts
``primary.replay``), those dicts are empty and replay raises are empty and replay raises ``KeyError: bs``. Reproduces the
``KeyError: bs``. Reproduces the deepep-4-gpu-h100 failure where deepep-4-gpu-h100 failure where
``flashattention_backend.target_verify_metadata[bs]`` lookup blew up ``flashattention_backend.target_verify_metadata[bs]`` lookup blew up
during ``init_device_graphs``. during ``init_device_graphs``.
@@ -126,19 +127,100 @@ class TestTboAttnDenseAttentionBackendCorrectness(CustomTestCase):
capture_bs = case.batch_size capture_bs = case.batch_size
num_tokens = sum(case.extend_lens) num_tokens = sum(case.extend_lens)
wrapper.init_cuda_graph_state(max_bs=capture_bs, max_num_tokens=num_tokens) 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 wrapper.init_forward_metadata_out_graph(batch, in_capture=True)
# primary.replay (instead of primary.capture) reads an unpopulated
# ``target_verify_metadata[bs]`` dict and raises KeyError. @unittest.skipIf(
wrapper.init_forward_metadata_capture_cuda_graph( get_device_sm() >= 100 or get_device_sm() < 80,
bs=capture_bs, "FA3 backend requires SM 80-90",
num_tokens=num_tokens, )
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, req_pool_indices=batch.req_pool_indices,
seq_lens=batch.seq_lens, seq_lens=batch.seq_lens,
encoder_lens=batch.encoder_lens, seq_lens_sum=int(batch.seq_lens_cpu.sum()),
forward_mode=batch.forward_mode, seq_lens_cpu=batch.seq_lens_cpu,
encoder_lens=None,
out_cache_loc=batch.out_cache_loc,
spec_info=batch.spec_info, 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__": if __name__ == "__main__":
unittest.main() unittest.main()
@@ -399,38 +399,35 @@ class TestTritonGDNBackendCorrectness(CustomTestCase):
linear_attn_backend.init_forward_metadata, sentinel_forward_batch 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): def test_hybrid_dispatch_replay_init_forward_metadata_fan_out(self):
backend, full_attn_backend, linear_attn_backend = ( backend, full_attn_backend, linear_attn_backend = (
self._make_dispatch_spy_backend() self._make_dispatch_spy_backend()
) )
sentinel_req_pool = object() fb = self._make_sentinel_fb()
sentinel_seq_lens = object() backend.init_forward_metadata_out_graph(fb)
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,
)
# We assert sentinel identity rather than exact (args, kwargs) shape # We assert sentinel identity rather than exact (args, kwargs) shape
# so a positional↔keyword refactor inside `HybridLinearAttnBackend` # 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): for sub_backend in (full_attn_backend, linear_attn_backend):
self._assert_fanout_forwarded( self._assert_fanout_forwarded(
sub_backend.init_forward_metadata_replay_cuda_graph, sub_backend.init_forward_metadata_out_graph, fb
sentinel_req_pool,
sentinel_seq_lens,
sentinel_seq_lens_cpu,
sentinel_spec_info,
ForwardMode.DECODE,
) )
def test_hybrid_dispatch_capture_init_forward_metadata_fan_out(self): 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 = ( backend, full_attn_backend, linear_attn_backend = (
self._make_dispatch_spy_backend() self._make_dispatch_spy_backend()
) )
sentinel_req_pool = object() fb = self._make_sentinel_fb()
sentinel_seq_lens = object() backend.init_forward_metadata_out_graph(fb, in_capture=True)
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,
)
for sub_backend in (full_attn_backend, linear_attn_backend): for sub_backend in (full_attn_backend, linear_attn_backend):
self._assert_fanout_forwarded( self._assert_fanout_forwarded(
sub_backend.init_forward_metadata_capture_cuda_graph, sub_backend.init_forward_metadata_out_graph, fb
sentinel_req_pool,
sentinel_seq_lens,
sentinel_spec_info,
ForwardMode.DECODE,
) )
@@ -216,7 +216,7 @@ class TestTritonMamba2BackendCorrectness(CustomTestCase):
# padding the `forward_batch.input_ids` / `out_cache_loc`. # padding the `forward_batch.input_ids` / `out_cache_loc`.
def test_mamba2_replay_metadata_padding_indices(self): 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 # `seq_lens_cpu=[5, 1, 1]` (two trailing rows at the cuda-graph
# fill value 1) so the padding-row count is observable in # fill value 1) so the padding-row count is observable in
# `state_indices_list[bs - 1]`. # `state_indices_list[bs - 1]`.
@@ -245,16 +245,17 @@ class TestTritonMamba2BackendCorrectness(CustomTestCase):
req_pool_indices req_pool_indices
] = torch.tensor([7, 0, 0], dtype=torch.int32, device=device) ] = torch.tensor([7, 0, 0], dtype=torch.int32, device=device)
backend.init_forward_metadata_replay_cuda_graph( fb = SimpleNamespace(
bs=bs, batch_size=bs,
forward_mode=ForwardMode.DECODE,
req_pool_indices=req_pool_indices, req_pool_indices=req_pool_indices,
seq_lens=seq_lens, 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_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() state_indices = backend.state_indices_list[bs - 1].cpu().tolist()
self.assertEqual( self.assertEqual(
@@ -326,62 +327,44 @@ class TestTritonMamba2BackendCorrectness(CustomTestCase):
linear_attn_backend.init_forward_metadata, sentinel_forward_batch 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): def test_hybrid_dispatch_replay_init_forward_metadata_fan_out(self):
backend, full_attn_backend, linear_attn_backend = ( backend, full_attn_backend, linear_attn_backend = (
self._make_dispatch_spy_backend() self._make_dispatch_spy_backend()
) )
sentinel_req_pool = object() fb = self._make_sentinel_fb()
sentinel_seq_lens = object() backend.init_forward_metadata_out_graph(fb)
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,
)
for sub_backend in (full_attn_backend, linear_attn_backend): for sub_backend in (full_attn_backend, linear_attn_backend):
self._assert_fanout_forwarded( self._assert_fanout_forwarded(
sub_backend.init_forward_metadata_replay_cuda_graph, sub_backend.init_forward_metadata_out_graph, fb
sentinel_req_pool,
sentinel_seq_lens,
sentinel_seq_lens_cpu,
sentinel_spec_info,
ForwardMode.DECODE,
) )
def test_hybrid_dispatch_capture_init_forward_metadata_fan_out(self): def test_hybrid_dispatch_capture_init_forward_metadata_fan_out(self):
backend, full_attn_backend, linear_attn_backend = ( backend, full_attn_backend, linear_attn_backend = (
self._make_dispatch_spy_backend() self._make_dispatch_spy_backend()
) )
sentinel_req_pool = object() fb = self._make_sentinel_fb()
sentinel_seq_lens = object() backend.init_forward_metadata_out_graph(fb, in_capture=True)
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,
)
for sub_backend in (full_attn_backend, linear_attn_backend): for sub_backend in (full_attn_backend, linear_attn_backend):
self._assert_fanout_forwarded( self._assert_fanout_forwarded(
sub_backend.init_forward_metadata_capture_cuda_graph, sub_backend.init_forward_metadata_out_graph, fb
sentinel_req_pool,
sentinel_seq_lens,
sentinel_spec_info,
ForwardMode.DECODE,
) )