[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:
co-authored by
Claude Opus 4.7
parent
6c69756fa8
commit
99da43b900
@@ -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,
|
||||
|
||||
+7
-3
@@ -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"
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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(
|
||||
|
||||
+22
-27
@@ -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`.
|
||||
|
||||
-28
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user