[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
def init_forward_metadata_out_graph(
self,
forward_batch: ForwardBatch,
in_capture: bool = False,
):
bs = forward_batch.batch_size
if in_capture:
self._init_cuda_graph_metadata(
bs, forward_batch.forward_mode, forward_batch.seq_lens
)
self._apply_cuda_graph_metadata(
bs=bs,
req_pool_indices=forward_batch.req_pool_indices,
seq_lens=forward_batch.seq_lens,
seq_lens_cpu=(
forward_batch.seq_lens.cpu()
if in_capture
else forward_batch.seq_lens_cpu
),
forward_mode=forward_batch.forward_mode,
spec_info=forward_batch.spec_info,
)
def init_forward_metadata(self, forward_batch: ForwardBatch):
"""Init the metadata for a forward pass."""
self.forward_metadata = ForwardMetadata()
@@ -538,39 +561,19 @@ class AscendAttnBackend(AttentionBackend):
self.graph_metadata[bs] = metadata
return metadata
def init_forward_metadata_capture_cuda_graph(
def _apply_cuda_graph_metadata(
self,
bs: int,
num_tokens: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
encoder_lens: Optional[torch.Tensor],
seq_lens_cpu: torch.Tensor,
forward_mode: ForwardMode,
spec_info: Optional[SpecInput],
):
self._init_cuda_graph_metadata(bs, forward_mode, seq_lens)
self.init_forward_metadata_replay_cuda_graph(
bs=bs,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
seq_lens_sum=None,
encoder_lens=encoder_lens,
forward_mode=forward_mode,
spec_info=spec_info,
seq_lens_cpu=seq_lens.cpu(),
)
"""Shared capture+replay body for the cuda-graph init path.
def init_forward_metadata_replay_cuda_graph(
self,
bs: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
seq_lens_sum: int,
encoder_lens: Optional[torch.Tensor],
forward_mode: ForwardMode,
spec_info: Optional[SpecInput],
seq_lens_cpu: Optional[torch.Tensor],
):
Public entry: :py:meth:`init_forward_metadata_out_graph`.
"""
metadata = self.graph_metadata[bs]
max_len = seq_lens_cpu[:bs].max().item()
if forward_mode.is_target_verify():
@@ -2395,6 +2398,32 @@ class AscendAttnMultiStepDraftBackend:
for i in range(self.speculative_num_steps - 1):
call_fn(i, forward_batch)
def init_forward_metadata_out_graph(
self,
forward_batch: ForwardBatch,
in_capture: bool = False,
):
from sglang.srt.model_executor.forward_batch_info import build_inner_fb_view
inner_fb = build_inner_fb_view(
forward_batch,
bs=forward_batch.batch_size,
forward_mode=ForwardMode.DECODE,
)
def call_fn(i, _forward_batch):
self.attn_backends[i].init_forward_metadata_out_graph(
inner_fb, in_capture=in_capture
)
self.common_template(forward_batch, call_fn)
def init_forward_metadata_in_graph(self, forward_batch: ForwardBatch) -> None:
def call_fn(i, _forward_batch):
self.attn_backends[i].init_forward_metadata_in_graph(forward_batch)
self.common_template(forward_batch, call_fn)
def init_forward_metadata(self, forward_batch: ForwardBatch):
def call_fn(i, forward_batch):
assert forward_batch.spec_info is not None
@@ -2405,34 +2434,3 @@ class AscendAttnMultiStepDraftBackend:
def init_cuda_graph_state(self, max_bs, max_num_tokens):
for i in range(self.speculative_num_steps):
self.attn_backends[i].init_cuda_graph_state(max_bs, max_num_tokens)
def init_forward_metadata_capture_cuda_graph(self, forward_batch: ForwardBatch):
def call_fn(i, forward_batch):
self.attn_backends[i].init_forward_metadata_capture_cuda_graph(
forward_batch.batch_size,
forward_batch.batch_size * self.topk,
forward_batch.req_pool_indices,
forward_batch.seq_lens,
encoder_lens=None,
forward_mode=ForwardMode.DECODE,
spec_info=forward_batch.spec_info,
)
self.common_template(forward_batch, call_fn)
def init_forward_metadata_replay_cuda_graph(
self, forward_batch: ForwardBatch, bs: int
):
def call_fn(i, forward_batch):
self.attn_backends[i].init_forward_metadata_replay_cuda_graph(
bs,
forward_batch.req_pool_indices,
forward_batch.seq_lens,
seq_lens_sum=-1,
encoder_lens=None,
forward_mode=ForwardMode.DECODE,
spec_info=forward_batch.spec_info,
seq_lens_cpu=forward_batch.seq_lens_cpu,
)
self.common_template(forward_batch, call_fn)
@@ -72,6 +72,21 @@ class AscendGDNAttnBackend(AscendMambaAttnBackendBase):
else:
self.ssm_state_indices = cache_indices
def init_forward_metadata_out_graph(
self,
forward_batch: ForwardBatch,
in_capture: bool = False,
):
if forward_batch.forward_mode.is_draft_extend(True):
return
super().init_forward_metadata_out_graph(forward_batch, in_capture=in_capture)
self.prepare_gdn_inputs(
forward_batch.batch_size,
forward_batch.forward_mode,
forward_batch.spec_info,
)
self.graph_mode = True
def init_forward_metadata(self, forward_batch: ForwardBatch):
if forward_batch.forward_mode.is_draft_extend(True):
return
@@ -83,53 +98,6 @@ class AscendGDNAttnBackend(AscendMambaAttnBackendBase):
)
self.graph_mode = False
def init_forward_metadata_capture_cuda_graph(
self,
bs: int,
num_tokens: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
encoder_lens: Optional[torch.Tensor],
forward_mode: ForwardMode,
spec_info: Optional[Union[EagleDraftInput, EagleVerifyInput]],
):
self.init_forward_metadata_replay_cuda_graph(
bs=bs,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
seq_lens_sum=None,
encoder_lens=encoder_lens,
forward_mode=forward_mode,
spec_info=spec_info,
seq_lens_cpu=seq_lens.cpu(),
)
def init_forward_metadata_replay_cuda_graph(
self,
bs: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
seq_lens_sum: int,
encoder_lens: Optional[torch.Tensor],
forward_mode: ForwardMode,
spec_info: Optional[Union[EagleDraftInput, EagleVerifyInput]],
seq_lens_cpu: Optional[torch.Tensor],
):
if forward_mode.is_draft_extend(True):
return
super().init_forward_metadata_replay_cuda_graph(
bs,
req_pool_indices,
seq_lens,
seq_lens_sum,
encoder_lens,
forward_mode,
spec_info,
seq_lens_cpu,
)
self.prepare_gdn_inputs(bs, forward_mode, spec_info)
self.graph_mode = True
def forward_decode(
self,
layer: RadixLinearAttention,
@@ -131,9 +131,13 @@ class AscendMambaAttnBackendBase(MambaAttnBackendBase):
spec_info: Optional[SpecInput],
seq_lens_cpu: Optional[torch.Tensor],
):
num_padding = torch.count_nonzero(
seq_lens_cpu == self.get_cuda_graph_seq_len_fill_value()
)
# out_graph passes seq_lens_cpu=None at capture; mirror the base guard.
if seq_lens_cpu is None:
num_padding = 0
else:
num_padding = torch.count_nonzero(
seq_lens_cpu == self.get_cuda_graph_seq_len_fill_value()
)
# Make sure forward metadata is correctly handled for padding reqs
req_pool_indices[bs - num_padding :] = 0
mamba_indices = self.req_to_token_pool.get_mamba_indices(req_pool_indices)
@@ -820,6 +820,25 @@ class AiterAttnBackend(AttentionBackend):
)
return output
def init_forward_metadata_out_graph(
self,
forward_batch: ForwardBatch,
in_capture: bool = False,
):
seq_lens_cpu = (
forward_batch.seq_lens.cpu() if in_capture else forward_batch.seq_lens_cpu
)
self._apply_cuda_graph_metadata(
bs=forward_batch.batch_size,
req_pool_indices=forward_batch.req_pool_indices,
seq_lens=forward_batch.seq_lens,
seq_lens_sum=None if in_capture else forward_batch.seq_lens_sum,
encoder_lens=forward_batch.encoder_lens,
forward_mode=forward_batch.forward_mode,
spec_info=forward_batch.spec_info,
seq_lens_cpu=seq_lens_cpu,
)
def init_forward_metadata(self, forward_batch: ForwardBatch):
"""Init auxiliary variables for aiter attention backend."""
@@ -1482,28 +1501,7 @@ class AiterAttnBackend(AttentionBackend):
device=self.device,
)
def init_forward_metadata_capture_cuda_graph(
self,
bs: int,
num_tokens: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
encoder_lens: Optional[torch.Tensor],
forward_mode: ForwardMode,
spec_info: Optional[SpecInput],
):
self.init_forward_metadata_replay_cuda_graph(
bs=bs,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
seq_lens_sum=None,
encoder_lens=encoder_lens,
forward_mode=forward_mode,
spec_info=spec_info,
seq_lens_cpu=seq_lens.cpu(),
)
def init_forward_metadata_replay_cuda_graph(
def _apply_cuda_graph_metadata(
self,
bs: int,
req_pool_indices: torch.Tensor,
@@ -2866,33 +2864,26 @@ class AiterMultiStepDraftBackend:
max_bs, max_num_tokens, kv_indices_buf=self.cuda_graph_kv_indices[i]
)
def init_forward_metadata_capture_cuda_graph(self, forward_batch: ForwardBatch):
def call_fn(i, forward_batch):
self.attn_backends[i].init_forward_metadata_capture_cuda_graph(
forward_batch.batch_size,
forward_batch.batch_size * self.topk,
forward_batch.req_pool_indices,
forward_batch.seq_lens,
encoder_lens=None,
forward_mode=ForwardMode.DECODE,
spec_info=forward_batch.spec_info,
)
self.common_template(forward_batch, self.cuda_graph_kv_indices, call_fn)
def init_forward_metadata_replay_cuda_graph(
self, forward_batch: ForwardBatch, bs: int
def init_forward_metadata_out_graph(
self,
forward_batch: ForwardBatch,
in_capture: bool = False,
):
def call_fn(i, forward_batch):
self.attn_backends[i].init_forward_metadata_replay_cuda_graph(
bs,
forward_batch.req_pool_indices,
forward_batch.seq_lens,
seq_lens_sum=-1,
encoder_lens=None,
forward_mode=ForwardMode.DECODE,
spec_info=forward_batch.spec_info,
seq_lens_cpu=forward_batch.seq_lens_cpu,
from sglang.srt.model_executor.forward_batch_info import build_inner_fb_view
inner_fb = build_inner_fb_view(
forward_batch,
bs=forward_batch.batch_size,
forward_mode=ForwardMode.DECODE,
)
def call_fn(i, _forward_batch):
self.attn_backends[i].init_forward_metadata_out_graph(
inner_fb, in_capture=in_capture
)
self.common_template(forward_batch, self.cuda_graph_kv_indices, call_fn)
def init_forward_metadata_in_graph(self, forward_batch: ForwardBatch) -> None:
for attn_backend in self.attn_backends:
attn_backend.init_forward_metadata_in_graph(forward_batch)
@@ -1,6 +1,6 @@
from __future__ import annotations
from abc import ABC, abstractmethod
from abc import ABC
from typing import TYPE_CHECKING, Optional
import torch
@@ -11,52 +11,84 @@ from sglang.srt.utils.common import is_npu
if TYPE_CHECKING:
from sglang.srt.layers.attention.dsa.dsa_indexer import BaseIndexerMetadata
from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.speculative.spec_info import SpecInput
class AttentionBackend(ABC):
"""The base class of attention backends"""
"""The base class of attention backends.
Forward-data init contract (3 methods):
- ``init_forward_metadata(fb)`` — eager entry point. Default is a wrapper
that calls ``_out_graph(fb)`` then ``_in_graph(fb)``. Backends may
override to keep an independent eager body.
- ``init_forward_metadata_out_graph(fb, in_capture=False)`` — per-iter
metadata prep, runs outside ``with graph.capture():``. Capture
sites pass ``in_capture=True``; replay/eager use the default
``False``. Backends read ``in_capture`` only when capture / replay
bodies diverge.
- ``init_forward_metadata_in_graph(fb)`` — graph-recordable static-shape
GPU op, runs inside ``with graph.capture():`` at capture time and
is auto-replayed by ``graph.replay()``. Default is no-op.
The legacy ``init_forward_metadata_capture_cuda_graph`` and
``init_forward_metadata_replay_cuda_graph`` overrides are fully
deprecated and removed from the ABC: out-of-tree backends overriding
those must migrate to ``init_forward_metadata_out_graph(fb, in_capture)``.
"""
def init_forward_metadata(self, forward_batch: ForwardBatch):
"""Eager entry point. Default = ``_out_graph(fb) + _in_graph(fb)``.
Backends may override to keep an independent eager body.
"""
self.init_forward_metadata_out_graph(forward_batch)
self.init_forward_metadata_in_graph(forward_batch)
def init_forward_metadata_out_graph(
self,
forward_batch: ForwardBatch,
in_capture: bool = False,
):
"""Per-iter metadata prep — runs outside ``with graph.capture():``.
Called at:
* capture: before ``with graph.capture():`` (caller passes
``in_capture=True``).
* replay: before ``graph.replay()`` (``in_capture=False``).
* eager: via :py:meth:`init_forward_metadata` default wrapper
(``in_capture=False``).
Backends read ``in_capture`` only when capture / replay bodies
diverge (e.g., snapshot metadata, swap buffer pointers, install
temp workspace). Host op / dynamic-shape / non-graph-recordable
logic lives here.
Default: no-op.
"""
def init_forward_metadata_in_graph(self, forward_batch: ForwardBatch):
"""Graph-recordable static-shape GPU op.
Runs inside ``with graph.capture():`` at capture time; recorded
ops auto-execute at replay via ``graph.replay()``.
Lint contract for overrides: body must NOT call ``.item()`` /
``.cpu()`` / ``.tolist()`` / dynamic-shape ``torch.empty()``.
Such ops belong in :py:meth:`init_forward_metadata_out_graph`; they
cannot be recorded into a cuda graph.
Default: no-op.
"""
# Opt out only when this backend never reads seq_lens_cpu / seq_lens_sum.
needs_cpu_seq_lens: bool = True
@abstractmethod
def init_forward_metadata(self, forward_batch: ForwardBatch):
"""Init the metadata for a forward pass."""
raise NotImplementedError()
def init_cuda_graph_state(self, max_bs: int, max_num_tokens: int):
"""Init the global shared states for cuda graph."""
raise NotImplementedError()
def init_forward_metadata_capture_cuda_graph(
self,
bs: int,
num_tokens: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
encoder_lens: Optional[torch.Tensor],
forward_mode: ForwardMode,
spec_info: Optional[SpecInput],
):
"""Init the metadata for a forward pass for capturing a cuda graph."""
raise NotImplementedError()
def init_forward_metadata_replay_cuda_graph(
self,
bs: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
seq_lens_sum: int,
encoder_lens: Optional[torch.Tensor],
forward_mode: ForwardMode,
spec_info: Optional[SpecInput],
seq_lens_cpu: Optional[torch.Tensor],
):
"""Init the metadata for a forward pass for replaying a cuda graph."""
raise NotImplementedError()
def get_cuda_graph_seq_len_fill_value(self):
"""Get the fill value for padded seq lens. Typically, it is 0 or 1."""
raise NotImplementedError()
@@ -14,13 +14,12 @@ import triton
from sglang.srt.layers.attention.flashinfer_mla_backend import FlashInferMLAAttnBackend
from sglang.srt.layers.attention.utils import create_flashmla_kv_indices_triton
from sglang.srt.layers.dp_attention import get_attention_tp_size
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.utils import is_cuda
if TYPE_CHECKING:
from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.model_executor.model_runner import ModelRunner
from sglang.srt.speculative.spec_info import SpecInput
_is_cuda = is_cuda()
if _is_cuda:
@@ -79,6 +78,37 @@ class CutlassMLABackend(FlashInferMLAAttnBackend):
self.q_data_type = model_runner.dtype
self.kv_cache_dim = self.kv_lora_rank + self.qk_rope_head_dim
def init_forward_metadata_out_graph(
self,
forward_batch: ForwardBatch,
in_capture: bool = False,
):
bs = forward_batch.batch_size
forward_mode = forward_batch.forward_mode
spec_info = forward_batch.spec_info
if forward_mode.is_decode_or_idle() and spec_info is None:
create_flashmla_kv_indices_triton[(bs,)](
self.req_to_token,
forward_batch.req_pool_indices[:bs],
forward_batch.seq_lens[:bs],
None,
self.cuda_graph_kv_indices,
self.req_to_token.stride(0),
self.cuda_graph_kv_indices.stride(0),
PAGED_SIZE=PAGE_SIZE,
)
if in_capture:
max_seqlen_pad = self.cuda_graph_kv_indices.shape[1]
self.forward_metadata = CutlassMLADecodeMetadata(
self.cuda_graph_mla_workspace,
self.cuda_graph_kv_indices[:bs, :max_seqlen_pad],
)
else:
super().init_forward_metadata_out_graph(
forward_batch, in_capture=in_capture
)
def init_forward_metadata(self, forward_batch: ForwardBatch):
bs = forward_batch.batch_size
@@ -143,77 +173,6 @@ class CutlassMLABackend(FlashInferMLAAttnBackend):
)
self.cuda_graph_kv_indices = cuda_graph_kv_indices
def init_forward_metadata_capture_cuda_graph(
self,
bs: int,
num_tokens: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
encoder_lens: Optional[torch.Tensor],
forward_mode: ForwardMode,
spec_info: Optional[SpecInput],
):
if forward_mode.is_decode_or_idle() and spec_info is None:
self.init_forward_metadata_replay_cuda_graph(
bs=bs,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
seq_lens_sum=None,
encoder_lens=encoder_lens,
forward_mode=forward_mode,
spec_info=spec_info,
seq_lens_cpu=None,
)
max_seqlen_pad = self.cuda_graph_kv_indices.shape[1]
self.forward_metadata = CutlassMLADecodeMetadata(
self.cuda_graph_mla_workspace,
self.cuda_graph_kv_indices[:bs, :max_seqlen_pad],
)
else:
super().init_forward_metadata_capture_cuda_graph(
bs,
num_tokens,
req_pool_indices,
seq_lens,
encoder_lens,
forward_mode,
spec_info,
)
def init_forward_metadata_replay_cuda_graph(
self,
bs: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
seq_lens_sum: int,
encoder_lens: Optional[torch.Tensor],
forward_mode: ForwardMode,
spec_info: Optional[SpecInput],
seq_lens_cpu: Optional[torch.Tensor],
):
if forward_mode.is_decode_or_idle():
create_flashmla_kv_indices_triton[(bs,)](
self.req_to_token,
req_pool_indices[:bs],
seq_lens[:bs],
None,
self.cuda_graph_kv_indices,
self.req_to_token.stride(0),
self.cuda_graph_kv_indices.stride(0),
PAGED_SIZE=PAGE_SIZE,
)
else:
super().init_forward_metadata_replay_cuda_graph(
bs,
req_pool_indices,
seq_lens,
seq_lens_sum,
encoder_lens,
forward_mode,
spec_info,
seq_lens_cpu,
)
def get_cuda_graph_seq_len_fill_value(self):
return 1
@@ -54,7 +54,6 @@ from sglang.srt.layers.dp_attention import (
from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.speculative.eagle_utils import per_step_draft_out_cache_loc
from sglang.srt.speculative.spec_info import SpecInput
from sglang.srt.utils import ceil_align
from sglang.srt.utils.common import is_sm120_supported
@@ -387,7 +386,6 @@ class DeepseekV4AttnBackend(
DSV4RawVerifyMetadata,
DSV4RawDecodeMetadata,
] = None
self._replay_forward_batch: Optional[ForwardBatch] = None # FIXME: out-of-band
def _move_to_device(self, x: List[int]) -> torch.Tensor:
pin_tensor = torch.tensor(x, dtype=torch.int32, pin_memory=True)
@@ -671,6 +669,136 @@ class DeepseekV4AttnBackend(
use_prefill_cuda_graph=use_prefill_cuda_graph,
)
def init_forward_metadata_in_graph(self, forward_batch: ForwardBatch) -> None:
# Upgrade Raw->Full so the c4/c128 compress + core_attn + indexer
# materialization is recorded inside the cuda graph; a no-op (Full
# already) when PREP_IN_CUDA_GRAPH=0.
if isinstance(self.forward_metadata, DSV4RawVerifyMetadata):
self.forward_metadata = self.make_forward_metadata_from_raw_verify(
raw_metadata=self.forward_metadata,
)
elif isinstance(self.forward_metadata, DSV4RawDecodeMetadata):
self.forward_metadata = self.make_forward_metadata_from_raw_decode(
raw_metadata=self.forward_metadata,
)
def init_forward_metadata_out_graph(
self,
forward_batch: ForwardBatch,
in_capture: bool = False,
) -> None:
bucket = _GraphBucket.of(forward_batch.forward_mode)
bs = forward_batch.batch_size
req_pool_indices = forward_batch.req_pool_indices
seq_lens = forward_batch.seq_lens
if in_capture:
# Captured graph does no real cache writes, so synthesize a dummy
# out_cache_loc per bucket (replay supplies the real value).
assert req_pool_indices.size(0) == bs
assert seq_lens.size(0) == bs
num_tokens = forward_batch.positions.numel()
if bucket == _GraphBucket.DECODE_OR_IDLE:
out_cache_loc = torch.zeros_like(seq_lens)
elif bucket == _GraphBucket.TARGET_VERIFY:
out_cache_loc = torch.zeros(num_tokens, **self.cuda_int32_kwargs)
else:
out_cache_loc = None
actual_forward_mode = forward_batch.forward_mode
seq_lens_sum = int(seq_lens.sum().item())
seq_lens_cpu = seq_lens.cpu()
else:
out_cache_loc = forward_batch.out_cache_loc
actual_forward_mode = getattr(
forward_batch, "actual_forward_mode", forward_batch.forward_mode
)
seq_lens_sum = forward_batch.seq_lens_sum
seq_lens_cpu = forward_batch.seq_lens_cpu
if actual_forward_mode == ForwardMode.IDLE:
logger.debug(
f"[IDLE replay] bs={bs}, "
f"local_seq_lens_len={len(seq_lens)}, "
f"has_graph={bs in self.cuda_graph_metadata_of_bucket_and_bs[_GraphBucket.DECODE_OR_IDLE]}"
)
device = seq_lens.device
seq_lens = torch.ones(bs, dtype=seq_lens.dtype, device=device)
seq_lens_cpu = torch.ones(bs, dtype=torch.int64)
seq_lens_sum = bs
req_pool_indices = torch.zeros(
bs, dtype=req_pool_indices.dtype, device=device
)
out_cache_loc = torch.zeros(bs, dtype=torch.int64, device=device)
assert seq_lens_cpu is not None
seq_lens = seq_lens[:bs]
seq_lens_cpu = seq_lens_cpu[:bs]
req_pool_indices = req_pool_indices[:bs]
actual_max_seq_len = seq_lens_cpu.max().item()
chosen_max_seq_len = self.MAX_SEQ_LEN_FOR_CAPTURE
assert actual_max_seq_len <= chosen_max_seq_len
if bucket == _GraphBucket.DECODE_OR_IDLE:
assert out_cache_loc is not None
assert len(out_cache_loc.shape) == 1, f"{out_cache_loc.shape=}"
out_cache_loc_padded = torch.nn.functional.pad(
out_cache_loc,
pad=(0, bs - len(out_cache_loc)),
mode="constant",
value=0,
)
temp_metadata = self.init_forward_metadata_decode(
max_seq_len=chosen_max_seq_len,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
out_cache_loc=out_cache_loc_padded,
)
elif bucket == _GraphBucket.TARGET_VERIFY:
assert out_cache_loc is not None
num_tokens_v = self.speculative_num_draft_tokens * bs
out_cache_loc_padded = torch.nn.functional.pad(
out_cache_loc,
pad=(0, num_tokens_v - len(out_cache_loc)),
mode="constant",
value=0,
)
temp_metadata = self.init_forward_metadata_target_verify(
max_seq_len=chosen_max_seq_len,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
out_cache_loc=out_cache_loc_padded,
use_prefill_cuda_graph=True,
)
elif bucket == _GraphBucket.DRAFT_EXTEND:
num_tokens_per_bs = self.draft_extend_num_tokens_per_bs
temp_metadata = self.init_forward_metadata_draft_extend(
max_seq_len=chosen_max_seq_len,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
seq_lens_cpu=seq_lens_cpu.tolist(),
num_tokens_per_bs=num_tokens_per_bs,
use_prefill_cuda_graph=True,
)
else:
raise NotImplementedError
self.replay_cuda_graph_metadata_from(
bs=bs, temp_metadata=temp_metadata, bucket=bucket
)
if in_capture:
# Preserve _current_capture_raw for on_after_cuda_graph_warmup
metadata = self.forward_metadata
self._current_capture_raw = (
metadata
if isinstance(
metadata,
(DSV4RawDecodeMetadata, DSV4RawVerifyMetadata),
)
else None
)
def init_forward_metadata(self, forward_batch: ForwardBatch) -> None:
if self.mtp_enabled and forward_batch.forward_mode.is_idle():
return
@@ -686,8 +814,7 @@ class DeepseekV4AttnBackend(
if forward_batch.forward_mode.is_decode_or_idle():
# DSv4 bakes this step's KV write target (c4/c128) into metadata,
# so slice the shared multi-step out_cache_loc now rather than at
# forward time.
# so slice the shared multi-step out_cache_loc now, not at forward time.
out_cache_loc = forward_batch.out_cache_loc
if self.topk > 0 and self.speculative_num_steps > 1:
out_cache_loc = per_step_draft_out_cache_loc(
@@ -734,154 +861,24 @@ class DeepseekV4AttnBackend(
raise NotImplementedError(f"unsupported mode {forward_batch.forward_mode=}")
self.forward_metadata = metadata
self.init_forward_metadata_in_graph(forward_batch)
def init_cuda_graph_state(self, max_bs: int, max_num_tokens: int) -> None:
self.cuda_graph_metadata_of_bucket_and_bs: Dict[
_GraphBucket,
Dict[
int,
Union[DSV4Metadata, DSV4RawDecodeMetadata, DSV4RawVerifyMetadata],
Union[
DSV4Metadata,
DSV4RawDecodeMetadata,
DSV4RawVerifyMetadata,
],
],
] = {bucket: {} for bucket in _GraphBucket}
self.draft_extend_num_tokens_per_bs = (
max_num_tokens // max_bs if max_bs > 0 else 1
)
def init_forward_metadata_capture_cuda_graph(
self,
bs: int,
num_tokens: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
encoder_lens: Optional[torch.Tensor],
forward_mode: ForwardMode,
spec_info: Optional[SpecInput],
) -> None:
from types import SimpleNamespace
assert req_pool_indices.size(0) == bs
assert seq_lens.size(0) == bs
bucket = _GraphBucket.of(forward_mode)
if bucket == _GraphBucket.DECODE_OR_IDLE:
dummy_cache_loc = torch.zeros_like(seq_lens)
elif bucket == _GraphBucket.TARGET_VERIFY:
dummy_cache_loc = torch.zeros(num_tokens, **self.cuda_int32_kwargs)
else:
dummy_cache_loc = None
self._replay_forward_batch = SimpleNamespace(
out_cache_loc=dummy_cache_loc,
forward_mode=forward_mode,
)
self.init_forward_metadata_replay_cuda_graph(
bs=bs,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
seq_lens_sum=int(seq_lens.sum().item()),
encoder_lens=encoder_lens,
forward_mode=forward_mode,
spec_info=spec_info,
seq_lens_cpu=seq_lens.cpu(),
)
# Preserve _current_capture_raw for on_after_cuda_graph_warmup
metadata = self.forward_metadata
self._current_capture_raw = (
metadata
if isinstance(metadata, (DSV4RawDecodeMetadata, DSV4RawVerifyMetadata))
else None
)
def init_forward_metadata_replay_cuda_graph(
self,
bs: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
seq_lens_sum: int,
encoder_lens: Optional[torch.Tensor],
forward_mode: ForwardMode,
spec_info: Optional[SpecInput],
seq_lens_cpu: Optional[torch.Tensor],
) -> None:
bucket = _GraphBucket.of(forward_mode)
# FIXME: see cuda_graph_runner — this attribute is set out-of-band.
fb = self._replay_forward_batch
out_cache_loc = fb.out_cache_loc
actual_forward_mode = fb.forward_mode
if actual_forward_mode == ForwardMode.IDLE:
logger.debug(
f"[IDLE replay] bs={bs}, "
f"local_seq_lens_len={len(seq_lens)}, "
f"has_graph={bs in self.cuda_graph_metadata_of_bucket_and_bs[_GraphBucket.DECODE_OR_IDLE]}"
)
device = seq_lens.device
seq_lens = torch.ones(bs, dtype=seq_lens.dtype, device=device)
seq_lens_cpu = torch.ones(bs, dtype=torch.int64)
seq_lens_sum = bs
req_pool_indices = torch.zeros(
bs, dtype=req_pool_indices.dtype, device=device
)
out_cache_loc = torch.zeros(bs, dtype=torch.int64, device=device)
assert seq_lens_cpu is not None
seq_lens = seq_lens[:bs]
seq_lens_cpu = seq_lens_cpu[:bs]
req_pool_indices = req_pool_indices[:bs]
actual_max_seq_len = seq_lens_cpu.max().item()
chosen_max_seq_len = self.MAX_SEQ_LEN_FOR_CAPTURE
assert actual_max_seq_len <= chosen_max_seq_len
if bucket == _GraphBucket.DECODE_OR_IDLE:
assert out_cache_loc is not None
assert len(out_cache_loc.shape) == 1, f"{out_cache_loc.shape=}"
out_cache_loc_padded = torch.nn.functional.pad(
out_cache_loc,
pad=(0, bs - len(out_cache_loc)),
mode="constant",
value=0,
)
temp_metadata = self.init_forward_metadata_decode(
max_seq_len=chosen_max_seq_len,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
out_cache_loc=out_cache_loc_padded,
)
elif bucket == _GraphBucket.TARGET_VERIFY:
assert out_cache_loc is not None
num_tokens = self.speculative_num_draft_tokens * bs
out_cache_loc_padded = torch.nn.functional.pad(
out_cache_loc,
pad=(0, num_tokens - len(out_cache_loc)),
mode="constant",
value=0,
)
temp_metadata = self.init_forward_metadata_target_verify(
max_seq_len=chosen_max_seq_len,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
out_cache_loc=out_cache_loc_padded,
use_prefill_cuda_graph=True,
)
elif bucket == _GraphBucket.DRAFT_EXTEND:
num_tokens_per_bs = self.draft_extend_num_tokens_per_bs
temp_metadata = self.init_forward_metadata_draft_extend(
max_seq_len=chosen_max_seq_len,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
seq_lens_cpu=seq_lens_cpu.tolist(),
num_tokens_per_bs=num_tokens_per_bs,
use_prefill_cuda_graph=True,
)
else:
raise NotImplementedError
self.replay_cuda_graph_metadata_from(
bs=bs, temp_metadata=temp_metadata, bucket=bucket
)
def replay_cuda_graph_metadata_from(
self,
bs: int,
@@ -938,24 +935,6 @@ class DeepseekV4AttnBackend(
cache_nope_fp8_rope_bf16_pack=swa_k_pack,
)
def _maybe_upgrade_forward_metadata(self) -> None:
# With SGLANG_PREP_IN_CUDA_GRAPH=1, init_forward_metadata_*
# returns a Raw metadata that only carries a few tensors. The
# full DSV4Metadata (including c4/c128 compress + core_attn +
# indexer metadata) must be materialized before any caller that
# touches those fields. For 1.6T the first two layers have
# compress_ratio=128, so forward_core_compressor / forward_c4_indexer
# can fire before attn_backend.forward(), and must trigger the
# upgrade themselves.
if isinstance(self.forward_metadata, DSV4RawVerifyMetadata):
self.forward_metadata = self.make_forward_metadata_from_raw_verify(
raw_metadata=self.forward_metadata,
)
elif isinstance(self.forward_metadata, DSV4RawDecodeMetadata):
self.forward_metadata = self.make_forward_metadata_from_raw_decode(
raw_metadata=self.forward_metadata,
)
def forward(
self,
q: torch.Tensor,
@@ -968,8 +947,6 @@ class DeepseekV4AttnBackend(
attn_sink: Optional[torch.Tensor] = None,
**_,
) -> torch.Tensor:
self._maybe_upgrade_forward_metadata()
if self.mtp_enabled and forward_batch.forward_mode.is_idle():
return q.new_empty(q.shape[0], q.shape[1], layer.v_head_dim)
@@ -1233,6 +1210,52 @@ class DeepseekV4MultiStepBackend(DeepseekV4AttnBackend):
)
)
def init_forward_metadata_in_graph(self, forward_batch: ForwardBatch) -> None:
for attn_backend in self.attn_backends:
attn_backend.init_forward_metadata_in_graph(forward_batch)
def init_forward_metadata_out_graph(
self,
forward_batch: ForwardBatch,
in_capture: bool = False,
):
from types import SimpleNamespace
inner_fb = SimpleNamespace(
batch_size=forward_batch.batch_size,
forward_mode=ForwardMode.DECODE,
# Propagate the real runtime mode so inner backends can detect IDLE
# and apply their idle substitution.
actual_forward_mode=getattr(
forward_batch, "actual_forward_mode", forward_batch.forward_mode
),
input_ids=getattr(forward_batch, "input_ids", None),
positions=getattr(forward_batch, "positions", None),
req_pool_indices=forward_batch.req_pool_indices,
seq_lens=forward_batch.seq_lens,
seq_lens_sum=forward_batch.seq_lens_sum,
seq_lens_cpu=forward_batch.seq_lens_cpu,
encoder_lens=None,
out_cache_loc=getattr(forward_batch, "out_cache_loc", None),
spec_info=forward_batch.spec_info,
)
if in_capture:
for i in range(self.speculative_num_steps):
self.attn_backends[i].init_forward_metadata_out_graph(
inner_fb, in_capture=True
)
else:
if self.speculative_num_steps == 1:
return
self.attn_backends[0].init_forward_metadata_out_graph(inner_fb)
temp_metadata = self.attn_backends[0].forward_metadata
for i in range(1, self.speculative_num_steps - 1):
self.attn_backends[i].replay_cuda_graph_metadata_from(
bs=forward_batch.batch_size,
temp_metadata=temp_metadata,
bucket=_GraphBucket.DECODE_OR_IDLE,
)
def init_forward_metadata(self, forward_batch: ForwardBatch):
for i in range(self.speculative_num_steps - 1):
self.attn_backends[i].init_forward_metadata(forward_batch)
@@ -1241,49 +1264,10 @@ class DeepseekV4MultiStepBackend(DeepseekV4AttnBackend):
for i in range(self.speculative_num_steps):
self.attn_backends[i].init_cuda_graph_state(max_bs, max_num_tokens)
def init_forward_metadata_capture_cuda_graph(self, forward_batch: ForwardBatch):
for i in range(self.speculative_num_steps):
self.attn_backends[i].init_forward_metadata_capture_cuda_graph(
forward_batch.batch_size,
forward_batch.batch_size * self.topk,
forward_batch.req_pool_indices,
forward_batch.seq_lens,
encoder_lens=None,
forward_mode=ForwardMode.DECODE,
spec_info=forward_batch.spec_info,
)
def on_after_cuda_graph_warmup(self):
for backend in self.attn_backends:
backend.on_after_cuda_graph_warmup()
def init_forward_metadata_replay_cuda_graph(
self, forward_batch: ForwardBatch, bs: int
):
if self.speculative_num_steps == 1:
return
self.attn_backends[0]._replay_forward_batch = forward_batch
self.attn_backends[0].init_forward_metadata_replay_cuda_graph(
bs=bs,
req_pool_indices=forward_batch.req_pool_indices,
seq_lens=forward_batch.seq_lens,
seq_lens_sum=forward_batch.seq_lens_sum,
encoder_lens=None,
forward_mode=ForwardMode.DECODE,
spec_info=forward_batch.spec_info,
seq_lens_cpu=forward_batch.seq_lens_cpu,
)
self.attn_backends[0]._replay_forward_batch = None
temp_metadata = self.attn_backends[0].forward_metadata
for i in range(1, self.speculative_num_steps - 1):
self.attn_backends[i].replay_cuda_graph_metadata_from(
bs=bs,
temp_metadata=temp_metadata,
bucket=_GraphBucket.DECODE_OR_IDLE,
)
def _pad_tensor_to_size(tensor: torch.Tensor, size: int, *, value: int = 0):
if value == 0:
@@ -52,7 +52,6 @@ from sglang.srt.layers.dp_attention import (
from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.speculative.eagle_utils import per_step_draft_out_cache_loc
from sglang.srt.speculative.spec_info import SpecInput
from sglang.srt.utils import ceil_align
if TYPE_CHECKING:
@@ -381,7 +380,6 @@ class DeepseekV4HipRadixBackend(
DSV4RawVerifyMetadata,
DSV4RawDecodeMetadata,
] = None
self._replay_forward_batch: Optional[ForwardBatch] = None # FIXME: out-of-band
def _move_to_device(self, x: List[int]) -> torch.Tensor:
pin_tensor = torch.tensor(x, dtype=torch.int32, pin_memory=True)
@@ -661,6 +659,133 @@ class DeepseekV4HipRadixBackend(
use_prefill_cuda_graph=use_prefill_cuda_graph,
)
def init_forward_metadata_in_graph(self, forward_batch: ForwardBatch) -> None:
# Upgrade Raw->Full so the c4/c128 compress + core_attn + indexer
# materialization is recorded inside the cuda graph; a no-op (Full
# already) when PREP_IN_CUDA_GRAPH=0.
if isinstance(self.forward_metadata, DSV4RawVerifyMetadata):
self.forward_metadata = self.make_forward_metadata_from_raw_verify(
raw_metadata=self.forward_metadata,
)
elif isinstance(self.forward_metadata, DSV4RawDecodeMetadata):
self.forward_metadata = self.make_forward_metadata_from_raw_decode(
raw_metadata=self.forward_metadata,
)
def init_forward_metadata_out_graph(
self,
forward_batch: ForwardBatch,
in_capture: bool = False,
) -> None:
bucket = _GraphBucket.of(forward_batch.forward_mode)
bs = forward_batch.batch_size
req_pool_indices = forward_batch.req_pool_indices
seq_lens = forward_batch.seq_lens
if in_capture:
assert req_pool_indices.size(0) == bs
assert seq_lens.size(0) == bs
num_tokens = forward_batch.positions.numel()
if bucket == _GraphBucket.DECODE_OR_IDLE:
out_cache_loc = torch.zeros_like(seq_lens)
elif bucket == _GraphBucket.TARGET_VERIFY:
out_cache_loc = torch.zeros(num_tokens, **self.cuda_int32_kwargs)
else:
out_cache_loc = None
actual_forward_mode = forward_batch.forward_mode
seq_lens_sum = int(seq_lens.sum().item())
seq_lens_cpu = seq_lens.cpu()
else:
out_cache_loc = forward_batch.out_cache_loc
actual_forward_mode = getattr(
forward_batch, "actual_forward_mode", forward_batch.forward_mode
)
seq_lens_sum = forward_batch.seq_lens_sum
seq_lens_cpu = forward_batch.seq_lens_cpu
if actual_forward_mode == ForwardMode.IDLE:
logger.debug(
f"[IDLE replay] bs={bs}, "
f"local_seq_lens_len={len(seq_lens)}, "
f"has_graph={bs in self.cuda_graph_metadata_of_bucket_and_bs[_GraphBucket.DECODE_OR_IDLE]}"
)
device = seq_lens.device
seq_lens = torch.ones(bs, dtype=seq_lens.dtype, device=device)
seq_lens_cpu = torch.ones(bs, dtype=torch.int64)
seq_lens_sum = bs
req_pool_indices = torch.zeros(
bs, dtype=req_pool_indices.dtype, device=device
)
out_cache_loc = torch.zeros(bs, dtype=torch.int64, device=device)
assert seq_lens_cpu is not None
seq_lens = seq_lens[:bs]
seq_lens_cpu = seq_lens_cpu[:bs]
req_pool_indices = req_pool_indices[:bs]
actual_max_seq_len = seq_lens_cpu.max().item()
chosen_max_seq_len = self.MAX_SEQ_LEN_FOR_CAPTURE
assert actual_max_seq_len <= chosen_max_seq_len
if bucket == _GraphBucket.DECODE_OR_IDLE:
assert out_cache_loc is not None
assert len(out_cache_loc.shape) == 1, f"{out_cache_loc.shape=}"
out_cache_loc_padded = torch.nn.functional.pad(
out_cache_loc,
pad=(0, bs - len(out_cache_loc)),
mode="constant",
value=0,
)
temp_metadata = self.init_forward_metadata_decode(
max_seq_len=chosen_max_seq_len,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
out_cache_loc=out_cache_loc_padded,
)
elif bucket == _GraphBucket.TARGET_VERIFY:
assert out_cache_loc is not None
num_tokens_v = self.speculative_num_draft_tokens * bs
out_cache_loc_padded = torch.nn.functional.pad(
out_cache_loc,
pad=(0, num_tokens_v - len(out_cache_loc)),
mode="constant",
value=0,
)
temp_metadata = self.init_forward_metadata_target_verify(
max_seq_len=chosen_max_seq_len,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
out_cache_loc=out_cache_loc_padded,
use_prefill_cuda_graph=True,
)
elif bucket == _GraphBucket.DRAFT_EXTEND:
num_tokens_per_bs = self.draft_extend_num_tokens_per_bs
temp_metadata = self.init_forward_metadata_draft_extend(
max_seq_len=chosen_max_seq_len,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
seq_lens_cpu=seq_lens_cpu.tolist(),
num_tokens_per_bs=num_tokens_per_bs,
use_prefill_cuda_graph=True,
)
else:
raise NotImplementedError
self.replay_cuda_graph_metadata_from(
bs=bs, temp_metadata=temp_metadata, bucket=bucket
)
if in_capture:
metadata = self.forward_metadata
self._current_capture_raw = (
metadata
if isinstance(
metadata,
(DSV4RawDecodeMetadata, DSV4RawVerifyMetadata),
)
else None
)
def init_forward_metadata(self, forward_batch: ForwardBatch) -> None:
if self.mtp_enabled and forward_batch.forward_mode.is_idle():
return
@@ -676,8 +801,7 @@ class DeepseekV4HipRadixBackend(
if forward_batch.forward_mode.is_decode_or_idle():
# DSv4 bakes this step's KV write target (c4/c128) into metadata,
# so slice the shared multi-step out_cache_loc now rather than at
# forward time.
# so slice the shared multi-step out_cache_loc now, not at forward time.
out_cache_loc = forward_batch.out_cache_loc
if self.topk > 0 and self.speculative_num_steps > 1:
out_cache_loc = per_step_draft_out_cache_loc(
@@ -725,154 +849,24 @@ class DeepseekV4HipRadixBackend(
raise NotImplementedError(f"unsupported mode {forward_batch.forward_mode=}")
self.forward_metadata = metadata
self.init_forward_metadata_in_graph(forward_batch)
def init_cuda_graph_state(self, max_bs: int, max_num_tokens: int) -> None:
self.cuda_graph_metadata_of_bucket_and_bs: Dict[
_GraphBucket,
Dict[
int,
Union[DSV4Metadata, DSV4RawDecodeMetadata, DSV4RawVerifyMetadata],
Union[
DSV4Metadata,
DSV4RawDecodeMetadata,
DSV4RawVerifyMetadata,
],
],
] = {bucket: {} for bucket in _GraphBucket}
self.draft_extend_num_tokens_per_bs = (
max_num_tokens // max_bs if max_bs > 0 else 1
)
def init_forward_metadata_capture_cuda_graph(
self,
bs: int,
num_tokens: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
encoder_lens: Optional[torch.Tensor],
forward_mode: ForwardMode,
spec_info: Optional[SpecInput],
) -> None:
from types import SimpleNamespace
assert req_pool_indices.size(0) == bs
assert seq_lens.size(0) == bs
bucket = _GraphBucket.of(forward_mode)
if bucket == _GraphBucket.DECODE_OR_IDLE:
dummy_cache_loc = torch.zeros_like(seq_lens)
elif bucket == _GraphBucket.TARGET_VERIFY:
dummy_cache_loc = torch.zeros(num_tokens, **self.cuda_int32_kwargs)
else:
dummy_cache_loc = None
self._replay_forward_batch = SimpleNamespace(
out_cache_loc=dummy_cache_loc,
forward_mode=forward_mode,
)
self.init_forward_metadata_replay_cuda_graph(
bs=bs,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
seq_lens_sum=int(seq_lens.sum().item()),
encoder_lens=encoder_lens,
forward_mode=forward_mode,
spec_info=spec_info,
seq_lens_cpu=seq_lens.cpu(),
)
# Preserve _current_capture_raw for on_after_cuda_graph_warmup
metadata = self.forward_metadata
self._current_capture_raw = (
metadata
if isinstance(metadata, (DSV4RawDecodeMetadata, DSV4RawVerifyMetadata))
else None
)
def init_forward_metadata_replay_cuda_graph(
self,
bs: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
seq_lens_sum: int,
encoder_lens: Optional[torch.Tensor],
forward_mode: ForwardMode,
spec_info: Optional[SpecInput],
seq_lens_cpu: Optional[torch.Tensor],
) -> None:
bucket = _GraphBucket.of(forward_mode)
# FIXME: see cuda_graph_runner — this attribute is set out-of-band.
fb = self._replay_forward_batch
out_cache_loc = fb.out_cache_loc
actual_forward_mode = fb.forward_mode
if actual_forward_mode == ForwardMode.IDLE:
logger.debug(
f"[IDLE replay] bs={bs}, "
f"local_seq_lens_len={len(seq_lens)}, "
f"has_graph={bs in self.cuda_graph_metadata_of_bucket_and_bs[_GraphBucket.DECODE_OR_IDLE]}"
)
device = seq_lens.device
seq_lens = torch.ones(bs, dtype=seq_lens.dtype, device=device)
seq_lens_cpu = torch.ones(bs, dtype=torch.int64)
seq_lens_sum = bs
req_pool_indices = torch.zeros(
bs, dtype=req_pool_indices.dtype, device=device
)
out_cache_loc = torch.zeros(bs, dtype=torch.int64, device=device)
assert seq_lens_cpu is not None
seq_lens = seq_lens[:bs]
seq_lens_cpu = seq_lens_cpu[:bs]
req_pool_indices = req_pool_indices[:bs]
actual_max_seq_len = seq_lens_cpu.max().item()
chosen_max_seq_len = self.MAX_SEQ_LEN_FOR_CAPTURE
assert actual_max_seq_len <= chosen_max_seq_len
if bucket == _GraphBucket.DECODE_OR_IDLE:
assert out_cache_loc is not None
assert len(out_cache_loc.shape) == 1, f"{out_cache_loc.shape=}"
out_cache_loc_padded = torch.nn.functional.pad(
out_cache_loc,
pad=(0, bs - len(out_cache_loc)),
mode="constant",
value=0,
)
temp_metadata = self.init_forward_metadata_decode(
max_seq_len=chosen_max_seq_len,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
out_cache_loc=out_cache_loc_padded,
)
elif bucket == _GraphBucket.TARGET_VERIFY:
assert out_cache_loc is not None
num_tokens = self.speculative_num_draft_tokens * bs
out_cache_loc_padded = torch.nn.functional.pad(
out_cache_loc,
pad=(0, num_tokens - len(out_cache_loc)),
mode="constant",
value=0,
)
temp_metadata = self.init_forward_metadata_target_verify(
max_seq_len=chosen_max_seq_len,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
out_cache_loc=out_cache_loc_padded,
use_prefill_cuda_graph=True,
)
elif bucket == _GraphBucket.DRAFT_EXTEND:
num_tokens_per_bs = self.draft_extend_num_tokens_per_bs
temp_metadata = self.init_forward_metadata_draft_extend(
max_seq_len=chosen_max_seq_len,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
seq_lens_cpu=seq_lens_cpu.tolist(),
num_tokens_per_bs=num_tokens_per_bs,
use_prefill_cuda_graph=True,
)
else:
raise NotImplementedError
self.replay_cuda_graph_metadata_from(
bs=bs, temp_metadata=temp_metadata, bucket=bucket
)
def replay_cuda_graph_metadata_from(
self,
bs: int,
@@ -929,24 +923,6 @@ class DeepseekV4HipRadixBackend(
cache_nope_fp8_rope_bf16_pack=swa_k_pack,
)
def _maybe_upgrade_forward_metadata(self) -> None:
# With SGLANG_PREP_IN_CUDA_GRAPH=1, init_forward_metadata_*
# returns a Raw metadata that only carries a few tensors. The
# full DSV4Metadata (including c4/c128 compress + core_attn +
# indexer metadata) must be materialized before any caller that
# touches those fields. For 1.6T the first two layers have
# compress_ratio=128, so forward_core_compressor / forward_c4_indexer
# can fire before attn_backend.forward(), and must trigger the
# upgrade themselves.
if isinstance(self.forward_metadata, DSV4RawVerifyMetadata):
self.forward_metadata = self.make_forward_metadata_from_raw_verify(
raw_metadata=self.forward_metadata,
)
elif isinstance(self.forward_metadata, DSV4RawDecodeMetadata):
self.forward_metadata = self.make_forward_metadata_from_raw_decode(
raw_metadata=self.forward_metadata,
)
def forward(
self,
q: torch.Tensor,
@@ -959,8 +935,6 @@ class DeepseekV4HipRadixBackend(
attn_sink: Optional[torch.Tensor] = None,
**_,
) -> torch.Tensor:
self._maybe_upgrade_forward_metadata()
if self.mtp_enabled and forward_batch.forward_mode.is_idle():
return q.new_empty(q.shape[0], q.shape[1], layer.v_head_dim)
@@ -1212,6 +1186,52 @@ class DeepseekV4MultiStepBackend(DeepseekV4HipRadixBackend):
)
)
def init_forward_metadata_in_graph(self, forward_batch: ForwardBatch) -> None:
for attn_backend in self.attn_backends:
attn_backend.init_forward_metadata_in_graph(forward_batch)
def init_forward_metadata_out_graph(
self,
forward_batch: ForwardBatch,
in_capture: bool = False,
):
from types import SimpleNamespace
inner_fb = SimpleNamespace(
batch_size=forward_batch.batch_size,
forward_mode=ForwardMode.DECODE,
# Propagate the real runtime mode so inner backends can detect IDLE
# and apply their idle substitution.
actual_forward_mode=getattr(
forward_batch, "actual_forward_mode", forward_batch.forward_mode
),
input_ids=getattr(forward_batch, "input_ids", None),
positions=getattr(forward_batch, "positions", None),
req_pool_indices=forward_batch.req_pool_indices,
seq_lens=forward_batch.seq_lens,
seq_lens_sum=forward_batch.seq_lens_sum,
seq_lens_cpu=forward_batch.seq_lens_cpu,
encoder_lens=None,
out_cache_loc=getattr(forward_batch, "out_cache_loc", None),
spec_info=forward_batch.spec_info,
)
if in_capture:
for i in range(self.speculative_num_steps):
self.attn_backends[i].init_forward_metadata_out_graph(
inner_fb, in_capture=True
)
else:
if self.speculative_num_steps == 1:
return
self.attn_backends[0].init_forward_metadata_out_graph(inner_fb)
temp_metadata = self.attn_backends[0].forward_metadata
for i in range(1, self.speculative_num_steps - 1):
self.attn_backends[i].replay_cuda_graph_metadata_from(
bs=forward_batch.batch_size,
temp_metadata=temp_metadata,
bucket=_GraphBucket.DECODE_OR_IDLE,
)
def init_forward_metadata(self, forward_batch: ForwardBatch):
for i in range(self.speculative_num_steps - 1):
self.attn_backends[i].init_forward_metadata(forward_batch)
@@ -1220,49 +1240,10 @@ class DeepseekV4MultiStepBackend(DeepseekV4HipRadixBackend):
for i in range(self.speculative_num_steps):
self.attn_backends[i].init_cuda_graph_state(max_bs, max_num_tokens)
def init_forward_metadata_capture_cuda_graph(self, forward_batch: ForwardBatch):
for i in range(self.speculative_num_steps):
self.attn_backends[i].init_forward_metadata_capture_cuda_graph(
forward_batch.batch_size,
forward_batch.batch_size * self.topk,
forward_batch.req_pool_indices,
forward_batch.seq_lens,
encoder_lens=None,
forward_mode=ForwardMode.DECODE,
spec_info=forward_batch.spec_info,
)
def on_after_cuda_graph_warmup(self):
for backend in self.attn_backends:
backend.on_after_cuda_graph_warmup()
def init_forward_metadata_replay_cuda_graph(
self, forward_batch: ForwardBatch, bs: int
):
if self.speculative_num_steps == 1:
return
self.attn_backends[0]._replay_forward_batch = forward_batch
self.attn_backends[0].init_forward_metadata_replay_cuda_graph(
bs=bs,
req_pool_indices=forward_batch.req_pool_indices,
seq_lens=forward_batch.seq_lens,
seq_lens_sum=forward_batch.seq_lens_sum,
encoder_lens=None,
forward_mode=ForwardMode.DECODE,
spec_info=forward_batch.spec_info,
seq_lens_cpu=forward_batch.seq_lens_cpu,
)
self.attn_backends[0]._replay_forward_batch = None
temp_metadata = self.attn_backends[0].forward_metadata
for i in range(1, self.speculative_num_steps - 1):
self.attn_backends[i].replay_cuda_graph_metadata_from(
bs=bs,
temp_metadata=temp_metadata,
bucket=_GraphBucket.DECODE_OR_IDLE,
)
def _pad_tensor_to_size(tensor: torch.Tensor, size: int, *, value: int = 0):
if value == 0:
@@ -408,6 +408,25 @@ class DeepseekSparseAttnBackend(
)
return page_table[:, strided_indices] // page_size
def init_forward_metadata_out_graph(
self,
forward_batch: ForwardBatch,
in_capture: bool = False,
):
seq_lens_cpu = (
forward_batch.seq_lens.cpu() if in_capture else forward_batch.seq_lens_cpu
)
self._apply_cuda_graph_metadata(
bs=forward_batch.batch_size,
req_pool_indices=forward_batch.req_pool_indices,
seq_lens=forward_batch.seq_lens,
seq_lens_cpu=seq_lens_cpu,
forward_mode=forward_batch.forward_mode,
spec_info=forward_batch.spec_info,
out_cache_loc=getattr(forward_batch, "out_cache_loc", None),
actual_forward_mode=getattr(forward_batch, "actual_forward_mode", None),
)
def init_forward_metadata(self, forward_batch: ForwardBatch):
"""Init the metadata for a forward pass."""
batch_size = forward_batch.batch_size
@@ -971,42 +990,23 @@ class DeepseekSparseAttnBackend(
self.decode_cuda_graph_metadata[bs] = metadata
self.forward_metadata = metadata
def init_forward_metadata_capture_cuda_graph(
self,
bs: int,
num_tokens: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
encoder_lens: Optional[torch.Tensor],
forward_mode: ForwardMode,
spec_info: Optional[SpecInput],
):
"""Initialize forward metadata for capturing CUDA graph."""
self.init_forward_metadata_replay_cuda_graph(
bs=bs,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
seq_lens_sum=None,
encoder_lens=encoder_lens,
forward_mode=forward_mode,
spec_info=spec_info,
seq_lens_cpu=seq_lens.cpu(),
)
def init_forward_metadata_replay_cuda_graph(
def _apply_cuda_graph_metadata(
self,
bs: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
seq_lens_sum: int,
encoder_lens: Optional[torch.Tensor],
seq_lens_cpu: torch.Tensor,
forward_mode: ForwardMode,
spec_info: Optional[SpecInput],
seq_lens_cpu: Optional[torch.Tensor],
out_cache_loc: Optional[torch.Tensor] = None,
actual_forward_mode: Optional[ForwardMode] = None,
):
"""Initialize forward metadata for replaying CUDA graph."""
"""Shared capture+replay body for the cuda-graph init path.
Public entry: :py:meth:`init_forward_metadata_out_graph`. Spec runners
also call this directly via _apply_cuda_graph_metadata when they
need to pass out_cache_loc / actual_forward_mode explicitly.
"""
assert seq_lens_cpu is not None
if bs not in self.decode_cuda_graph_metadata:
@@ -2370,21 +2370,26 @@ class DeepseekSparseAttnMultiStepBackend:
for i in range(self.speculative_num_steps - 1):
self.attn_backends[i].init_cuda_graph_state(max_bs, max_num_tokens)
def init_forward_metadata_capture_cuda_graph(self, forward_batch: ForwardBatch):
for i in range(self.speculative_num_steps - 1):
self.attn_backends[i].init_forward_metadata_capture_cuda_graph(
forward_batch.batch_size,
forward_batch.batch_size * self.topk,
forward_batch.req_pool_indices,
forward_batch.seq_lens,
encoder_lens=None,
forward_mode=ForwardMode.DECODE,
spec_info=forward_batch.spec_info,
)
def init_forward_metadata_replay_cuda_graph(
self, forward_batch: ForwardBatch, bs: int
def init_forward_metadata_out_graph(
self,
forward_batch: ForwardBatch,
in_capture: bool = False,
):
from sglang.srt.model_executor.forward_batch_info import build_inner_fb_view
if in_capture:
inner_fb = build_inner_fb_view(
forward_batch,
bs=forward_batch.batch_size,
forward_mode=ForwardMode.DECODE,
)
for i in range(self.speculative_num_steps - 1):
self.attn_backends[i].init_forward_metadata_out_graph(
inner_fb, in_capture=True
)
return
bs = forward_batch.batch_size
if envs.SGLANG_DSA_ENABLE_MTP_PRECOMPUTE_METADATA.get():
# Precompute metadata once (shared across all backends)
precomputed = self.attn_backends[0]._precompute_replay_metadata(
@@ -2542,20 +2547,21 @@ class DeepseekSparseAttnMultiStepBackend:
forward_mode=ForwardMode.DECODE,
)
else:
# Fallback: compute metadata separately for each backend
for i in range(self.speculative_num_steps - 1):
self.attn_backends[i].init_forward_metadata_replay_cuda_graph(
self.attn_backends[i]._apply_cuda_graph_metadata(
bs=bs,
req_pool_indices=forward_batch.req_pool_indices,
seq_lens=forward_batch.seq_lens,
seq_lens_sum=forward_batch.seq_lens_sum,
encoder_lens=None,
seq_lens_cpu=forward_batch.seq_lens_cpu,
forward_mode=ForwardMode.DECODE,
spec_info=forward_batch.spec_info,
seq_lens_cpu=forward_batch.seq_lens_cpu,
out_cache_loc=None,
)
def init_forward_metadata_in_graph(self, forward_batch: ForwardBatch) -> None:
for i in range(self.speculative_num_steps - 1):
self.attn_backends[i].init_forward_metadata_in_graph(forward_batch)
# Backward-compat aliases (deprecated: use DSA class names)
DeepseekSparseAttnBackend = DeepseekSparseAttnBackend
@@ -57,9 +57,6 @@ class CompressorBackendMixin:
assert isinstance(metadata, FusedCompressMetadata)
return metadata
def _maybe_upgrade_forward_metadata(self) -> None:
pass
def forward_compress(
self,
*,
@@ -153,11 +150,6 @@ class CompressorBackendMixin:
) -> None:
if forward_batch.forward_mode.is_idle():
return
# PREP_IN_CG lazy upgrade: the concrete backend (DeepseekV4AttnBackend)
# owns this helper. MQALayer._forward_prepare calls us before
# attn_backend.forward(), so Raw -> DSV4Metadata must happen here too
# (e.g. 1.6T layer 0 has compress_ratio=128 and needs cX_compress_metadata).
self._maybe_upgrade_forward_metadata()
token_to_kv_pool = self.token_to_kv_pool
if TYPE_CHECKING:
assert isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool)
@@ -187,8 +179,6 @@ class CompressorBackendMixin:
compressor: Compressor,
) -> None:
assert is_overlap_compress(compressor.ratio)
# PREP_IN_CG lazy upgrade (see forward_core_compressor for rationale).
self._maybe_upgrade_forward_metadata()
token_to_kv_pool = self.token_to_kv_pool
if TYPE_CHECKING:
assert isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool)
@@ -404,9 +404,6 @@ class CompressorBackendMixin:
super().__init__()
self.forward_metadata: DSV4Metadata
# NOTE: Will be overridden
def _maybe_upgrade_forward_metadata(self): ...
def _get_paged_compress_metadata(self, compress_ratio: int) -> CompressMetadata:
attr_name = f"c{compress_ratio}_compress_metadata"
return getattr(self.forward_metadata, attr_name)
@@ -483,7 +480,6 @@ class CompressorBackendMixin:
if forward_batch.forward_mode.is_idle():
return
self._maybe_upgrade_forward_metadata()
token_to_kv_pool = self.token_to_kv_pool
token_to_kv_pool = cast("DeepSeekV4TokenToKVPool", token_to_kv_pool)
kv_score_input = compressor.compute_kv_score(x, forward_batch)
@@ -443,9 +443,6 @@ class C4IndexerBackendMixin:
) -> None:
if forward_batch.forward_mode.is_idle():
return
# PREP_IN_CG lazy upgrade: this runs from MQALayer._forward_prepare,
# before attn_backend.forward() would trigger the upgrade.
self._maybe_upgrade_forward_metadata()
token_to_kv_pool = self.token_to_kv_pool
if TYPE_CHECKING:
@@ -21,7 +21,9 @@ from sglang.jit_kernel.flash_attention import (
)
from sglang.srt.distributed.parallel_state import get_tensor_model_parallel_rank
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
from sglang.srt.layers.attention.flashattention_backend import FlashAttentionMetadata
from sglang.srt.layers.attention.flashattention_backend import (
FlashAttentionMetadata,
)
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
if TYPE_CHECKING:
@@ -172,6 +174,35 @@ class DualChunkFlashAttentionBackend(AttentionBackend):
end_head = start_head + self.num_heads
return [layer_sparse_attention_config[i] for i in range(start_head, end_head)]
def init_forward_metadata_out_graph(
self,
forward_batch: ForwardBatch,
in_capture: bool = False,
):
bs = forward_batch.batch_size
req_pool_indices = forward_batch.req_pool_indices
seq_lens = forward_batch.seq_lens
forward_mode = forward_batch.forward_mode
if in_capture:
self._bind_metadata_buffers(bs, req_pool_indices, forward_mode)
self._apply_cuda_graph_metadata(
bs=bs,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
forward_mode=forward_mode,
)
if in_capture and forward_mode.is_decode_or_idle():
# Restore max_seq_len scalars — replay sets actual values but CUDA
# graph needs the safe upper bound baked in at capture time.
md = self.forward_metadata
md.max_seq_len = self.max_context_len
md.max_seq_len_intra = self.max_context_len
md.max_seq_len_succ = self.max_context_len
md.max_seq_len_inter = self.max_context_len
def init_forward_metadata(self, forward_batch: ForwardBatch):
"""Initialize forward metadata hence all layers in the forward pass can reuse it."""
@@ -577,49 +608,17 @@ class DualChunkFlashAttentionBackend(AttentionBackend):
self.forward_metadata = metadata
def init_forward_metadata_capture_cuda_graph(
def _apply_cuda_graph_metadata(
self,
bs: int,
num_tokens: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
encoder_lens: Optional[torch.Tensor],
forward_mode: ForwardMode,
spec_info: Optional[None],
):
self._bind_metadata_buffers(bs, req_pool_indices, forward_mode)
self.init_forward_metadata_replay_cuda_graph(
bs=bs,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
seq_lens_sum=None,
encoder_lens=encoder_lens,
forward_mode=forward_mode,
spec_info=spec_info,
seq_lens_cpu=seq_lens.cpu(),
)
# Restore max_seq_len scalars — replay sets actual values but CUDA graph
# needs the safe upper bound baked in at capture time.
if forward_mode.is_decode_or_idle():
md = self.forward_metadata
md.max_seq_len = self.max_context_len
md.max_seq_len_intra = self.max_context_len
md.max_seq_len_succ = self.max_context_len
md.max_seq_len_inter = self.max_context_len
"""Shared capture+replay body for the cuda-graph init path.
def init_forward_metadata_replay_cuda_graph(
self,
bs: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
seq_lens_sum: int,
encoder_lens: Optional[torch.Tensor],
forward_mode: ForwardMode,
spec_info: Optional[None],
seq_lens_cpu: Optional[torch.Tensor],
out_cache_loc: torch.Tensor = None,
):
"""Initialize forward metadata for replaying CUDA graph."""
Public entry: :py:meth:`init_forward_metadata_out_graph`.
"""
assert forward_mode.is_decode()
seq_lens = seq_lens[:bs]
req_pool_indices = req_pool_indices[:bs]
@@ -273,6 +273,102 @@ class FlashAttentionBackend(AttentionBackend):
num_splits=self.num_splits,
)
def init_forward_metadata_out_graph(
self,
forward_batch: ForwardBatch,
in_capture: bool = False,
):
bs = forward_batch.batch_size
req_pool_indices = forward_batch.req_pool_indices
seq_lens = forward_batch.seq_lens
encoder_lens = forward_batch.encoder_lens
forward_mode = forward_batch.forward_mode
spec_info = forward_batch.spec_info
out_cache_loc = getattr(forward_batch, "out_cache_loc", None)
if in_capture:
num_tokens = forward_batch.positions.numel()
seq_lens_cpu = seq_lens.cpu()
self._bind_metadata_buffers(
bs,
num_tokens,
encoder_lens,
forward_mode,
spec_info,
seq_lens.device,
)
if (
forward_mode.is_decode_or_idle()
and spec_info is not None
and self.topk > 1
):
# topk>1 draft decode: replay needs out_cache_loc which capture doesn't have;
# set forward_metadata directly and let actual CUDA graph replay fill data.
self.forward_metadata = self.draft_decode_metadata_topk_normal[bs]
self.forward_metadata_spec_decode_expand = (
self.draft_decode_metadata_topk_expand[bs]
)
return
if forward_mode.is_target_verify() and self.topk > 1:
# topk>1 target verify: replay needs spec_info.positions and .custom_mask
# which are not populated at capture time.
self.forward_metadata = self.target_verify_metadata_topk_normal[bs]
self.forward_metadata_spec_decode_expand = (
self.target_verify_metadata_topk_expand[bs]
)
return
self._apply_cuda_graph_metadata(
bs=bs,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
seq_lens_sum=None,
encoder_lens=encoder_lens,
forward_mode=forward_mode,
spec_info=spec_info,
seq_lens_cpu=seq_lens_cpu,
out_cache_loc=out_cache_loc,
)
if forward_mode.is_decode_or_idle() and spec_info is None:
# Local attention and scheduler metadata require capture-time slice sizing.
# Both depend on data already filled by replay above.
metadata = self.decode_cuda_graph_metadata[bs]
self._maybe_update_local_attn_metadata_for_capture(metadata, bs)
if self._sched_meta_buf is not None:
sched = self._compute_scheduler_metadata(
bs,
max(metadata.max_seq_len_k, 1),
metadata.cache_seqlens_int32,
metadata.cu_seqlens_q,
)
if sched is not None:
n = sched.shape[0]
self._sched_meta_buf[:n] = sched
self._sched_meta_buf[n:] = 0
metadata.scheduler_metadata = self._sched_meta_buf[:n]
if forward_mode.is_draft_extend(include_v2=True):
# CUDA graph bakes max_seq_len_q as a constant. replay() sets it to
# max(num_accept_tokens_cpu) which is None/empty at capture time,
# falling back to 1. Restore the correct upper bound so the kernel
# sees num_tokens_per_bs (not 1) for all replays of this graph.
self.forward_metadata.max_seq_len_q = num_tokens // bs
else:
self._apply_cuda_graph_metadata(
bs=bs,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
seq_lens_sum=forward_batch.seq_lens_sum,
encoder_lens=encoder_lens,
forward_mode=forward_mode,
spec_info=spec_info,
seq_lens_cpu=forward_batch.seq_lens_cpu,
out_cache_loc=out_cache_loc,
)
def init_forward_metadata(self, forward_batch: ForwardBatch):
"""Initialize forward metadata hence all layers in the forward pass can reuse it."""
metadata = FlashAttentionMetadata()
@@ -1901,77 +1997,7 @@ class FlashAttentionBackend(AttentionBackend):
return metadata, metadata_expand
def init_forward_metadata_capture_cuda_graph(
self,
bs: int,
num_tokens: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
encoder_lens: Optional[torch.Tensor],
forward_mode: ForwardMode,
spec_info: Optional[SpecInput],
):
"""Initialize forward metadata for capturing CUDA graph."""
seq_lens_cpu = seq_lens.cpu()
self._bind_metadata_buffers(
bs, num_tokens, encoder_lens, forward_mode, spec_info, seq_lens.device
)
if forward_mode.is_decode_or_idle() and spec_info is not None and self.topk > 1:
# topk>1 draft decode: replay needs out_cache_loc which capture doesn't have;
# set forward_metadata directly and let actual CUDA graph replay fill data.
self.forward_metadata = self.draft_decode_metadata_topk_normal[bs]
self.forward_metadata_spec_decode_expand = (
self.draft_decode_metadata_topk_expand[bs]
)
return
if forward_mode.is_target_verify() and self.topk > 1:
# topk>1 target verify: replay needs spec_info.positions and .custom_mask
# which are not populated at capture time.
self.forward_metadata = self.target_verify_metadata_topk_normal[bs]
self.forward_metadata_spec_decode_expand = (
self.target_verify_metadata_topk_expand[bs]
)
return
self.init_forward_metadata_replay_cuda_graph(
bs=bs,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
seq_lens_sum=None,
encoder_lens=encoder_lens,
forward_mode=forward_mode,
spec_info=spec_info,
seq_lens_cpu=seq_lens_cpu,
)
if forward_mode.is_decode_or_idle() and spec_info is None:
# Local attention and scheduler metadata require capture-time slice sizing.
# Both depend on data already filled by replay above.
metadata = self.decode_cuda_graph_metadata[bs]
self._maybe_update_local_attn_metadata_for_capture(metadata, bs)
if self._sched_meta_buf is not None:
sched = self._compute_scheduler_metadata(
bs,
max(metadata.max_seq_len_k, 1),
metadata.cache_seqlens_int32,
metadata.cu_seqlens_q,
)
if sched is not None:
n = sched.shape[0]
self._sched_meta_buf[:n] = sched
self._sched_meta_buf[n:] = 0
metadata.scheduler_metadata = self._sched_meta_buf[:n]
if forward_mode.is_draft_extend(include_v2=True):
# CUDA graph bakes max_seq_len_q as a constant. replay() sets it to
# max(num_accept_tokens_cpu) which is None/empty at capture time,
# falling back to 1. Restore the correct upper bound so the kernel
# sees num_tokens_per_bs (not 1) for all replays of this graph.
self.forward_metadata.max_seq_len_q = num_tokens // bs
def init_forward_metadata_replay_cuda_graph(
def _apply_cuda_graph_metadata(
self,
bs: int,
req_pool_indices: torch.Tensor,
@@ -1983,7 +2009,13 @@ class FlashAttentionBackend(AttentionBackend):
seq_lens_cpu: Optional[torch.Tensor],
out_cache_loc: Optional[torch.Tensor] = None,
):
"""Initialize forward metadata for replaying CUDA graph."""
"""Shared capture+replay body for the cuda-graph init path.
Public entry: :py:meth:`init_forward_metadata_out_graph`. This helper
formerly lived as the legacy ``init_forward_metadata_replay_cuda_graph``;
the capture path used to wrap it. Both legacy method overrides
are gone.
"""
seq_lens = seq_lens[:bs]
seq_lens_cpu = seq_lens_cpu[:bs]
req_pool_indices = req_pool_indices[:bs]
@@ -2324,6 +2356,16 @@ class FlashAttentionBackend(AttentionBackend):
)
metadata.page_table[:, :max_seq_pages].copy_(page_indices // self.page_size)
else:
raise ValueError(
f"FA3 `_apply_cuda_graph_metadata` only supports the modes the "
f"full cuda-graph runner captures (decode / idle / target_verify "
f"/ draft_extend / draft_extend_v2). Got {forward_mode=}. "
f"Piecewise / breakable capture must route through "
f"`init_forward_metadata(fb)` (the eager entry) instead of "
f"`init_forward_metadata_out_graph(fb, in_capture=True)`."
)
if encoder_lens is not None:
# Per-request varlen encoder support (e.g. MossVL different images).
metadata.encoder_max_seq_len_k = int(encoder_lens.max().item())
@@ -2353,7 +2395,10 @@ class FlashAttentionBackend(AttentionBackend):
return 1
def _maybe_init_local_attn_metadata(
self, forwardbatch: ForwardBatch, metadata: FlashAttentionMetadata, device
self,
forwardbatch: ForwardBatch,
metadata: FlashAttentionMetadata,
device,
):
"""Centralized utility to initialize local_attn_metadata if chunked attention is enabled."""
if not self.has_local_attention:
@@ -2628,45 +2673,33 @@ class FlashAttentionMultiStepBackend:
for i in range(self.speculative_num_steps - 1):
self.attn_backends[i].init_cuda_graph_state(max_bs, max_num_tokens)
def init_forward_metadata_capture_cuda_graph(
def init_forward_metadata_out_graph(
self,
forward_batch: ForwardBatch,
in_capture: bool = False,
):
from sglang.srt.model_executor.forward_batch_info import build_inner_fb_view
assert forward_batch.spec_info is not None
assert forward_batch.spec_info.is_draft_input()
for i in range(self.speculative_num_steps - 1):
self.attn_backends[i].init_forward_metadata_capture_cuda_graph(
forward_batch.batch_size,
forward_batch.batch_size * self.topk,
forward_batch.req_pool_indices,
forward_batch.seq_lens,
encoder_lens=forward_batch.encoder_lens,
forward_mode=ForwardMode.DECODE,
spec_info=forward_batch.spec_info,
)
def init_forward_metadata_replay_cuda_graph(
self, forward_batch: ForwardBatch, bs: int
):
assert forward_batch.spec_info is not None
assert forward_batch.spec_info.is_draft_input()
inner_fb = build_inner_fb_view(
forward_batch,
bs=forward_batch.batch_size,
forward_mode=ForwardMode.DECODE,
encoder_lens=forward_batch.encoder_lens,
)
for i in range(self.speculative_num_steps - 1):
# TODO: incrementally update the metadata for the later steps,
# so that they do not need to recompute everything from scratch.
self.attn_backends[i].init_forward_metadata_replay_cuda_graph(
bs,
forward_batch.req_pool_indices,
forward_batch.seq_lens,
forward_batch.seq_lens_sum,
encoder_lens=forward_batch.encoder_lens,
forward_mode=ForwardMode.DECODE,
spec_info=forward_batch.spec_info,
seq_lens_cpu=forward_batch.seq_lens_cpu,
out_cache_loc=forward_batch.out_cache_loc,
self.attn_backends[i].init_forward_metadata_out_graph(
inner_fb, in_capture=in_capture
)
def init_forward_metadata_in_graph(self, forward_batch: ForwardBatch) -> None:
for i in range(self.speculative_num_steps - 1):
self.attn_backends[i].init_forward_metadata_in_graph(forward_batch)
@torch.compile(dynamic=True, backend=get_compiler_backend())
def draft_decode_set_expand_metadata(
@@ -461,6 +461,69 @@ class FlashInferAttnBackend(AttentionBackend):
),
)
def init_forward_metadata_out_graph(
self,
forward_batch: ForwardBatch,
in_capture: bool = False,
):
bs = forward_batch.batch_size
req_pool_indices = forward_batch.req_pool_indices
seq_lens = forward_batch.seq_lens
seq_lens_cpu = forward_batch.seq_lens_cpu
seq_lens_sum = forward_batch.seq_lens_sum
encoder_lens = forward_batch.encoder_lens
forward_mode = forward_batch.forward_mode
spec_info = forward_batch.spec_info
if in_capture:
num_tokens = forward_batch.positions.numel()
self._prepare_cuda_graph_metadata(bs, num_tokens, forward_mode, spec_info)
if forward_mode.is_decode_or_idle():
self.indices_updater_decode.update(
req_pool_indices[:bs],
seq_lens[:bs],
seq_lens_cpu[:bs] if seq_lens_cpu is not None else None,
seq_lens_sum,
decode_wrappers=self.decode_cuda_graph_metadata[bs],
encoder_lens=encoder_lens[:bs] if encoder_lens is not None else None,
spec_info=spec_info,
fixed_split_size=None,
disable_split_kv=self.disable_cuda_graph_kv_split,
)
elif forward_mode.is_target_verify() or forward_mode.is_draft_extend():
self.indices_updater_prefill.update(
req_pool_indices[:bs],
seq_lens[:bs],
seq_lens_cpu[:bs] if seq_lens_cpu is not None else None,
seq_lens_sum,
prefix_lens=None,
prefill_wrappers=self.prefill_cuda_graph_metadata[bs],
use_ragged=False,
encoder_lens=encoder_lens[:bs] if encoder_lens is not None else None,
spec_info=spec_info,
)
elif forward_mode.is_dllm_extend():
self.indices_updater_prefill.update(
req_pool_indices[:bs],
seq_lens[:bs],
seq_lens_cpu[:bs] if seq_lens_cpu is not None else None,
seq_lens_sum,
prefix_lens=seq_lens - self.dllm_config.block_size,
prefill_wrappers=self.prefill_cuda_graph_metadata[bs],
use_ragged=not self.use_paged,
encoder_lens=encoder_lens[:bs] if encoder_lens is not None else None,
spec_info=None,
)
else:
raise ValueError("Invalid forward mode")
if in_capture and forward_mode.is_decode_or_idle():
# fast_decode_plan needs _cached_module from the initial begin_forward
# above, so install it only after that first plan has run.
for w in self.decode_cuda_graph_metadata[bs]:
w.begin_forward = partial(fast_decode_plan, w)
def init_forward_metadata(self, forward_batch: ForwardBatch):
if forward_batch.forward_mode.is_decode_or_idle():
self.indices_updater_decode.update(
@@ -662,85 +725,6 @@ class FlashInferAttnBackend(AttentionBackend):
else:
raise ValueError(f"Invalid mode: {forward_mode=}")
def init_forward_metadata_capture_cuda_graph(
self,
bs: int,
num_tokens: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
encoder_lens: Optional[torch.Tensor],
forward_mode: ForwardMode,
spec_info: Optional[SpecInput],
):
seq_lens_sum = seq_lens.sum().item()
seq_lens_cpu = seq_lens.cpu()
self._prepare_cuda_graph_metadata(bs, num_tokens, forward_mode, spec_info)
self.init_forward_metadata_replay_cuda_graph(
bs=bs,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
seq_lens_sum=seq_lens_sum,
encoder_lens=encoder_lens,
forward_mode=forward_mode,
spec_info=spec_info,
seq_lens_cpu=seq_lens_cpu,
)
# fast_decode_plan requires _cached_module set by the initial full
# begin_forward call above; install it only after that first plan runs.
if forward_mode.is_decode_or_idle():
for w in self.decode_cuda_graph_metadata[bs]:
w.begin_forward = partial(fast_decode_plan, w)
def init_forward_metadata_replay_cuda_graph(
self,
bs: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
seq_lens_sum: int,
encoder_lens: Optional[torch.Tensor],
forward_mode: ForwardMode,
spec_info: Optional[SpecInput],
seq_lens_cpu: Optional[torch.Tensor],
):
if forward_mode.is_decode_or_idle():
self.indices_updater_decode.update(
req_pool_indices[:bs],
seq_lens[:bs],
seq_lens_cpu[:bs] if seq_lens_cpu is not None else None,
seq_lens_sum,
decode_wrappers=self.decode_cuda_graph_metadata[bs],
encoder_lens=encoder_lens[:bs] if encoder_lens is not None else None,
spec_info=spec_info,
fixed_split_size=None,
disable_split_kv=self.disable_cuda_graph_kv_split,
)
elif forward_mode.is_target_verify() or forward_mode.is_draft_extend():
self.indices_updater_prefill.update(
req_pool_indices[:bs],
seq_lens[:bs],
seq_lens_cpu[:bs] if seq_lens_cpu is not None else None,
seq_lens_sum,
prefix_lens=None,
prefill_wrappers=self.prefill_cuda_graph_metadata[bs],
use_ragged=False,
encoder_lens=encoder_lens[:bs] if encoder_lens is not None else None,
spec_info=spec_info,
)
elif forward_mode.is_dllm_extend():
self.indices_updater_prefill.update(
req_pool_indices[:bs],
seq_lens[:bs],
seq_lens_cpu[:bs] if seq_lens_cpu is not None else None,
seq_lens_sum,
prefix_lens=seq_lens - self.dllm_config.block_size,
prefill_wrappers=self.prefill_cuda_graph_metadata[bs],
use_ragged=not self.use_paged,
encoder_lens=encoder_lens[:bs] if encoder_lens is not None else None,
spec_info=None,
)
else:
raise ValueError("Invalid forward mode")
def get_cuda_graph_seq_len_fill_value(self):
return 1
@@ -1668,37 +1652,27 @@ class FlashInferMultiStepDraftBackend:
max_bs, max_num_tokens, kv_indices_buf=self.cuda_graph_kv_indices[i]
)
def init_forward_metadata_capture_cuda_graph(self, forward_batch: ForwardBatch):
def call_fn(i, forward_batch):
self.attn_backends[i].init_forward_metadata_capture_cuda_graph(
forward_batch.batch_size,
forward_batch.batch_size * self.topk,
forward_batch.req_pool_indices,
forward_batch.seq_lens,
encoder_lens=None,
forward_mode=ForwardMode.DECODE,
spec_info=forward_batch.spec_info,
)
self.common_template(forward_batch, self.cuda_graph_kv_indices, call_fn)
def init_forward_metadata_replay_cuda_graph(
self, forward_batch: ForwardBatch, bs: int
def init_forward_metadata_out_graph(
self,
forward_batch: ForwardBatch,
in_capture: bool = False,
):
def call_fn(i, forward_batch):
self.attn_backends[i].init_forward_metadata_replay_cuda_graph(
bs,
forward_batch.req_pool_indices,
forward_batch.seq_lens,
seq_lens_sum=-1,
encoder_lens=None,
forward_mode=ForwardMode.DECODE,
spec_info=forward_batch.spec_info,
seq_lens_cpu=forward_batch.seq_lens_cpu,
from sglang.srt.model_executor.forward_batch_info import build_inner_fb_view
bs = forward_batch.batch_size
def call_fn(i, fb):
inner_fb = build_inner_fb_view(fb, bs=bs, forward_mode=ForwardMode.DECODE)
self.attn_backends[i].init_forward_metadata_out_graph(
inner_fb, in_capture=in_capture
)
self.common_template(forward_batch, self.cuda_graph_kv_indices, call_fn)
def init_forward_metadata_in_graph(self, forward_batch: ForwardBatch) -> None:
for attn_backend in self.attn_backends:
attn_backend.init_forward_metadata_in_graph(forward_batch)
def should_use_tensor_core(
kv_cache_dtype: torch.dtype,
@@ -289,6 +289,80 @@ class FlashInferMLAAttnBackend(AttentionBackend):
self.decode_cuda_graph_metadata = {}
self.prefill_cuda_graph_metadata = {} # For verify
def init_forward_metadata_out_graph(
self,
forward_batch: ForwardBatch,
in_capture: bool = False,
):
bs = forward_batch.batch_size
req_pool_indices = forward_batch.req_pool_indices
seq_lens = forward_batch.seq_lens
forward_mode = forward_batch.forward_mode
spec_info = forward_batch.spec_info
if in_capture:
num_tokens = forward_batch.positions.numel()
seq_lens_sum = seq_lens.sum().item()
seq_lens_cpu = seq_lens.cpu()
if forward_mode.is_decode_or_idle():
decode_wrapper = BatchMLAPagedAttentionWrapper(
self.workspace_buffer,
use_cuda_graph=True,
qo_indptr=self.cuda_graph_qo_indptr[: num_tokens + 1],
kv_indptr=self.cuda_graph_kv_indptr[: num_tokens + 1],
kv_indices=self.cuda_graph_kv_indices,
kv_len_arr=self.cuda_graph_kv_lens[:num_tokens],
backend="auto",
)
self.indices_updater_decode.update(
req_pool_indices,
seq_lens,
seq_lens_sum,
decode_wrapper=decode_wrapper,
init_metadata_replay=False,
spec_info=spec_info,
)
self.decode_cuda_graph_metadata[bs] = decode_wrapper
self.forward_metadata = DecodeMetadata(decode_wrapper)
# fast_mla_decode_plan needs _cached_module from the initial
# begin_forward above, so install it only after that call completes.
decode_wrapper.plan = partial(fast_mla_decode_plan, decode_wrapper)
elif forward_mode.is_target_verify() or forward_mode.is_draft_extend():
prefill_wrapper = BatchMLAPagedAttentionWrapper(
self.workspace_buffer,
use_cuda_graph=True,
qo_indptr=self.cuda_graph_qo_indptr[: bs + 1],
kv_indptr=self.cuda_graph_kv_indptr[: bs + 1],
kv_indices=self.cuda_graph_kv_indices,
kv_len_arr=self.cuda_graph_kv_lens[:bs],
backend="auto",
)
self.prefill_cuda_graph_metadata[bs] = prefill_wrapper
self.forward_metadata = PrefillMetadata(prefill_wrapper, False)
else:
raise ValueError(f"Invalid mode: {forward_mode=}")
self._apply_cuda_graph_metadata(
bs=bs,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
seq_lens_sum=seq_lens_sum,
forward_mode=forward_mode,
spec_info=spec_info,
seq_lens_cpu=seq_lens_cpu,
)
else:
self._apply_cuda_graph_metadata(
bs=bs,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
seq_lens_sum=forward_batch.seq_lens_sum,
forward_mode=forward_mode,
spec_info=spec_info,
seq_lens_cpu=forward_batch.seq_lens_cpu,
)
def init_forward_metadata(self, forward_batch: ForwardBatch):
if forward_batch.forward_mode.is_decode_or_idle():
self.indices_updater_decode.update(
@@ -374,83 +448,20 @@ class FlashInferMLAAttnBackend(AttentionBackend):
"kv_indices": self.cuda_graph_kv_indices,
}
def init_forward_metadata_capture_cuda_graph(
self,
bs: int,
num_tokens: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
encoder_lens: Optional[torch.Tensor],
forward_mode: ForwardMode,
spec_info: Optional[SpecInput],
):
seq_lens_sum = seq_lens.sum().item()
seq_lens_cpu = seq_lens.cpu()
if forward_mode.is_decode_or_idle():
# Decode: create wrapper, run the initial full begin_forward (False),
# then install the fast plan. After that, call replay so the
# data-update path (update(True)) is also exercised during capture.
decode_wrapper = BatchMLAPagedAttentionWrapper(
self.workspace_buffer,
use_cuda_graph=True,
qo_indptr=self.cuda_graph_qo_indptr[: num_tokens + 1],
kv_indptr=self.cuda_graph_kv_indptr[: num_tokens + 1],
kv_indices=self.cuda_graph_kv_indices,
kv_len_arr=self.cuda_graph_kv_lens[:num_tokens],
backend="auto",
)
self.indices_updater_decode.update(
req_pool_indices,
seq_lens,
seq_lens_sum,
decode_wrapper=decode_wrapper,
init_metadata_replay=False,
spec_info=spec_info,
)
self.decode_cuda_graph_metadata[bs] = decode_wrapper
self.forward_metadata = DecodeMetadata(decode_wrapper)
# fast_mla_decode_plan requires _cached_module set by the initial
# begin_forward above; install it only after that call completes.
decode_wrapper.plan = partial(fast_mla_decode_plan, decode_wrapper)
elif forward_mode.is_target_verify() or forward_mode.is_draft_extend():
# Prefill: create wrapper and store — replay handles the update call.
prefill_wrapper = BatchMLAPagedAttentionWrapper(
self.workspace_buffer,
use_cuda_graph=True,
qo_indptr=self.cuda_graph_qo_indptr[: bs + 1],
kv_indptr=self.cuda_graph_kv_indptr[: bs + 1],
kv_indices=self.cuda_graph_kv_indices,
kv_len_arr=self.cuda_graph_kv_lens[:bs],
backend="auto",
)
self.prefill_cuda_graph_metadata[bs] = prefill_wrapper
self.forward_metadata = PrefillMetadata(prefill_wrapper, False)
else:
raise ValueError(f"Invalid mode: {forward_mode=}")
self.init_forward_metadata_replay_cuda_graph(
bs=bs,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
seq_lens_sum=seq_lens_sum,
encoder_lens=encoder_lens,
forward_mode=forward_mode,
spec_info=spec_info,
seq_lens_cpu=seq_lens_cpu,
)
def init_forward_metadata_replay_cuda_graph(
def _apply_cuda_graph_metadata(
self,
bs: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
seq_lens_sum: int,
encoder_lens: Optional[torch.Tensor],
forward_mode: ForwardMode,
spec_info: Optional[SpecInput],
seq_lens_cpu: Optional[torch.Tensor],
):
"""Shared capture+replay body for the cuda-graph init path.
Public entry: :py:meth:`init_forward_metadata_out_graph`.
"""
if forward_mode.is_decode_or_idle():
assert seq_lens_cpu is not None
kv_len_arr_cpu = seq_lens_cpu[:bs]
@@ -993,37 +1004,30 @@ class FlashInferMLAMultiStepDraftBackend:
max_bs, max_num_tokens, kv_indices_buf=self.cuda_graph_kv_indices[i]
)
def init_forward_metadata_capture_cuda_graph(self, forward_batch: ForwardBatch):
def call_fn(i, forward_batch):
self.attn_backends[i].init_forward_metadata_capture_cuda_graph(
forward_batch.batch_size,
forward_batch.batch_size * self.topk,
forward_batch.req_pool_indices,
forward_batch.seq_lens,
encoder_lens=None,
forward_mode=ForwardMode.DECODE,
spec_info=forward_batch.spec_info,
)
self.common_template(forward_batch, self.cuda_graph_kv_indices, call_fn)
def init_forward_metadata_replay_cuda_graph(
self, forward_batch: ForwardBatch, bs: int
def init_forward_metadata_out_graph(
self,
forward_batch: ForwardBatch,
in_capture: bool = False,
):
def call_fn(i, forward_batch):
self.attn_backends[i].init_forward_metadata_replay_cuda_graph(
bs,
forward_batch.req_pool_indices,
forward_batch.seq_lens,
seq_lens_sum=-1,
encoder_lens=None,
forward_mode=ForwardMode.DECODE,
spec_info=forward_batch.spec_info,
seq_lens_cpu=forward_batch.seq_lens_cpu,
from sglang.srt.model_executor.forward_batch_info import build_inner_fb_view
inner_fb = build_inner_fb_view(
forward_batch,
bs=forward_batch.batch_size,
forward_mode=ForwardMode.DECODE,
)
def call_fn(i, _forward_batch):
self.attn_backends[i].init_forward_metadata_out_graph(
inner_fb, in_capture=in_capture
)
self.common_template(forward_batch, self.cuda_graph_kv_indices, call_fn)
def init_forward_metadata_in_graph(self, forward_batch: ForwardBatch) -> None:
for attn_backend in self.attn_backends:
attn_backend.init_forward_metadata_in_graph(forward_batch)
def fast_mla_decode_plan(
self,
@@ -21,7 +21,6 @@ from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMo
if TYPE_CHECKING:
from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.model_executor.model_runner import ModelRunner
from sglang.srt.speculative.spec_info import SpecInput
logger = logging.getLogger(__name__)
@@ -86,6 +85,25 @@ class FlashMLABackend(FlashInferMLAAttnBackend):
self.cuda_graph_mla_metadata_view = None
self.cuda_graph_num_splits_view = None
def init_forward_metadata_out_graph(
self,
forward_batch: ForwardBatch,
in_capture: bool = False,
):
forward_mode = forward_batch.forward_mode
if forward_mode.is_decode_or_idle() or forward_mode.is_target_verify():
self._apply_decode_target_verify_metadata(
bs=forward_batch.batch_size,
req_pool_indices=forward_batch.req_pool_indices,
seq_lens=forward_batch.seq_lens,
seq_lens_cpu=forward_batch.seq_lens_cpu,
forward_mode=forward_mode,
)
else:
super().init_forward_metadata_out_graph(
forward_batch, in_capture=in_capture
)
def init_forward_metadata(self, forward_batch: ForwardBatch):
bs = forward_batch.batch_size
if forward_batch.forward_mode.is_decode_or_idle():
@@ -185,50 +203,21 @@ class FlashMLABackend(FlashInferMLAAttnBackend):
self.cuda_graph_mla_metadata_view = None
self.cuda_graph_num_splits_view = None
def init_forward_metadata_capture_cuda_graph(
self,
bs: int,
num_tokens: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
encoder_lens: Optional[torch.Tensor],
forward_mode: ForwardMode,
spec_info: Optional[SpecInput],
):
if forward_mode.is_decode_or_idle() or forward_mode.is_target_verify():
self.init_forward_metadata_replay_cuda_graph(
bs=bs,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
seq_lens_sum=None,
encoder_lens=encoder_lens,
forward_mode=forward_mode,
spec_info=spec_info,
seq_lens_cpu=None,
)
else:
super().init_forward_metadata_capture_cuda_graph(
bs,
num_tokens,
req_pool_indices,
seq_lens,
encoder_lens,
forward_mode,
spec_info,
)
def init_forward_metadata_replay_cuda_graph(
def _apply_decode_target_verify_metadata(
self,
bs: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
seq_lens_sum: int,
encoder_lens: Optional[torch.Tensor],
forward_mode: ForwardMode,
spec_info: Optional[SpecInput],
seq_lens_cpu: Optional[torch.Tensor],
forward_mode: ForwardMode,
):
if forward_mode.is_decode_or_idle() or forward_mode.is_target_verify():
"""Shared decode/target-verify capture+replay body.
Public entry: :py:meth:`init_forward_metadata_out_graph` (which routes
to this helper for decode/target-verify and falls back to the
FlashInferMLA parent for prefill/draft-extend).
"""
if True:
seq_lens = seq_lens[:bs]
seq_lens_cpu = seq_lens_cpu[:bs] if seq_lens_cpu is not None else None
@@ -295,17 +284,6 @@ class FlashMLABackend(FlashInferMLAAttnBackend):
self.cuda_graph_num_splits_view,
self.cuda_graph_kv_indices[:bs, :max_seqlen_pad],
)
else:
super().init_forward_metadata_replay_cuda_graph(
bs,
req_pool_indices,
seq_lens,
seq_lens_sum,
encoder_lens,
forward_mode,
spec_info,
seq_lens_cpu,
)
def get_cuda_graph_seq_len_fill_value(self):
return 1
@@ -516,39 +494,29 @@ class FlashMLAMultiStepDraftBackend:
max_bs, max_num_tokens, block_kv_indices=None
)
def init_forward_metadata_capture_cuda_graph(self, forward_batch: ForwardBatch):
def call_fn(i, forward_batch):
# EAGLE draft worker uses DECODE mode for draft steps
from sglang.srt.model_executor.forward_batch_info import ForwardMode
# Create a dummy forward_mode for draft step
self.attn_backends[i].init_forward_metadata_capture_cuda_graph(
forward_batch.batch_size,
forward_batch.batch_size * self.topk,
forward_batch.req_pool_indices,
forward_batch.seq_lens,
encoder_lens=None,
forward_mode=ForwardMode.DECODE,
spec_info=forward_batch.spec_info,
)
self.common_template(forward_batch, call_fn)
def init_forward_metadata_replay_cuda_graph(
self, forward_batch: ForwardBatch, bs: int
def init_forward_metadata_out_graph(
self,
forward_batch: ForwardBatch,
in_capture: bool = False,
):
def call_fn(i, forward_batch):
from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.srt.model_executor.forward_batch_info import (
ForwardMode,
build_inner_fb_view,
)
self.attn_backends[i].init_forward_metadata_replay_cuda_graph(
bs,
forward_batch.req_pool_indices,
forward_batch.seq_lens,
seq_lens_sum=-1,
encoder_lens=None,
forward_mode=ForwardMode.DECODE,
spec_info=forward_batch.spec_info,
seq_lens_cpu=forward_batch.seq_lens_cpu,
inner_fb = build_inner_fb_view(
forward_batch,
bs=forward_batch.batch_size,
forward_mode=ForwardMode.DECODE,
)
def call_fn(i, _forward_batch):
self.attn_backends[i].init_forward_metadata_out_graph(
inner_fb, in_capture=in_capture
)
self.common_template(forward_batch, call_fn)
def init_forward_metadata_in_graph(self, forward_batch: ForwardBatch) -> None:
for attn_backend in self.attn_backends:
attn_backend.init_forward_metadata_in_graph(forward_batch)
@@ -7,7 +7,6 @@ from sglang.srt.layers.attention.dsa.dsa_indexer import BaseIndexerMetadata
from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.model_executor.model_runner import ModelRunner
from sglang.srt.speculative.spec_info import SpecInput
class HybridAttnBackend(AttentionBackend):
@@ -52,6 +51,14 @@ class HybridAttnBackend(AttentionBackend):
else:
return self.prefill_backend
def init_forward_metadata_out_graph(
self,
forward_batch: ForwardBatch,
in_capture: bool = False,
):
backend = self._select_backend(forward_batch.forward_mode)
backend.init_forward_metadata_out_graph(forward_batch, in_capture=in_capture)
def init_forward_metadata(self, forward_batch: ForwardBatch):
backend = self._select_backend(forward_batch.forward_mode)
backend.init_forward_metadata(forward_batch)
@@ -66,50 +73,6 @@ class HybridAttnBackend(AttentionBackend):
# that will be used for target_verify.
self.prefill_backend.init_cuda_graph_state(max_bs, max_num_tokens)
def init_forward_metadata_capture_cuda_graph(
self,
bs: int,
num_tokens: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
encoder_lens: Optional[torch.Tensor],
forward_mode: ForwardMode,
spec_info: Optional[SpecInput],
):
backend = self._select_backend(forward_mode)
backend.init_forward_metadata_capture_cuda_graph(
bs,
num_tokens,
req_pool_indices,
seq_lens,
encoder_lens,
forward_mode,
spec_info,
)
def init_forward_metadata_replay_cuda_graph(
self,
bs: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
seq_lens_sum: int,
encoder_lens: Optional[torch.Tensor],
forward_mode: ForwardMode,
spec_info: Optional[SpecInput],
seq_lens_cpu: Optional[torch.Tensor],
):
backend = self._select_backend(forward_mode)
backend.init_forward_metadata_replay_cuda_graph(
bs,
req_pool_indices,
seq_lens,
seq_lens_sum,
encoder_lens,
forward_mode,
spec_info,
seq_lens_cpu,
)
def get_cuda_graph_seq_len_fill_value(self):
return self.decode_backend.get_cuda_graph_seq_len_fill_value()
@@ -258,6 +258,21 @@ class MambaAttnBackendBase(AttentionBackend):
has_mamba_track_mask=has_mamba_track_mask,
)
def init_forward_metadata_out_graph(
self,
forward_batch: ForwardBatch,
in_capture: bool = False,
):
# seq_lens_cpu is unused by _replay_metadata for the non-target-verify
# case but kept in the contract for compatibility.
self.forward_metadata = self._replay_metadata(
forward_batch.batch_size,
forward_batch.req_pool_indices,
forward_batch.forward_mode,
forward_batch.spec_info,
forward_batch.seq_lens_cpu if not in_capture else None,
)
def init_forward_metadata(self, forward_batch: ForwardBatch):
self._execute_deferred_mamba_cow_and_clear(forward_batch)
self.forward_metadata = self._forward_metadata(forward_batch)
@@ -393,42 +408,6 @@ class MambaAttnBackendBase(AttentionBackend):
track_ssm_final_dst.to(self.device, non_blocking=True),
)
def init_forward_metadata_capture_cuda_graph(
self,
bs: int,
num_tokens: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
encoder_lens: Optional[torch.Tensor],
forward_mode: ForwardMode,
spec_info: Optional[Union[EagleDraftInput, EagleVerifyInput]],
):
self.init_forward_metadata_replay_cuda_graph(
bs=bs,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
seq_lens_sum=None,
encoder_lens=encoder_lens,
forward_mode=forward_mode,
spec_info=spec_info,
seq_lens_cpu=None,
)
def init_forward_metadata_replay_cuda_graph(
self,
bs: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
seq_lens_sum: int,
encoder_lens: Optional[torch.Tensor],
forward_mode: ForwardMode,
spec_info: Optional[Union[EagleDraftInput, EagleVerifyInput]],
seq_lens_cpu: Optional[torch.Tensor],
):
self.forward_metadata = self._replay_metadata(
bs, req_pool_indices, forward_mode, spec_info, seq_lens_cpu
)
def init_forward_metadata_capture_cpu_graph(
self,
bs: int,
@@ -698,6 +677,27 @@ class Mamba2AttnBackend(MambaAttnBackendBase):
model_runner.server_args.mamba_track_interval >= self.mamba_chunk_size
), f"mamba_track_interval ({model_runner.server_args.mamba_track_interval}) must be >= mamba_chunk_size ({self.mamba_chunk_size})"
def init_forward_metadata_out_graph(
self,
forward_batch: ForwardBatch,
in_capture: bool = False,
):
metadata = self._replay_metadata(
forward_batch.batch_size,
forward_batch.req_pool_indices,
forward_batch.forward_mode,
forward_batch.spec_info,
forward_batch.seq_lens_cpu if not in_capture else None,
)
spec_info = forward_batch.spec_info
draft_token_num = spec_info.draft_token_num if spec_info is not None else 1
self.forward_metadata = Mamba2Metadata.prepare_decode(
metadata,
forward_batch.seq_lens,
is_target_verify=forward_batch.forward_mode.is_target_verify(),
draft_token_num=draft_token_num,
)
def init_forward_metadata(self, forward_batch: ForwardBatch):
self._execute_deferred_mamba_cow_and_clear(forward_batch)
metadata = self._forward_metadata(forward_batch)
@@ -707,49 +707,6 @@ class Mamba2AttnBackend(MambaAttnBackendBase):
forward_batch,
)
def init_forward_metadata_capture_cuda_graph(
self,
bs: int,
num_tokens: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
encoder_lens: Optional[torch.Tensor],
forward_mode: ForwardMode,
spec_info: Optional[Union[EagleDraftInput, EagleVerifyInput]],
):
self.init_forward_metadata_replay_cuda_graph(
bs=bs,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
seq_lens_sum=None,
encoder_lens=encoder_lens,
forward_mode=forward_mode,
spec_info=spec_info,
seq_lens_cpu=None,
)
def init_forward_metadata_replay_cuda_graph(
self,
bs: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
seq_lens_sum: int,
encoder_lens: Optional[torch.Tensor],
forward_mode: ForwardMode,
spec_info: Optional[Union[EagleDraftInput, EagleVerifyInput]],
seq_lens_cpu: Optional[torch.Tensor],
):
metadata = self._replay_metadata(
bs, req_pool_indices, forward_mode, spec_info, seq_lens_cpu
)
draft_token_num = spec_info.draft_token_num if spec_info is not None else 1
self.forward_metadata = Mamba2Metadata.prepare_decode(
metadata,
seq_lens,
is_target_verify=forward_mode.is_target_verify(),
draft_token_num=draft_token_num,
)
def forward(
self,
mixer: MambaMixer2,
@@ -847,6 +804,16 @@ class HybridLinearAttnBackend(AttentionBackend):
assert layer_id is not None, "either layer or layer_id must be provided"
return layer_id in self.full_attn_layers
def init_forward_metadata_out_graph(
self,
forward_batch: ForwardBatch,
in_capture: bool = False,
):
for attn_backend in self.attn_backend_list:
attn_backend.init_forward_metadata_out_graph(
forward_batch, in_capture=in_capture
)
def init_forward_metadata(self, forward_batch: ForwardBatch):
if forward_batch.forward_mode.is_draft_extend_v2():
# DRAFT_EXTEND_V2 only runs full-attn layers in the draft model,
@@ -864,27 +831,6 @@ class HybridLinearAttnBackend(AttentionBackend):
for attn_backend in self.attn_backend_list:
attn_backend.init_cpu_graph_state(max_bs, max_num_tokens)
def init_forward_metadata_capture_cuda_graph(
self,
bs: int,
num_tokens: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
encoder_lens: Optional[torch.Tensor],
forward_mode: ForwardMode,
spec_info: Optional[SpecInput],
):
for attn_backend in self.attn_backend_list:
attn_backend.init_forward_metadata_capture_cuda_graph(
bs,
num_tokens,
req_pool_indices,
seq_lens,
encoder_lens,
forward_mode,
spec_info,
)
def init_forward_metadata_capture_cpu_graph(
self,
bs: int,
@@ -906,29 +852,6 @@ class HybridLinearAttnBackend(AttentionBackend):
spec_info,
)
def init_forward_metadata_replay_cuda_graph(
self,
bs: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
seq_lens_sum: int,
encoder_lens: Optional[torch.Tensor],
forward_mode: ForwardMode,
spec_info: Optional[SpecInput],
seq_lens_cpu: Optional[torch.Tensor],
):
for attn_backend in self.attn_backend_list:
attn_backend.init_forward_metadata_replay_cuda_graph(
bs,
req_pool_indices,
seq_lens,
seq_lens_sum,
encoder_lens,
forward_mode,
spec_info,
seq_lens_cpu,
)
def get_cuda_graph_seq_len_fill_value(self):
return self.full_attn_backend.get_cuda_graph_seq_len_fill_value()
@@ -1,6 +1,5 @@
import logging
import math
from typing import Optional, Union
import torch
@@ -9,12 +8,13 @@ from sglang.srt.layers.attention.linear.lightning_attn import (
BailingLinearKernel,
linear_decode_forward_triton,
)
from sglang.srt.layers.attention.linear.linear_metadata import BailingLinearMetadata
from sglang.srt.layers.attention.linear.linear_metadata import (
BailingLinearMetadata,
)
from sglang.srt.layers.attention.linear.seg_la import SegLaMeta, seg_la_fwd
from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_executor.model_runner import ModelRunner
from sglang.srt.speculative.eagle_info import EagleDraftInput, EagleVerifyInput
logger = logging.getLogger(__name__)
@@ -71,6 +71,28 @@ class LightningAttentionBackend(MambaAttnBackendBase):
f"linear_backend for linear attention in hybrid_linear_backend: {self.linear_backend}"
)
def init_forward_metadata_out_graph(
self,
forward_batch: ForwardBatch,
in_capture: bool = False,
):
# seq_lens_cpu is unused by the underlying _replay_metadata for
# non-target-verify modes; pass it through for compatibility.
bs = forward_batch.batch_size
metadata = self._replay_metadata(
bs,
forward_batch.req_pool_indices,
forward_batch.forward_mode,
forward_batch.spec_info,
forward_batch.seq_lens_cpu if not in_capture else None,
)
self.forward_metadata = BailingLinearMetadata.prepare_decode(
metadata.query_start_loc,
metadata.mamba_cache_indices,
bs,
forward_batch.seq_lens,
)
def init_forward_metadata(self, forward_batch: ForwardBatch):
metadata = self._forward_metadata(forward_batch)
self.forward_metadata = BailingLinearMetadata.prepare_mixed(
@@ -79,45 +101,6 @@ class LightningAttentionBackend(MambaAttnBackendBase):
forward_batch,
)
def init_forward_metadata_capture_cuda_graph(
self,
bs: int,
num_tokens: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
encoder_lens: Optional[torch.Tensor],
forward_mode: ForwardMode,
spec_info: Optional[Union[EagleDraftInput, EagleVerifyInput]],
):
self.init_forward_metadata_replay_cuda_graph(
bs=bs,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
seq_lens_sum=None,
encoder_lens=encoder_lens,
forward_mode=forward_mode,
spec_info=spec_info,
seq_lens_cpu=None,
)
def init_forward_metadata_replay_cuda_graph(
self,
bs: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
seq_lens_sum: int,
encoder_lens: Optional[torch.Tensor],
forward_mode: ForwardMode,
spec_info: Optional[Union[EagleDraftInput, EagleVerifyInput]],
seq_lens_cpu: Optional[torch.Tensor],
):
metadata = self._replay_metadata(
bs, req_pool_indices, forward_mode, spec_info, seq_lens_cpu
)
self.forward_metadata = BailingLinearMetadata.prepare_decode(
metadata.query_start_loc, metadata.mamba_cache_indices, bs, seq_lens
)
@staticmethod
def _build_slope_tensor(
n_attention_heads: int, num_hidden_layers: int, device="cuda"
+138 -204
View File
@@ -1,13 +1,11 @@
from typing import TYPE_CHECKING, Callable, List, Optional
import torch
from types import SimpleNamespace
from typing import TYPE_CHECKING, Callable, List
from sglang.srt.batch_overlap import two_batch_overlap
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
from sglang.srt.speculative.spec_info import SpecInput
if TYPE_CHECKING:
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
class TboAttnBackend(AttentionBackend):
@@ -27,6 +25,88 @@ class TboAttnBackend(AttentionBackend):
children=[creator() for _ in range(2)],
)
def init_forward_metadata_out_graph(
self,
forward_batch: "ForwardBatch",
in_capture: bool = False,
):
self.primary.init_forward_metadata_out_graph(
forward_batch=forward_batch, in_capture=in_capture
)
tbo_children = getattr(forward_batch, "tbo_children", None)
if tbo_children is not None:
for child, forward_batch_child in zip(
self.children, tbo_children, strict=True
):
if forward_batch_child.batch_size > 0:
child.init_forward_metadata_out_graph(
forward_batch=forward_batch_child, in_capture=in_capture
)
return
if in_capture:
return
# Replay path: build_replay_fb_view returns a SimpleNamespace and
# tbo_plugin.replay_prepare does not call prepare_raw, so split the
# padded buffers here using the same indices the eager path would.
self._dispatch_children_from_replay_view(forward_batch)
def _dispatch_children_from_replay_view(self, fb_view) -> None:
bs = fb_view.batch_size
forward_mode = fb_view.forward_mode
spec_info = fb_view.spec_info
token_num_per_seq = two_batch_overlap.get_token_num_per_seq(
forward_mode=forward_mode, spec_info=spec_info
)
num_tokens = bs * token_num_per_seq
(
tbo_split_seq_index,
tbo_split_token_index,
) = two_batch_overlap.compute_split_indices_for_cuda_graph_replay(
forward_mode=forward_mode,
cuda_graph_num_tokens=num_tokens,
spec_info=spec_info,
)
bs_left = tbo_split_seq_index
bs_right = bs - bs_left
for child, child_bs, seq_slice, tok_slice in (
(
self.children[0],
bs_left,
slice(None, tbo_split_seq_index),
slice(None, tbo_split_token_index),
),
(
self.children[1],
bs_right,
slice(tbo_split_seq_index, None),
slice(tbo_split_token_index, None),
),
):
if child_bs == 0:
continue
child_fb_view = _build_tbo_child_replay_fb_view(
fb_view,
child_bs=child_bs,
seq_slice=seq_slice,
tok_slice=tok_slice,
token_num_per_seq=token_num_per_seq,
)
child.init_forward_metadata_out_graph(
forward_batch=child_fb_view, in_capture=False
)
def init_forward_metadata_in_graph(self, forward_batch: "ForwardBatch"):
self.primary.init_forward_metadata_in_graph(forward_batch=forward_batch)
tbo_children = getattr(forward_batch, "tbo_children", None)
if tbo_children is not None:
for child, forward_batch_child in zip(
self.children, tbo_children, strict=True
):
if forward_batch_child.batch_size > 0:
child.init_forward_metadata_in_graph(
forward_batch=forward_batch_child
)
def init_forward_metadata(self, forward_batch: "ForwardBatch"):
self.primary.init_forward_metadata(forward_batch=forward_batch)
if forward_batch.tbo_children is not None:
@@ -42,140 +122,10 @@ class TboAttnBackend(AttentionBackend):
# TODO for children, maybe can provide *smaller* max_bs to optimize
item.init_cuda_graph_state(max_bs=max_bs, max_num_tokens=max_num_tokens)
def init_forward_metadata_capture_cuda_graph(
self,
bs: int,
num_tokens: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
encoder_lens: Optional[torch.Tensor],
forward_mode: "ForwardMode",
spec_info: Optional[SpecInput],
):
self.primary.init_forward_metadata_capture_cuda_graph(
bs=bs,
num_tokens=num_tokens,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
encoder_lens=encoder_lens,
forward_mode=forward_mode,
spec_info=spec_info,
)
self._init_forward_metadata_cuda_graph_children(
fn_name="init_forward_metadata_capture_cuda_graph",
bs=bs,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
encoder_lens=encoder_lens,
forward_mode=forward_mode,
spec_info=spec_info,
capture_num_tokens=num_tokens,
)
def init_forward_metadata_replay_cuda_graph(
self,
bs: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
seq_lens_sum: int,
encoder_lens: Optional[torch.Tensor],
forward_mode: "ForwardMode",
spec_info: Optional[SpecInput],
seq_lens_cpu: Optional[torch.Tensor],
):
self.primary.init_forward_metadata_replay_cuda_graph(
bs=bs,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
seq_lens_sum=seq_lens_sum,
encoder_lens=encoder_lens,
forward_mode=forward_mode,
spec_info=spec_info,
seq_lens_cpu=seq_lens_cpu,
)
self._init_forward_metadata_cuda_graph_children(
fn_name="init_forward_metadata_replay_cuda_graph",
bs=bs,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
encoder_lens=encoder_lens,
forward_mode=forward_mode,
spec_info=spec_info,
replay_seq_lens_sum=seq_lens_sum,
replay_seq_lens_cpu=seq_lens_cpu,
)
def _init_forward_metadata_cuda_graph_children(
self,
fn_name: str,
# common args
bs: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
encoder_lens: Optional[torch.Tensor],
forward_mode: "ForwardMode",
spec_info: Optional[SpecInput],
# capture args
capture_num_tokens: int = None,
# replay args
replay_seq_lens_sum: int = None,
replay_seq_lens_cpu: Optional[torch.Tensor] = None,
):
token_num_per_seq = two_batch_overlap.get_token_num_per_seq(
forward_mode=forward_mode, spec_info=spec_info
)
if fn_name == "init_forward_metadata_capture_cuda_graph":
assert (
capture_num_tokens == bs * token_num_per_seq
), "For target-verify or decode mode, num_tokens should be equal to token_num_per_seq * bs"
num_tokens = bs * token_num_per_seq
tbo_split_seq_index, tbo_split_token_index = (
two_batch_overlap.compute_split_indices_for_cuda_graph_replay(
forward_mode=forward_mode,
cuda_graph_num_tokens=num_tokens,
spec_info=spec_info,
)
)
num_tokens_child_left = tbo_split_token_index
num_tokens_child_right = num_tokens - tbo_split_token_index
bs_child_left = tbo_split_seq_index
bs_child_right = bs - bs_child_left
assert (
num_tokens_child_left > 0 and num_tokens_child_right > 0
), f"{num_tokens_child_left=} {num_tokens_child_right=} {forward_mode=} {num_tokens=}"
common_pre_split_args = dict(
fn_name=fn_name,
bs=bs,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
encoder_lens=encoder_lens,
forward_mode=forward_mode,
spec_info=spec_info,
capture_num_tokens=capture_num_tokens,
replay_seq_lens_sum=replay_seq_lens_sum,
replay_seq_lens_cpu=replay_seq_lens_cpu,
)
args_left = _init_forward_metadata_cuda_graph_split(
output_bs=bs_child_left,
seq_slice=slice(None, tbo_split_seq_index),
**common_pre_split_args,
)
args_right = _init_forward_metadata_cuda_graph_split(
output_bs=bs_child_right,
seq_slice=slice(tbo_split_seq_index, None),
**common_pre_split_args,
)
child_left, child_right = self.children
getattr(child_left, fn_name)(**args_left)
getattr(child_right, fn_name)(**args_right)
def on_after_cuda_graph_warmup(self):
self.primary.on_after_cuda_graph_warmup()
for child in self.children:
child.on_after_cuda_graph_warmup()
def get_cuda_graph_seq_len_fill_value(self):
ans = self.primary.get_cuda_graph_seq_len_fill_value()
@@ -196,75 +146,59 @@ class TboAttnBackend(AttentionBackend):
return self.primary.get_indexer_metadata(layer_id, forward_batch)
def _init_forward_metadata_cuda_graph_split(
fn_name: str,
def _build_tbo_child_replay_fb_view(
fb_view,
*,
child_bs: int,
seq_slice: slice,
output_bs: int,
# common args
bs: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
encoder_lens: Optional[torch.Tensor],
forward_mode: "ForwardMode",
spec_info: Optional[SpecInput],
# capture args
capture_num_tokens: int = None,
# replay args
replay_seq_lens_sum: int = None,
replay_seq_lens_cpu: Optional[torch.Tensor] = None,
):
token_num_per_seq = two_batch_overlap.get_token_num_per_seq(
forward_mode=forward_mode, spec_info=spec_info
)
assert encoder_lens is None, "encoder_lens is not supported yet"
tok_slice: slice,
token_num_per_seq: int,
) -> SimpleNamespace:
"""Slice a parent replay fb_view into a per-child view.
Mirrors the legacy ``_init_forward_metadata_cuda_graph_split`` (deleted
along with the cuda_graph variants) for the new
``init_forward_metadata_out_graph(fb_view)`` contract: padded
capture-time buffers are sliced per child, spec_info is split, and
seq_lens_sum is recomputed from the sliced ``seq_lens_cpu``.
"""
assert (
getattr(fb_view, "encoder_lens", None) is None
), "TBO replay split does not support encoder_lens yet"
spec_info = getattr(fb_view, "spec_info", None)
if spec_info is not None:
output_spec_info = two_batch_overlap.split_spec_info(
start_seq = seq_slice.start or 0
end_seq = seq_slice.stop if seq_slice.stop is not None else start_seq + child_bs
child_spec_info = two_batch_overlap.split_spec_info(
spec_info=spec_info,
start_seq_index=seq_slice.start if seq_slice.start is not None else 0,
end_seq_index=seq_slice.stop if seq_slice.stop is not None else bs,
start_token_index=(
seq_slice.start * token_num_per_seq
if seq_slice.start is not None
else 0
),
end_token_index=(
seq_slice.stop * token_num_per_seq
if seq_slice.stop is not None
else bs * token_num_per_seq
),
start_seq_index=start_seq,
end_seq_index=end_seq,
start_token_index=start_seq * token_num_per_seq,
end_token_index=end_seq * token_num_per_seq,
)
else:
output_spec_info = None
ans = dict(
bs=output_bs,
req_pool_indices=req_pool_indices[seq_slice],
seq_lens=seq_lens[seq_slice],
# directly forward
forward_mode=forward_mode,
# ignore
child_spec_info = None
child_seq_lens_cpu = fb_view.seq_lens_cpu[seq_slice]
parent_input_ids = getattr(fb_view, "input_ids", None)
parent_out_cache_loc = getattr(fb_view, "out_cache_loc", None)
return SimpleNamespace(
batch_size=child_bs,
forward_mode=fb_view.forward_mode,
actual_forward_mode=getattr(
fb_view, "actual_forward_mode", fb_view.forward_mode
),
input_ids=(
parent_input_ids[tok_slice] if parent_input_ids is not None else None
),
req_pool_indices=fb_view.req_pool_indices[seq_slice],
seq_lens=fb_view.seq_lens[seq_slice],
seq_lens_sum=int(child_seq_lens_cpu.sum()),
seq_lens_cpu=child_seq_lens_cpu,
encoder_lens=None,
spec_info=output_spec_info,
out_cache_loc=(
parent_out_cache_loc[tok_slice]
if parent_out_cache_loc is not None
else None
),
spec_info=child_spec_info,
)
if fn_name == "init_forward_metadata_capture_cuda_graph":
assert (
capture_num_tokens == bs * token_num_per_seq
), "Only support num_tokens==bs * token_num_per_seq for target-verify or decode mode"
ans.update(
dict(
num_tokens=output_bs * token_num_per_seq,
)
)
elif fn_name == "init_forward_metadata_replay_cuda_graph":
output_seq_lens_cpu = replay_seq_lens_cpu[seq_slice]
ans.update(
dict(
seq_lens_sum=output_seq_lens_cpu.sum().item(),
seq_lens_cpu=output_seq_lens_cpu,
)
)
else:
raise NotImplementedError
return ans
@@ -458,6 +458,59 @@ class TritonAttnBackend(AttentionBackend):
)
return qo_indptr, kv_indptr, num_tokens_per_bs
def init_forward_metadata_out_graph(
self,
forward_batch: ForwardBatch,
in_capture: bool = False,
):
bs = forward_batch.batch_size
req_pool_indices = forward_batch.req_pool_indices
seq_lens = forward_batch.seq_lens
forward_mode = forward_batch.forward_mode
spec_info = forward_batch.spec_info
if in_capture:
assert forward_batch.encoder_lens is None, "Not supported"
# Multi-step speculative decode: kv buffers come from spec_info
# rather than the cuda-graph pool, so replay is not involved.
if forward_mode.is_decode_or_idle() and spec_info is not None:
self.forward_metadata = ForwardMetadata(
attn_logits=self.cuda_graph_attn_logits,
attn_lse=self.cuda_graph_attn_lse,
max_extend_len=None,
num_kv_splits=self.cuda_graph_num_kv_splits,
kv_indptr=spec_info.kv_indptr,
kv_indices=spec_info.kv_indices,
qo_indptr=None,
custom_mask=None,
mask_indptr=None,
window_kv_indptr=self.window_kv_indptr,
window_kv_indices=None,
window_num_kv_splits=None,
window_kv_offsets=None,
swa_attn_logits=self.cuda_graph_swa_attn_logits,
)
return
self._apply_cuda_graph_metadata(
bs=bs,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
forward_mode=forward_mode,
spec_info=spec_info,
)
self.forward_metadata = self._build_cuda_graph_forward_metadata(
bs, forward_mode, spec_info
)
else:
self._apply_cuda_graph_metadata(
bs=bs,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
forward_mode=forward_mode,
spec_info=spec_info,
)
def init_forward_metadata(self, forward_batch: ForwardBatch):
"""Init auxiliary variables for triton attention backend."""
@@ -835,66 +888,18 @@ class TritonAttnBackend(AttentionBackend):
else:
raise ValueError(f"Invalid forward mode: {forward_mode=} for CUDA Graph.")
def init_forward_metadata_capture_cuda_graph(
self,
bs: int,
num_tokens: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
encoder_lens: Optional[torch.Tensor],
forward_mode: ForwardMode,
spec_info: Optional[SpecInput],
):
assert encoder_lens is None, "Not supported"
# Multi-step speculative decode: kv buffers come from spec_info rather
# than the cuda-graph pool, so replay is not involved for this path.
if forward_mode.is_decode_or_idle() and spec_info is not None:
self.forward_metadata = ForwardMetadata(
attn_logits=self.cuda_graph_attn_logits,
attn_lse=self.cuda_graph_attn_lse,
max_extend_len=None,
num_kv_splits=self.cuda_graph_num_kv_splits,
kv_indptr=spec_info.kv_indptr,
kv_indices=spec_info.kv_indices,
qo_indptr=None,
custom_mask=None,
mask_indptr=None,
window_kv_indptr=self.window_kv_indptr,
window_kv_indices=None,
window_num_kv_splits=None,
window_kv_offsets=None,
swa_attn_logits=self.cuda_graph_swa_attn_logits,
)
return
# Run the same buffer update as replay, then freeze the result into
# a ForwardMetadata whose tensor fields point into the cuda-graph buffers.
self.init_forward_metadata_replay_cuda_graph(
bs=bs,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
seq_lens_sum=None,
encoder_lens=encoder_lens,
forward_mode=forward_mode,
spec_info=spec_info,
seq_lens_cpu=None,
)
self.forward_metadata = self._build_cuda_graph_forward_metadata(
bs, forward_mode, spec_info
)
def init_forward_metadata_replay_cuda_graph(
def _apply_cuda_graph_metadata(
self,
bs: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
seq_lens_sum: int,
encoder_lens: Optional[torch.Tensor],
forward_mode: ForwardMode,
spec_info: Optional[SpecInput],
seq_lens_cpu: Optional[torch.Tensor],
):
"""Shared capture+replay body for the cuda-graph init path.
Public entry: :py:meth:`init_forward_metadata_out_graph`.
"""
# NOTE: encoder_lens expected to be zeros or None
if forward_mode.is_decode_or_idle():
assert spec_info is None, "Multi-step cuda graph init is not done here."
@@ -1417,36 +1422,45 @@ class TritonMultiStepDraftBackend:
cuda_graph_num_kv_splits_buf=self.cuda_graph_num_kv_splits,
)
def init_forward_metadata_capture_cuda_graph(self, forward_batch: ForwardBatch):
def call_fn(i, forward_batch):
self.attn_backends[i].init_forward_metadata_capture_cuda_graph(
forward_batch.batch_size,
forward_batch.batch_size * self.topk,
forward_batch.req_pool_indices,
forward_batch.seq_lens,
encoder_lens=None,
def init_forward_metadata_out_graph(
self,
forward_batch: ForwardBatch,
in_capture: bool = False,
):
from sglang.srt.model_executor.forward_batch_info import build_inner_fb_view
if in_capture:
inner_fb = build_inner_fb_view(
forward_batch,
bs=forward_batch.batch_size,
forward_mode=ForwardMode.DECODE,
spec_info=forward_batch.spec_info,
)
self.common_template(forward_batch, None, call_fn)
def call_fn(i, _forward_batch):
self.attn_backends[i].init_forward_metadata_out_graph(
inner_fb, in_capture=True
)
def init_forward_metadata_replay_cuda_graph(
self, forward_batch: ForwardBatch, bs: int
):
self.common_template(forward_batch, None, None)
self.common_template(forward_batch, None, call_fn)
else:
bs = forward_batch.batch_size
self.common_template(forward_batch, None, None)
# NOTE: Multi-step's attention backends use the slice of
# - kv_indptr buffer (cuda graph and non-cuda graph)
# - kv_indices buffer (cuda graph only)
# So we don't need to assign the KV indices inside the attention backend.
# NOTE: Multi-step's attention backends use the slice of
# - kv_indptr buffer (cuda graph and non-cuda graph)
# - kv_indices buffer (cuda graph only)
# So we don't need to assign the KV indices inside the attention backend.
# Compute num_kv_splits only once
num_token = forward_batch.batch_size * self.topk
self.attn_backends[-1].get_num_kv_splits(
self.attn_backends[-1].cuda_graph_num_kv_splits[:num_token],
forward_batch.seq_lens[:bs],
)
# Compute num_kv_splits only once
num_token = bs * self.topk
self.attn_backends[-1].get_num_kv_splits(
self.attn_backends[-1].cuda_graph_num_kv_splits[:num_token],
forward_batch.seq_lens[:bs],
)
def init_forward_metadata_in_graph(self, forward_batch: ForwardBatch) -> None:
for attn_backend in self.attn_backends:
attn_backend.init_forward_metadata_in_graph(forward_batch)
def update_sliding_window_buffer(
@@ -397,50 +397,19 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
return metadata
def init_forward_metadata_capture_cuda_graph(
self,
bs: int,
num_tokens: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
encoder_lens: Optional[torch.Tensor],
forward_mode: ForwardMode,
spec_info: Optional[SpecInput],
):
"""Initialize metadata for CUDA graph capture."""
seq_lens_cpu = seq_lens.cpu()
self._build_cuda_graph_metadata(
bs, num_tokens, forward_mode, spec_info, seq_lens.device
)
self.init_forward_metadata_replay_cuda_graph(
bs=bs,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
seq_lens_sum=None,
encoder_lens=encoder_lens,
forward_mode=forward_mode,
spec_info=spec_info,
seq_lens_cpu=seq_lens_cpu,
)
if forward_mode.is_draft_extend():
# CUDA graph bakes max_seq_len_q as a constant. replay() sets it to
# max(num_accept_tokens_cpu) which is None/empty at capture time,
# falling back to 1. Restore the correct upper bound so the kernel
# sees num_tokens_per_bs (not 1) for all replays of this graph.
self.forward_metadata.max_seq_len_q = num_tokens // bs
def init_forward_metadata_replay_cuda_graph(
def _apply_cuda_graph_metadata(
self,
bs: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
seq_lens_sum: int,
encoder_lens: Optional[torch.Tensor],
forward_mode: ForwardMode,
spec_info: Optional[SpecInput],
seq_lens_cpu: Optional[torch.Tensor],
):
"""Replay CUDA graph with new inputs."""
"""Shared capture+replay body for the cuda-graph init path.
Public entry: :py:meth:`init_forward_metadata_out_graph`.
"""
seq_lens = seq_lens[:bs]
seq_lens_cpu = seq_lens_cpu[:bs]
req_pool_indices = req_pool_indices[:bs]
@@ -571,6 +540,48 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
page_size=self.page_size,
)
def init_forward_metadata_out_graph(
self,
forward_batch: ForwardBatch,
in_capture: bool = False,
):
bs = forward_batch.batch_size
req_pool_indices = forward_batch.req_pool_indices
seq_lens = forward_batch.seq_lens
encoder_lens = forward_batch.encoder_lens
forward_mode = forward_batch.forward_mode
spec_info = forward_batch.spec_info
if in_capture:
num_tokens = forward_batch.positions.numel()
seq_lens_cpu = seq_lens.cpu()
self._build_cuda_graph_metadata(
bs, num_tokens, forward_mode, spec_info, seq_lens.device
)
self._apply_cuda_graph_metadata(
bs=bs,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
forward_mode=forward_mode,
spec_info=spec_info,
seq_lens_cpu=seq_lens_cpu,
)
if forward_mode.is_draft_extend():
# CUDA graph bakes max_seq_len_q as a constant. replay() sets it
# to max(num_accept_tokens_cpu) which is None/empty at capture
# time, falling back to 1. Restore the correct upper bound so
# the kernel sees num_tokens_per_bs (not 1) for all replays.
self.forward_metadata.max_seq_len_q = num_tokens // bs
else:
self._apply_cuda_graph_metadata(
bs=bs,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
forward_mode=forward_mode,
spec_info=spec_info,
seq_lens_cpu=forward_batch.seq_lens_cpu,
)
def init_forward_metadata(self, forward_batch: ForwardBatch):
"""Initialize the metadata for a forward pass."""
@@ -899,39 +910,25 @@ class TRTLLMHAAttnMultiStepDraftBackend(FlashInferMultiStepDraftBackend):
for i in range(self.speculative_num_steps - 1):
self.attn_backends[i].init_cuda_graph_state(max_bs, max_num_tokens)
def init_forward_metadata_capture_cuda_graph(
def init_forward_metadata_out_graph(
self,
forward_batch: ForwardBatch,
in_capture: bool = False,
):
from sglang.srt.model_executor.forward_batch_info import build_inner_fb_view
assert forward_batch.spec_info is not None
assert forward_batch.spec_info.is_draft_input()
# TRTLLM-MHA uses encoder_lens from the original fb for inner dispatch
# (FlashInfer parent forces encoder_lens=None instead).
inner_fb = build_inner_fb_view(
forward_batch,
bs=forward_batch.batch_size,
forward_mode=ForwardMode.DECODE,
encoder_lens=forward_batch.encoder_lens,
)
for i in range(self.speculative_num_steps - 1):
self.attn_backends[i].init_forward_metadata_capture_cuda_graph(
forward_batch.batch_size,
forward_batch.batch_size * self.topk,
forward_batch.req_pool_indices,
forward_batch.seq_lens,
encoder_lens=forward_batch.encoder_lens,
forward_mode=ForwardMode.DECODE,
spec_info=forward_batch.spec_info,
)
def init_forward_metadata_replay_cuda_graph(
self, forward_batch: ForwardBatch, bs: int
):
assert forward_batch.spec_info is not None
assert forward_batch.spec_info.is_draft_input()
for i in range(self.speculative_num_steps - 1):
self.attn_backends[i].init_forward_metadata_replay_cuda_graph(
bs,
forward_batch.req_pool_indices,
forward_batch.seq_lens,
forward_batch.seq_lens_sum,
encoder_lens=forward_batch.encoder_lens,
forward_mode=ForwardMode.DECODE,
spec_info=forward_batch.spec_info,
seq_lens_cpu=forward_batch.seq_lens_cpu,
self.attn_backends[i].init_forward_metadata_out_graph(
inner_fb, in_capture=in_capture
)
@@ -38,7 +38,6 @@ if is_flashinfer_available():
if TYPE_CHECKING:
from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.model_executor.model_runner import ModelRunner
from sglang.srt.speculative.spec_info import SpecInput
logger = logging.getLogger(__name__)
@@ -499,83 +498,23 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
self.decode_cuda_graph_metadata[bs] = metadata
self.forward_decode_metadata = metadata
def init_forward_metadata_capture_cuda_graph(
self,
bs: int,
num_tokens: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
encoder_lens: Optional[torch.Tensor],
forward_mode: ForwardMode,
spec_info: Optional[SpecInput],
):
"""Initialize metadata for CUDA graph capture."""
# Delegate to parent for non-decode modes.
if (
not forward_mode.is_decode_or_idle()
and not forward_mode.is_target_verify()
and not forward_mode.is_draft_extend(include_v2=True)
):
return super().init_forward_metadata_capture_cuda_graph(
bs,
num_tokens,
req_pool_indices,
seq_lens,
encoder_lens,
forward_mode,
spec_info,
)
self._init_cuda_graph_metadata(
bs, num_tokens, forward_mode, seq_lens, seq_lens.device
)
self.init_forward_metadata_replay_cuda_graph(
bs=bs,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
seq_lens_sum=None,
encoder_lens=encoder_lens,
forward_mode=forward_mode,
spec_info=spec_info,
seq_lens_cpu=seq_lens.cpu(),
)
def init_forward_metadata_replay_cuda_graph(
def _apply_cuda_graph_metadata(
self,
bs: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
seq_lens_sum: int,
encoder_lens: Optional[torch.Tensor],
forward_mode: ForwardMode,
spec_info: Optional[SpecInput],
seq_lens_cpu: Optional[torch.Tensor],
):
"""Replay CUDA graph with new inputs."""
# Delegate to parent for non-decode modes.
if (
not forward_mode.is_decode_or_idle()
and not forward_mode.is_target_verify()
and not forward_mode.is_draft_extend(include_v2=True)
):
return super().init_forward_metadata_replay_cuda_graph(
bs,
req_pool_indices,
seq_lens,
seq_lens_sum,
encoder_lens,
forward_mode,
spec_info,
seq_lens_cpu,
)
"""Shared decode / target-verify / draft-extend capture+replay body.
Public entry: :py:meth:`init_forward_metadata_out_graph` (which routes
the non-decode-family modes to the FlashInferMLA parent).
"""
metadata = self.decode_cuda_graph_metadata[bs]
if forward_mode.is_target_verify():
seq_lens = seq_lens[:bs] + self.num_draft_tokens
metadata.seq_lens_k.copy_(seq_lens.to(dtype=torch.int32))
del seq_lens_sum # not handle "num_draft_tokens" but we do not need it
elif forward_mode.is_draft_extend(include_v2=True):
num_tokens_per_bs = self.num_draft_tokens
metadata.max_seq_len_q = num_tokens_per_bs
@@ -618,6 +557,47 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
if fallback_to_flashinfer_impl:
super().init_mha_chunk_metadata(forward_batch)
def init_forward_metadata_out_graph(
self,
forward_batch: ForwardBatch,
in_capture: bool = False,
):
forward_mode = forward_batch.forward_mode
if (
not forward_mode.is_decode_or_idle()
and not forward_mode.is_target_verify()
and not forward_mode.is_draft_extend(include_v2=True)
):
return super().init_forward_metadata_out_graph(
forward_batch, in_capture=in_capture
)
bs = forward_batch.batch_size
if in_capture:
num_tokens = forward_batch.positions.numel()
seq_lens_cpu = forward_batch.seq_lens.cpu()
self._init_cuda_graph_metadata(
bs,
num_tokens,
forward_mode,
forward_batch.seq_lens,
forward_batch.seq_lens.device,
)
self._apply_cuda_graph_metadata(
bs=bs,
req_pool_indices=forward_batch.req_pool_indices,
seq_lens=forward_batch.seq_lens,
forward_mode=forward_mode,
)
else:
self._apply_cuda_graph_metadata(
bs=bs,
req_pool_indices=forward_batch.req_pool_indices,
seq_lens=forward_batch.seq_lens,
forward_mode=forward_mode,
)
def init_forward_metadata(self, forward_batch: ForwardBatch):
"""Initialize the metadata for a forward pass."""
# Delegate to parent for non-decode modes.
@@ -1284,17 +1264,21 @@ class TRTLLMMLAMultiStepDraftBackend(FlashInferMLAMultiStepDraftBackend):
for i in range(self.speculative_num_steps - 1):
self.attn_backends[i].init_forward_metadata(forward_batch)
def init_forward_metadata_replay_cuda_graph(
self, forward_batch: ForwardBatch, bs: int
def init_forward_metadata_out_graph(
self,
forward_batch: ForwardBatch,
in_capture: bool = False,
):
for i in range(self.speculative_num_steps - 1):
self.attn_backends[i].init_forward_metadata_replay_cuda_graph(
bs,
forward_batch.req_pool_indices,
forward_batch.seq_lens,
seq_lens_sum=None,
encoder_lens=None,
forward_mode=ForwardMode.DECODE,
spec_info=forward_batch.spec_info,
seq_lens_cpu=forward_batch.seq_lens_cpu,
from sglang.srt.model_executor.forward_batch_info import build_inner_fb_view
if in_capture:
return super().init_forward_metadata_out_graph(
forward_batch, in_capture=in_capture
)
inner_fb = build_inner_fb_view(
forward_batch,
bs=forward_batch.batch_size,
forward_mode=ForwardMode.DECODE,
)
for i in range(self.speculative_num_steps - 1):
self.attn_backends[i].init_forward_metadata_out_graph(inner_fb)
@@ -147,6 +147,53 @@ class WaveAttnBackend(AttentionBackend):
MAX_NUM_SEQ=SCHEDULE_SEQ,
)
def init_forward_metadata_out_graph(
self,
forward_batch: ForwardBatch,
in_capture: bool = False,
):
bs = forward_batch.batch_size
req_pool_indices = forward_batch.req_pool_indices
seq_lens = forward_batch.seq_lens
forward_mode = forward_batch.forward_mode
spec_info = forward_batch.spec_info
if in_capture:
assert forward_batch.encoder_lens is None, "Not supported"
# kv buffers come from spec_info rather than the cuda-graph pool.
if forward_mode.is_decode_or_idle() and spec_info is not None:
self.forward_metadata = ForwardMetadata(
attn_logits=self.cuda_graph_attn_logits,
attn_lse=self.cuda_graph_attn_lse,
max_extend_len=None,
num_kv_splits=self.cuda_graph_num_kv_splits,
kv_indptr=spec_info.kv_indptr,
kv_indices=spec_info.kv_indices,
qo_indptr=None,
custom_mask=None,
mask_indptr=None,
)
return
self._apply_cuda_graph_metadata(
bs=bs,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
forward_mode=forward_mode,
spec_info=spec_info,
)
self.forward_metadata = self._build_cuda_graph_forward_metadata(
bs, forward_mode, spec_info
)
else:
self._apply_cuda_graph_metadata(
bs=bs,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
forward_mode=forward_mode,
spec_info=spec_info,
)
def init_forward_metadata(self, forward_batch: ForwardBatch):
"""Init auxiliary variables for wave attention backend."""
@@ -373,59 +420,18 @@ class WaveAttnBackend(AttentionBackend):
else:
raise ValueError(f"Invalid forward mode: {forward_mode=} for CUDA Graph.")
def init_forward_metadata_capture_cuda_graph(
self,
bs: int,
num_tokens: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
encoder_lens: Optional[torch.Tensor],
forward_mode: ForwardMode,
spec_info: Optional[SpecInput],
):
assert encoder_lens is None, "Not supported"
# Multi-step speculative decode: kv buffers come from spec_info rather than
# the cuda-graph pool, so replay is not involved for this path.
if forward_mode.is_decode_or_idle() and spec_info is not None:
self.forward_metadata = ForwardMetadata(
attn_logits=self.cuda_graph_attn_logits,
attn_lse=self.cuda_graph_attn_lse,
max_extend_len=None,
num_kv_splits=self.cuda_graph_num_kv_splits,
kv_indptr=spec_info.kv_indptr,
kv_indices=spec_info.kv_indices,
qo_indptr=None,
custom_mask=None,
mask_indptr=None,
)
return
self.init_forward_metadata_replay_cuda_graph(
bs=bs,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
seq_lens_sum=None,
encoder_lens=encoder_lens,
forward_mode=forward_mode,
spec_info=spec_info,
seq_lens_cpu=None,
)
self.forward_metadata = self._build_cuda_graph_forward_metadata(
bs, forward_mode, spec_info
)
def init_forward_metadata_replay_cuda_graph(
def _apply_cuda_graph_metadata(
self,
bs: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
seq_lens_sum: int,
encoder_lens: Optional[torch.Tensor],
forward_mode: ForwardMode,
spec_info: Optional[SpecInput],
seq_lens_cpu: Optional[torch.Tensor],
):
"""Shared capture+replay body for the cuda-graph init path.
Public entry: :py:meth:`init_forward_metadata_out_graph`.
"""
if forward_mode.is_decode_or_idle():
kv_indptr = self.kv_indptr
kv_indices = self.cuda_graph_kv_indices
@@ -906,7 +906,10 @@ class XPUAttentionBackend(AttentionBackend):
return 1
def _init_local_attn_metadata(
self, forwardbatch: ForwardBatch, metadata: FlashAttentionMetadata, device
self,
forwardbatch: ForwardBatch,
metadata: FlashAttentionMetadata,
device,
):
"""Centralized utility to initialize local_attn_metadata if chunked attention is enabled."""
if self.attention_chunk_size is None:
@@ -24,6 +24,7 @@ import os
from contextlib import contextmanager
from dataclasses import dataclass
from functools import partial
from types import SimpleNamespace
from typing import TYPE_CHECKING, Callable, Dict, List, Optional, Tuple, Union
import torch
@@ -107,6 +108,57 @@ if TYPE_CHECKING:
_has_foreach_copy = hasattr(torch, "_foreach_copy_")
def build_replay_fb_view(
forward_batch: "ForwardBatch",
buffers: "DecodeInputBuffers",
bs: int,
raw_bs: int,
num_tokens: int,
seq_len_fill_value: int,
capture_forward_mode: "ForwardMode",
is_encoder_decoder: bool,
) -> SimpleNamespace:
"""Construct a ForwardBatch-like view for backend replay-side init.
Combines the original ``forward_batch`` (for unpadded / per-iter
fields like ``spec_info``, ``out_cache_loc``, and the runtime
``actual_forward_mode``) with the padded capture-time buffers from
``buffers`` (for ``req_pool_indices``, ``seq_lens``, ``seq_lens_cpu``,
``encoder_lens``).
Field semantics:
- ``forward_mode``: the capture-time mode (``capture_forward_mode``),
used by backends for bucket / dispatch decisions (e.g. choosing
between decode / target-verify / draft-extend code paths).
- ``actual_forward_mode``: the original runtime ``forward_batch
.forward_mode``, which may be ``IDLE`` even when the captured
graph corresponds to ``DECODE``. DSV4's replay metadata prep
uses this for IDLE-batch substitution; other backends ignore it.
This view subsumes the ``_replay_forward_batch`` side channel DSV4
previously read out-of-band step 04 swaps that mechanism for this
explicit fb_view field.
"""
return SimpleNamespace(
batch_size=bs,
forward_mode=capture_forward_mode,
actual_forward_mode=forward_batch.forward_mode,
input_ids=buffers.input_ids[:num_tokens],
req_pool_indices=buffers.req_pool_indices[:bs],
seq_lens=buffers.seq_lens[:bs],
seq_lens_sum=(
None
if forward_batch.seq_lens_sum is None
else forward_batch.seq_lens_sum + (bs - raw_bs) * seq_len_fill_value
),
seq_lens_cpu=buffers.seq_lens_cpu[:bs],
encoder_lens=buffers.encoder_lens[:bs] if is_encoder_decoder else None,
out_cache_loc=getattr(forward_batch, "out_cache_loc", None),
spec_info=forward_batch.spec_info,
)
def _grouped_foreach_copy_(dsts: List[torch.Tensor], srcs: List[torch.Tensor]) -> None:
"""Call torch._foreach_copy_ grouped by (dst_dtype, src_dtype) pairs."""
@@ -1099,15 +1151,7 @@ class CudaGraphRunner:
if lora_ids is not None:
self.model_runner.lora_manager.prepare_lora_batch(forward_batch)
attn_backend.init_forward_metadata_capture_cuda_graph(
bs,
num_tokens,
req_pool_indices,
seq_lens,
encoder_lens,
forward_batch.forward_mode,
forward_batch.spec_info,
)
attn_backend.init_forward_metadata_out_graph(forward_batch, in_capture=True)
def run_once():
# Without this, warmup-1 caches the translation; the capture
@@ -1116,6 +1160,10 @@ class CudaGraphRunner:
if self.model_runner.is_hybrid_swa:
self.model_runner.token_to_kv_pool.invalidate_loc_cache()
# Must run inside the capture block: warmup mutations here are
# undone by on_after_cuda_graph_warmup so capture starts clean.
attn_backend.init_forward_metadata_in_graph(forward_batch)
forward_batch.dp_local_start_pos = forward_batch.dp_local_num_tokens = (
None
)
@@ -1269,25 +1317,17 @@ class CudaGraphRunner:
attn_backend = self.model_runner.decode_attn_backend_group[stream_idx]
else:
attn_backend = self.attn_backend
# FIXME: implicit channel for backends (dsv4) that need forward_batch
# in replay metadata prep. Should become a real param on the interface.
attn_backend._replay_forward_batch = forward_batch
seq_lens_sum_arg = (
None
if forward_batch.seq_lens_sum is None
else forward_batch.seq_lens_sum + (bs - raw_bs) * self.seq_len_fill_value
fb_view = build_replay_fb_view(
forward_batch=forward_batch,
buffers=buffers,
bs=bs,
raw_bs=raw_bs,
num_tokens=bs * self.num_tokens_per_bs,
seq_len_fill_value=self.seq_len_fill_value,
capture_forward_mode=self.capture_forward_mode,
is_encoder_decoder=self.is_encoder_decoder,
)
attn_backend.init_forward_metadata_replay_cuda_graph(
bs,
buffers.req_pool_indices[:bs],
buffers.seq_lens[:bs],
seq_lens_sum_arg,
buffers.encoder_lens[:bs] if self.is_encoder_decoder else None,
self.capture_forward_mode,
forward_batch.spec_info,
seq_lens_cpu=buffers.seq_lens_cpu[:bs],
)
attn_backend._replay_forward_batch = None
attn_backend.init_forward_metadata_out_graph(fb_view)
# Store fields
self.raw_bs = raw_bs
@@ -1233,6 +1233,46 @@ def enable_num_token_non_padded():
return get_moe_expert_parallel_world_size() > 1
def build_inner_fb_view(
forward_batch: ForwardBatch,
*,
bs: int,
forward_mode: ForwardMode,
encoder_lens: Optional[torch.Tensor] = None,
):
"""Build a ForwardBatch-like view for MultiStep draft wrapper dispatch.
MultiStep draft wrappers (FlashInferMultiStepDraftBackend,
AiterMultiStepDraftBackend, TritonMultiStepDraftBackend, etc.) need
to dispatch to per-step inner backends'
:py:meth:`AttentionBackend.init_forward_metadata_out_graph` with an
overridden ``forward_mode`` (typically pinned to ``DECODE``) and
sometimes overridden ``encoder_lens``. The result is a thin
namespace mirroring just the fields backend init reads, avoiding
the cost of allocating a real ``ForwardBatch``.
``actual_forward_mode`` carries the original runtime
``forward_batch.forward_mode`` (e.g., spec-decode draft) so backends
that check it for IDLE substitution (DSV4) see the unaltered value.
"""
from types import SimpleNamespace
return SimpleNamespace(
batch_size=bs,
forward_mode=forward_mode,
actual_forward_mode=forward_batch.forward_mode,
input_ids=getattr(forward_batch, "input_ids", None),
positions=getattr(forward_batch, "positions", None),
req_pool_indices=forward_batch.req_pool_indices,
seq_lens=forward_batch.seq_lens,
seq_lens_sum=forward_batch.seq_lens_sum,
seq_lens_cpu=forward_batch.seq_lens_cpu,
encoder_lens=encoder_lens,
out_cache_loc=getattr(forward_batch, "out_cache_loc", None),
spec_info=forward_batch.spec_info,
)
class PPProxyTensors:
# adapted from https://github.com/vllm-project/vllm/blob/d14e98d924724b284dc5eaf8070d935e214e50c0/vllm/sequence.py#L1103
tensors: Dict[str, torch.Tensor]
@@ -419,7 +419,6 @@ class PiecewiseCudaGraphRunner:
return_pooled_hidden_states=self.capture_return_pooled_hidden_states,
)
# Attention backend
self.model_runner.attn_backend.init_forward_metadata(forward_batch)
forward_batch.dp_local_start_pos = forward_batch.dp_local_num_tokens = None
set_dp_buffer_len(None, num_tokens, forward_batch.dp_padding_mode.is_max_len())
@@ -798,7 +797,6 @@ class PiecewiseCudaGraphRunner:
self.moe_fusions,
dsa_indexers=self.dsa_indexers,
):
# Due to the dispatch kernel for MLA model, we init the metadata with original forward_batch
self.model_runner.attn_backend.init_forward_metadata(forward_batch)
output = self.model_runner.model.forward(
static_forward_batch.input_ids,
-3
View File
@@ -1582,9 +1582,6 @@ class DeepseekV4Model(nn.Module):
for _attr in ("freqs_cis_c4", "freqs_cis_c128"):
if hasattr(forward_batch, _attr):
delattr(forward_batch, _attr)
# Upgrade lazy raw metadata on the main stream once before any layer
# forks alt-streams; later per-layer calls become no-ops.
get_attn_backend()._maybe_upgrade_forward_metadata()
use_fused = self.use_fused_mhc_post_pre
prev_residual, prev_post, prev_comb = None, None, None
@@ -373,6 +373,8 @@ class EAGLEDraftCudaGraphRunner:
if self.model_runner.is_hybrid_swa:
self.model_runner.token_to_kv_pool.invalidate_loc_cache()
self.draft_attn_backend.init_forward_metadata_in_graph(forward_batch)
forward_batch.dp_local_start_pos = forward_batch.dp_local_num_tokens = None
set_dp_buffer_len(
global_dp_buffer_len,
@@ -392,8 +394,8 @@ class EAGLEDraftCudaGraphRunner:
return ret
with forward_context(ForwardContext(attn_backend=self.draft_attn_backend)):
self.draft_attn_backend.init_forward_metadata_capture_cuda_graph(
forward_batch
self.draft_attn_backend.init_forward_metadata_out_graph(
forward_batch, in_capture=True
)
self.deepep_adapter.capture(is_extend_in_batch=False)
self._capture_init(run_once)
@@ -509,9 +511,8 @@ class EAGLEDraftCudaGraphRunner:
buffers.seq_lens_cpu[:raw_bs].copy_(forward_batch.seq_lens_cpu)
forward_batch.seq_lens_cpu = buffers.seq_lens_cpu[:bs]
self.draft_attn_backend.init_forward_metadata_replay_cuda_graph(
forward_batch, bs
)
# forward_batch.batch_size was overwritten to bs above when padding.
self.draft_attn_backend.init_forward_metadata_out_graph(forward_batch)
self.raw_bs = raw_bs
self.bs = bs
# TODO: The forward_batch.seq_len_sum might need to be updated to reflect the padding in the cuda graph
@@ -423,14 +423,8 @@ class EAGLEDraftExtendCudaGraphRunner:
with forward_context(
ForwardContext(attn_backend=self.draft_extend_attn_backend)
):
self.draft_extend_attn_backend.init_forward_metadata_capture_cuda_graph(
bs=bs,
num_tokens=num_tokens,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
encoder_lens=None,
forward_mode=self.forward_mode,
spec_info=spec_info,
self.draft_extend_attn_backend.init_forward_metadata_out_graph(
forward_batch, in_capture=True
)
self.deepep_adapter.capture(is_extend_in_batch=True)
@@ -542,19 +536,24 @@ class EAGLEDraftExtendCudaGraphRunner:
forward_batch.spec_info.num_correct_drafts = buffers.num_correct_drafts[:bs]
forward_batch.spec_info.num_accept_tokens = buffers.num_accept_tokens[:bs]
from types import SimpleNamespace
seq_lens_sum = forward_batch.seq_lens_sum
if seq_lens_sum is not None:
seq_lens_sum = seq_lens_sum + (bs - raw_bs) * self.seq_len_fill_value
self.draft_extend_attn_backend.init_forward_metadata_replay_cuda_graph(
bs=bs,
fb_view = SimpleNamespace(
batch_size=bs,
forward_mode=self.forward_mode,
input_ids=getattr(forward_batch, "input_ids", None),
req_pool_indices=buffers.req_pool_indices,
seq_lens=buffers.seq_lens,
seq_lens_sum=seq_lens_sum,
encoder_lens=None,
forward_mode=self.forward_mode,
spec_info=forward_batch.spec_info,
seq_lens_cpu=buffers.seq_lens_cpu,
encoder_lens=None,
out_cache_loc=forward_batch.out_cache_loc,
spec_info=forward_batch.spec_info,
)
self.draft_extend_attn_backend.init_forward_metadata_out_graph(fb_view)
# Replay
self.raw_bs = raw_bs
@@ -288,34 +288,33 @@ class FrozenKVMTPWorker(TpModelWorker):
self, forward_batch: ForwardBatch
) -> None:
with self._frozen_kv_target_view(forward_batch):
self.draft_attn_backend.init_forward_metadata_capture_cuda_graph(
forward_batch.batch_size,
forward_batch.positions.numel(),
forward_batch.req_pool_indices,
forward_batch.seq_lens,
encoder_lens=None,
forward_mode=ForwardMode.DECODE,
spec_info=None,
self.draft_attn_backend.init_forward_metadata_out_graph(
forward_batch, in_capture=True
)
def _init_frozen_kv_metadata_replay_cuda_graph(
self, forward_batch: ForwardBatch, bs: int, seq_lens_sum: int
) -> None:
from types import SimpleNamespace
fb_view = SimpleNamespace(
batch_size=bs,
forward_mode=ForwardMode.DECODE,
input_ids=getattr(forward_batch, "input_ids", None),
req_pool_indices=forward_batch.req_pool_indices[:bs],
seq_lens=forward_batch.seq_lens[:bs],
seq_lens_sum=seq_lens_sum,
seq_lens_cpu=(
forward_batch.seq_lens_cpu[:bs]
if forward_batch.seq_lens_cpu is not None
else None
),
encoder_lens=None,
out_cache_loc=getattr(forward_batch, "out_cache_loc", None),
spec_info=None,
)
with self._frozen_kv_target_view(forward_batch):
self.draft_attn_backend.init_forward_metadata_replay_cuda_graph(
bs,
forward_batch.req_pool_indices[:bs],
forward_batch.seq_lens[:bs],
seq_lens_sum,
encoder_lens=None,
forward_mode=ForwardMode.DECODE,
spec_info=None,
seq_lens_cpu=(
forward_batch.seq_lens_cpu[:bs]
if forward_batch.seq_lens_cpu is not None
else None
),
)
self.draft_attn_backend.init_forward_metadata_out_graph(fb_view)
def init_cuda_graphs(self) -> None:
if self.server_args.disable_cuda_graph or self.speculative_num_steps <= 1:
@@ -485,15 +485,7 @@ class MultiLayerEagleDraftExtendCudaGraphRunner:
return ret
with forward_context(ForwardContext(attn_backend=attn_backend)):
attn_backend.init_forward_metadata_capture_cuda_graph(
bs=bs,
num_tokens=num_tokens,
req_pool_indices=forward_batch.req_pool_indices,
seq_lens=forward_batch.seq_lens,
encoder_lens=None,
forward_mode=self.forward_mode,
spec_info=forward_batch.spec_info,
)
attn_backend.init_forward_metadata_out_graph(forward_batch, in_capture=True)
self.deepep_adapter.capture(is_extend_in_batch=True)
self._capture_init(run_once)
out = self._capture_graph(
@@ -572,19 +564,24 @@ class MultiLayerEagleDraftExtendCudaGraphRunner:
forward_batch.spec_info.positions = buffers.positions[:num_tokens]
forward_batch.spec_info.extend_seq_lens_tensor = buffers.extend_seq_lens[:bs]
self.eagle_worker.draft_extend_attn_backend_list[
self.step
].init_forward_metadata_replay_cuda_graph(
bs=bs,
from types import SimpleNamespace
fb_view = SimpleNamespace(
batch_size=bs,
forward_mode=self.forward_mode,
input_ids=getattr(forward_batch, "input_ids", None),
req_pool_indices=buffers.req_pool_indices,
seq_lens=buffers.seq_lens,
seq_lens_sum=forward_batch.seq_lens_sum
+ (bs - raw_bs) * self.seq_len_fill_value,
encoder_lens=None,
forward_mode=self.forward_mode,
spec_info=forward_batch.spec_info,
seq_lens_cpu=buffers.seq_lens_cpu,
encoder_lens=None,
out_cache_loc=forward_batch.out_cache_loc,
spec_info=forward_batch.spec_info,
)
self.eagle_worker.draft_extend_attn_backend_list[
self.step
].init_forward_metadata_out_graph(fb_view)
# Replay
self.raw_bs = raw_bs
@@ -1079,7 +1079,6 @@ def _seed_c4_if_needed(fixture: DSV4AttentionFixture) -> None:
compress_ratios.
"""
if fixture.case.compress_ratio == 4:
fixture.backend._maybe_upgrade_forward_metadata()
_seed_c4_sparse_indices(fixture, num_entries=_DSV4_EXTRA_ENTRIES)
@@ -1273,7 +1272,6 @@ def _pure_torch_dsv4_combined_reference(
# `c4_sparse_page_indices` back to all -1 on the next upgrade) — the
# reference must observe the same seeded indices the backend forward saw.
_seed_c4_if_needed(fixture)
fixture.backend._maybe_upgrade_forward_metadata()
md = fixture.backend.forward_metadata.core_metadata
runner = fixture.runner
max_context_len = runner.req_to_token_pool.req_to_token.shape[1]
@@ -1554,9 +1552,6 @@ def run_dsv4_compress_attention_case(
q_input, _ = fixture.actual_module.project(fixture.input_hidden)
with torch.no_grad(), forward_context(ForwardContext(attn_backend=fixture.backend)):
fixture.backend.init_forward_metadata(fixture.forward_batch)
# Trigger lazy upgrade so we can patch the metadata that the smoke
# case relies on (specifically c4_sparse_page_indices).
fixture.backend._maybe_upgrade_forward_metadata()
if case.compress_ratio == 4:
_seed_c4_sparse_indices(fixture, num_entries=extra_entries)
actual = fixture.backend.forward(
@@ -291,36 +291,31 @@ def _init_cuda_graph_capture_metadata(backend, capture_batch_size: int, batch):
max_bs=capture_batch_size,
max_num_tokens=batch.input_ids.numel(),
)
backend.init_forward_metadata_capture_cuda_graph(
bs=capture_batch_size,
num_tokens=batch.input_ids.numel(),
req_pool_indices=batch.req_pool_indices,
seq_lens=batch.seq_lens,
encoder_lens=batch.encoder_lens,
forward_mode=batch.forward_mode,
spec_info=batch.spec_info,
)
backend.init_forward_metadata_out_graph(batch, in_capture=True)
backend.init_forward_metadata_in_graph(batch)
def _init_cuda_graph_replay_metadata(backend, capture_batch_size: int, batch):
# Some backends (e.g., `DeepseekV4AttnBackend`) read out-of-band attributes
# off the backend during replay metadata init — production wires this in
# `sglang/srt/model_executor/cuda_graph_runner.py:1234`. Mirror that
# contract so backends that don't use it just store-and-clear the field.
backend._replay_forward_batch = batch
try:
backend.init_forward_metadata_replay_cuda_graph(
bs=capture_batch_size,
req_pool_indices=batch.req_pool_indices,
seq_lens=batch.seq_lens,
seq_lens_sum=batch.seq_lens_sum,
encoder_lens=batch.encoder_lens,
forward_mode=batch.forward_mode,
spec_info=batch.spec_info,
seq_lens_cpu=batch.seq_lens_cpu,
)
finally:
backend._replay_forward_batch = None
from types import SimpleNamespace
fb_view = SimpleNamespace(
batch_size=capture_batch_size,
forward_mode=batch.forward_mode,
actual_forward_mode=batch.forward_mode,
input_ids=batch.input_ids,
positions=getattr(batch, "positions", None),
req_pool_indices=batch.req_pool_indices,
seq_lens=batch.seq_lens,
seq_lens_sum=batch.seq_lens_sum,
seq_lens_cpu=batch.seq_lens_cpu,
encoder_lens=batch.encoder_lens,
out_cache_loc=getattr(batch, "out_cache_loc", None),
spec_info=batch.spec_info,
)
backend.init_forward_metadata_out_graph(fb_view)
# No real cuda graph here, so run the in-graph step explicitly to produce
# the Full metadata the forward path expects (no-op for non-DSV4).
backend.init_forward_metadata_in_graph(fb_view)
# Best-effort metadata-shape sanity check — catches negative kv_lens and
# non-monotonic indptr that would otherwise leave real-row output correct
# but corrupt padded-row scratch state. See `metadata_invariants.py`.
@@ -1022,13 +1022,6 @@ class EagleDraftExtendCudaGraphRunnerAdapter:
make_forward_batch: Callable[
[Any, Any, Any, EagleDraftRunnerSettings], ForwardBatch
]
# Optional hook invoked with `(draft_extend_attn_backend, batch)` right
# before `graph_runner.replay(batch)`. DSV4 needs this to set the
# out-of-band `_replay_forward_batch` attribute that
# `DeepseekV4AttnBackend.init_forward_metadata_replay_cuda_graph` reads
# (the multi-step DECODE wrapper sets it internally, but the single-
# backend DRAFT_EXTEND path does not).
pre_replay: Callable[[Any, ForwardBatch], None] = None
check_case: Callable[[Any, EagleDraftRunnerSettings], None] = (
lambda _case, _settings: None
)
@@ -1268,12 +1261,7 @@ def run_eagle_draft_extend_cuda_graph_runner_case(
adapter.prepare_replay_state(graph_fixture, case, draft_inputs, settings)
testcase.assertTrue(graph_runner.can_run(graph_batch))
if adapter.pre_replay is not None:
adapter.pre_replay(graph_backend, graph_batch)
actual = graph_runner.replay(graph_batch)
if adapter.pre_replay is not None:
# Best-effort cleanup of any out-of-band state pre_replay set.
adapter.pre_replay(graph_backend, None)
adapter.assert_outputs_close(actual, expected, settings)
finally:
_reset_cuda_graph_test_buffers()
@@ -1934,21 +1922,6 @@ def _dsv4_assert_draft_extend_outputs_close(actual, expected, settings) -> None:
)
def _dsv4_draft_extend_pre_replay(
draft_extend_attn_backend,
batch: ForwardBatch | None,
) -> None:
"""Set/clear the out-of-band `_replay_forward_batch` attribute that
`DeepseekV4AttnBackend.init_forward_metadata_replay_cuda_graph` reads.
The DSV4 multi-step DECODE wrapper sets this internally
(`deepseek_v4_backend.py:1231,1242`), but the single-backend DRAFT_EXTEND
path used by `_create_dsv4_prefill_backend` does not. Set before
`replay()` and clear afterwards to mimic the multi-step pattern.
"""
draft_extend_attn_backend._replay_forward_batch = batch
def run_dsv4_eagle_draft_extend_cuda_graph_runner_case(
testcase,
case: DSV4AttentionCase,
@@ -2001,7 +1974,6 @@ def run_dsv4_eagle_draft_extend_cuda_graph_runner_case(
make_draft_inputs=_make_dsv4_draft_extend_inputs,
prepare_replay_state=_prepare_dsv4_draft_extend_replay_state,
make_forward_batch=_make_dsv4_eagle_draft_extend_forward_batch,
pre_replay=_dsv4_draft_extend_pre_replay,
assert_outputs_close=_dsv4_assert_draft_extend_outputs_close,
)
run_eagle_draft_extend_cuda_graph_runner_case(