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