[refactor] Unify CUDA graph runner input buffers behind CudaGraphBufferRegistry (#26742)

Co-authored-by: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
This commit is contained in:
Cheng Wan
2026-06-03 10:54:10 -07:00
committed by GitHub
co-authored by Claude Opus 4.8
parent b0f78bef97
commit 45604a0f4a
6 changed files with 2242 additions and 249 deletions
@@ -169,6 +169,9 @@ class BreakableCudaGraphRunner:
def _init_buffers(self, model_runner): def _init_buffers(self, model_runner):
"""Initialize input buffers.""" """Initialize input buffers."""
from sglang.srt.model_executor.cuda_graph_buffer_registry import (
build_prefill_registry,
)
from sglang.srt.model_executor.piecewise_cuda_graph_runner import ( from sglang.srt.model_executor.piecewise_cuda_graph_runner import (
PrefillInputBuffers, PrefillInputBuffers,
) )
@@ -215,6 +218,21 @@ class BreakableCudaGraphRunner:
) )
self.buffers.share_buffers() self.buffers.share_buffers()
# Token-axis FB-shared slot registry adopting the PrefillInputBuffers
# storage. Breakable has no mamba track and bs is not padded here, so
# there are no bs-axis slots (max_bs is unused).
self.buffer_registry = build_prefill_registry(
device=self.device,
max_bs=1,
max_num_token=self.max_num_tokens,
cache_loc_dtype=torch.int64 if not is_npu() else torch.int32,
is_multimodal=self.is_multimodal,
hidden_size=model_runner.model_config.hidden_size,
embed_dtype=model_runner.dtype,
enable_mamba_track=False,
source=self.buffers,
)
@torch.no_grad() @torch.no_grad()
def _run_forward(self, forward_batch, num_tokens): def _run_forward(self, forward_batch, num_tokens):
"""Run layer-stack forward with proper context. """Run layer-stack forward with proper context.
@@ -271,8 +289,12 @@ class BreakableCudaGraphRunner:
hidden_states=self.static_draft_hidden_states[:num_tokens], hidden_states=self.static_draft_hidden_states[:num_tokens],
) )
buffers = self.buffers registry = self.buffer_registry
bs = 1 bs = 1
def _slot(name):
return registry.get_slot(name).slice_for(bs, num_tokens)
with torch.device(self.device): with torch.device(self.device):
seq_lens = torch.full((bs,), num_tokens, dtype=torch.int64) seq_lens = torch.full((bs,), num_tokens, dtype=torch.int64)
extend_seq_lens = torch.full((bs,), num_tokens, dtype=torch.int64) extend_seq_lens = torch.full((bs,), num_tokens, dtype=torch.int64)
@@ -284,16 +306,16 @@ class BreakableCudaGraphRunner:
return ForwardBatch( return ForwardBatch(
forward_mode=ForwardMode.EXTEND, forward_mode=ForwardMode.EXTEND,
batch_size=bs, batch_size=bs,
input_ids=buffers.input_ids[:num_tokens], input_ids=_slot("input_ids"),
input_embeds=( input_embeds=(
buffers.input_embeds[:num_tokens] if self.is_multimodal else None _slot("input_embeds") if registry.has_slot("input_embeds") else None
), ),
req_pool_indices=req_pool_indices, req_pool_indices=req_pool_indices,
seq_lens=seq_lens, seq_lens=seq_lens,
next_token_logits_buffer=None, next_token_logits_buffer=None,
orig_seq_lens=orig_seq_lens, orig_seq_lens=orig_seq_lens,
seq_lens_cpu=torch.tensor([num_tokens], device="cpu"), seq_lens_cpu=torch.tensor([num_tokens], device="cpu"),
out_cache_loc=buffers.out_cache_loc[:num_tokens], out_cache_loc=_slot("out_cache_loc"),
seq_lens_sum=num_tokens, seq_lens_sum=num_tokens,
mamba_track_indices=None, mamba_track_indices=None,
mamba_track_mask=None, mamba_track_mask=None,
@@ -307,13 +329,15 @@ class BreakableCudaGraphRunner:
extend_prefix_lens_cpu=torch.tensor([0], device="cpu"), extend_prefix_lens_cpu=torch.tensor([0], device="cpu"),
extend_seq_lens_cpu=torch.tensor([num_tokens], device="cpu"), extend_seq_lens_cpu=torch.tensor([num_tokens], device="cpu"),
extend_logprob_start_lens_cpu=torch.tensor([num_tokens], device="cpu"), extend_logprob_start_lens_cpu=torch.tensor([num_tokens], device="cpu"),
positions=buffers.positions[:num_tokens], positions=_slot("positions"),
global_num_tokens_gpu=None, global_num_tokens_gpu=None,
global_num_tokens_for_logprob_gpu=None, global_num_tokens_for_logprob_gpu=None,
dp_padding_mode=DpPaddingMode.get_default_mode_in_cuda_graph(), dp_padding_mode=DpPaddingMode.get_default_mode_in_cuda_graph(),
global_dp_buffer_len=None, global_dp_buffer_len=None,
mrope_positions=( mrope_positions=(
buffers.mrope_positions[:, :num_tokens] if self.is_multimodal else None _slot("mrope_positions")
if registry.has_slot("mrope_positions")
else None
), ),
spec_algorithm=None, spec_algorithm=None,
spec_info=spec_info, spec_info=spec_info,
@@ -446,9 +470,9 @@ class BreakableCudaGraphRunner:
if self.use_input_embeds: if self.use_input_embeds:
if ie is None: if ie is None:
raise ValueError("BCG replay expects input_embeds but got None") raise ValueError("BCG replay expects input_embeds but got None")
self.buffers.input_embeds[:static_num_tokens].copy_( self.buffer_registry.get_slot("input_embeds").slice_for(
ie[:static_num_tokens] 1, static_num_tokens
) ).copy_(ie[:static_num_tokens])
else: else:
if ie is not None: if ie is not None:
raise ValueError( raise ValueError(
@@ -0,0 +1,848 @@
"""FB-shared slot registry for the CUDA graph forward paths.
``CudaGraphBufferRegistry`` is the ForwardBatch → graph-resident buffer mirror
used by capture / replay. It replaces the per-runner ``DecodeInputBuffers`` /
``PrefillInputBuffers`` dataclasses and their hand-written
``populate_from_forward_batch`` methods with a single ``GraphSlot``-driven
registry.
Backend-private buffers (kernel workspaces, derived page tables, etc.) stay
on ``AttentionBackend.cuda_graph_*`` — the registry only owns FB-shared
slots (FB attribute name maps 1:1 to slot name).
"""
from __future__ import annotations
from dataclasses import dataclass, field
from enum import Enum
from typing import TYPE_CHECKING, Any, Callable, Dict, List, Optional, Tuple
import torch
from sglang.srt.model_executor.input_buffers import share_input_buffer
if TYPE_CHECKING:
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
_has_foreach_copy = hasattr(torch, "_foreach_copy_")
def _grouped_foreach_copy_(dsts: List[torch.Tensor], srcs: List[torch.Tensor]) -> None:
"""Call torch._foreach_copy_ grouped by (dst_dtype, src_dtype) pairs
(a single foreach call requires a uniform dtype pair)."""
def _foreach_copy(
group_dsts: List[torch.Tensor], group_srcs: List[torch.Tensor]
) -> None:
if _has_foreach_copy:
torch._foreach_copy_(group_dsts, group_srcs)
else:
for dst, src in zip(group_dsts, group_srcs):
dst.copy_(src)
groups: Dict[Tuple[torch.dtype, torch.dtype], Tuple[List, List]] = {}
for dst, src in zip(dsts, srcs):
key = (dst.dtype, src.dtype)
if key not in groups:
groups[key] = ([], [])
groups[key][0].append(dst)
groups[key][1].append(src)
for group_dsts, group_srcs in groups.values():
_foreach_copy(group_dsts, group_srcs)
class PaddingPolicy(Enum):
"""How to handle ``raw_n < padded_n`` for a slot.
KEEP_PAD — Leave the padded region as-is (caller proves the
padded tail will not be read).
FILL_SENTINEL — Reset the padded region to ``slot.pad_value`` before
copy (e.g. ``seq_lens`` filled with
``seq_len_fill_value``).
ZERO — Reset the padded region to ``0`` (e.g.
``out_cache_loc`` / ``req_pool_indices`` — padded
rows must point at slot 0 so dummy attention reads
land harmlessly).
FOREACH_COPY — Always copy ``raw_n`` from src; padded region is
left as whatever the previous replay (or the init
zeros) wrote. Caller is responsible for proving
safety.
FILL_ONCE — Fill the whole buffer to ``pad_value`` once at alloc;
never reset per iter (e.g. ``encoder_lens`` init to
``encoder_len_fill_value``, copied head-only with the
tail kept).
"""
KEEP_PAD = "keep_pad"
FILL_SENTINEL = "fill_sentinel"
ZERO = "zero"
FOREACH_COPY = "foreach_copy"
FILL_ONCE = "fill_once"
@dataclass
class FillContext:
"""Per-iteration shape context passed to ``GraphSlot.post_fill``.
Carries both the bs-axis and tokens-axis raw/padded counts so a hook can
derive values regardless of its own slot's axis — e.g. the padded token
count (``padded_num_tokens`` == padded_bs * num_tokens_per_bs), which the
global-num-tokens fill and the local-num-token-non-padded transform need.
"""
raw_bs: int
padded_bs: int
raw_num_tokens: int
padded_num_tokens: int
# Side inputs that are not ForwardBatch attributes but are needed by a
# slot's source_fn — e.g. the pipeline-parallel proxy tensors, which the
# replay path receives as a separate argument rather than off the FB.
pp_proxy_tensors: Optional[Any] = None
@dataclass
class GraphSlot:
"""A single FB-mirrored buffer.
Each slot mirrors one ``ForwardBatch`` attribute. ``name`` MUST match
the FB attribute name so ``fill_from`` can ``getattr(fb, name)`` and
``extract_buffer`` can ``setattr`` the view back into a FB replace.
Fields:
name — the FB attribute name mirrored by this slot.
shape_fn — ``(max_bs, max_num_tokens) -> shape`` callable
used at ``register_slot`` time to allocate the
physical buffer.
dtype — buffer dtype.
axis — ``"bs"`` (slot is sliced ``[:bs]``) or
``"tokens"`` (sliced ``[:num_tokens]``) or
``"none"`` (no slicing — full buffer always
exposed; used for scalar buffers and global
counters).
device — buffer device. ``None`` means use registry
default; can be ``"cpu"`` for slots like
``seq_lens_cpu`` that must live on host.
padding_policy — see ``PaddingPolicy``.
pad_value — sentinel for ``FILL_SENTINEL``.
enabled — runtime gate; disabled slots are not allocated
and skipped during fill / extract.
copy_from_fb — when ``True`` (default), ``fill_from`` copies the
same-named FB tensor into the buffer head. Set
``False`` for computed slots whose value is not a
straight FB copy (e.g. ``global_num_tokens_*``,
filled by a ``post_fill`` instead).
post_fill — optional ``(buffer, forward_batch, FillContext)
-> None`` hook run after the grouped copy. Used for
compute-then-write slots (local-num-token-non-padded
transform, global-num-tokens fill).
slice_fn — optional ``(buffer, padded_n) -> Tensor``
override for slots with non-trivial slicing
(e.g. ``mrope_positions`` shape ``[3, T]`` is
sliced on axis 1 not 0).
source_fn — optional ``(forward_batch, FillContext) -> Tensor |
None`` override for the copy *source*. When set,
``fill_from`` copies ``source_fn(fb, ctx)`` (instead of
the same-named FB attribute) into
``buffer[:src.shape[0]]`` — a source-length slice for
structured / side-sourced fields whose data lives on a
nested FB dataclass (``ngram_embedding_info.*``) or an
out-of-band argument (``pp_proxy_tensors``, carried on
``FillContext``). Returning ``None`` skips the copy for
that iteration. Such slots use dotted names and are
skipped by ``extract_buffer``.
"""
name: str
shape_fn: Callable[[int, int], Tuple[int, ...]]
dtype: torch.dtype
axis: str = "tokens"
device: Optional[torch.device] = None
padding_policy: PaddingPolicy = PaddingPolicy.FOREACH_COPY
pad_value: Optional[Any] = None
enabled: bool = True
copy_from_fb: bool = True
post_fill: Optional[
Callable[[torch.Tensor, "ForwardBatch", "FillContext"], None]
] = None
slice_fn: Optional[Callable[[torch.Tensor, int], torch.Tensor]] = None
source_fn: Optional[
Callable[["ForwardBatch", "FillContext"], Optional[torch.Tensor]]
] = None
# runtime
buffer: Optional[torch.Tensor] = field(default=None, repr=False)
def __post_init__(self) -> None:
if self.axis not in ("bs", "tokens", "none"):
raise ValueError(
f"GraphSlot {self.name!r}: axis must be one of "
f"'bs'/'tokens'/'none', got {self.axis!r}"
)
def _padded_n(self, padded_bs: int, padded_num_tokens: int) -> int:
if self.axis == "bs":
return padded_bs
if self.axis == "tokens":
return padded_num_tokens
# axis == "none": no slicing
return self.buffer.shape[0] if self.buffer is not None else 0
def _raw_n(self, raw_bs: int, raw_num_tokens: int) -> int:
if self.axis == "bs":
return raw_bs
if self.axis == "tokens":
return raw_num_tokens
return self.buffer.shape[0] if self.buffer is not None else 0
def slice_for(self, padded_bs: int, padded_num_tokens: int) -> torch.Tensor:
"""Return the ``[:padded_n]`` slice of the buffer consumed by callers.
This truncates the (full-length) buffer to the active region for the
current iteration — it is a slice, not a tensor reshape.
"""
if self.buffer is None:
raise RuntimeError(f"GraphSlot {self.name!r}: buffer not allocated")
if self.slice_fn is not None:
return self.slice_fn(
self.buffer, self._padded_n(padded_bs, padded_num_tokens)
)
if self.axis == "none":
return self.buffer
return self.buffer[: self._padded_n(padded_bs, padded_num_tokens)]
def reset_padding(self, raw_n: int, padded_n: int) -> None:
"""Reset the padded tail according to ``padding_policy``."""
if self.buffer is None or raw_n >= padded_n:
return
if self.padding_policy in (
PaddingPolicy.KEEP_PAD,
PaddingPolicy.FOREACH_COPY,
PaddingPolicy.FILL_ONCE,
):
return
# slice_fn governs non-trivial layouts (e.g. mrope_positions [3, T]);
# the pad region is the same axis the slot exposes via slice_for().
if self.slice_fn is not None:
# slice_fn returns the [:padded_n] portion already; we need the
# tail [raw_n:padded_n]. We rely on slice_fn slicing the same
# axis used by slice_for(): take the padded slice first, then index
# the tail with the standard slice on axis 0 of the result.
padded_slice = self.slice_fn(self.buffer, padded_n)
tail = (
padded_slice[..., raw_n:padded_n]
if padded_slice.dim() > 1
else padded_slice[raw_n:padded_n]
)
else:
tail = self.buffer[raw_n:padded_n]
if self.padding_policy == PaddingPolicy.FILL_SENTINEL:
if self.pad_value is None:
raise RuntimeError(
f"GraphSlot {self.name!r}: FILL_SENTINEL requires pad_value"
)
tail.fill_(self.pad_value)
elif self.padding_policy == PaddingPolicy.ZERO:
tail.zero_()
class CudaGraphBufferRegistry:
"""FB → graph-resident buffer mirror, shared across eager / capture / replay.
The registry holds a dict of ``GraphSlot`` instances, each mirroring
one ``ForwardBatch`` attribute. Slots are registered up-front (during
runner init), allocated at ``register_slot``, then filled per-iter via
``fill_from(fb, ...)`` and consumed via ``extract_buffer(template) ->
ForwardBatch``. ``fill_from`` issues plain D2D copies on the caller's
current stream; cross-stream correctness (stream handoff) is handled by
the runners, not here.
Backend-private buffers (kernel workspace, derived page tables) are
NOT managed here — backends keep them on ``self.cuda_graph_*`` and
allocate via ``AttentionBackend.init_cuda_graph_state(...)``.
Usage::
registry = CudaGraphBufferRegistry(device=..., max_bs=..., max_num_tokens=...)
registry.register_slot(GraphSlot(name="input_ids", ...))
registry.register_slot(GraphSlot(name="seq_lens",
padding_policy=PaddingPolicy.FILL_SENTINEL,
pad_value=seq_len_fill_value, ...))
# per-iter:
registry.fill_from(fb, raw_bs=..., padded_bs=..., raw_num_tokens=...,
padded_num_tokens=...)
fb_view = registry.extract_buffer(padded_bs=..., padded_num_tokens=...,
forward_batch_template=fb)
attn_backend.init_forward_metadata(fb_view)
model.forward(fb_view.input_ids, fb_view.positions, fb_view)
"""
def __init__(
self,
*,
device: torch.device,
max_bs: int,
max_num_tokens: int,
share_pool: bool = False,
) -> None:
self.device = device
self.max_bs = max_bs
self.max_num_tokens = max_num_tokens
# When True, slot buffers are coalesced by name through the global
# ForwardInputBuffers pool, so a registry can share physical storage
# (and data_ptr) with the legacy DecodeInputBuffers during migration.
self.share_pool = share_pool
self._slots: Dict[str, GraphSlot] = {}
# ---- registration ------------------------------------------------------
def register_slot(
self, slot: GraphSlot, bind: Optional[torch.Tensor] = None
) -> GraphSlot:
"""Register a slot and allocate (or adopt) its physical buffer.
If ``bind`` is given, the slot adopts that existing tensor instead of
allocating a fresh one (and skips the pool / sentinel init — the bound
tensor is assumed already initialized). This lets a registry share
storage with the legacy ``DecodeInputBuffers`` by adopting its fields,
guaranteeing a stable, identical ``data_ptr`` for capture vs replay.
Returns the slot for caller convenience. Re-registering an existing
name raises.
"""
if slot.name in self._slots:
raise ValueError(
f"GraphSlot {slot.name!r} already registered; "
"use enable()/disable() to gate per-iter."
)
if not slot.enabled:
# Even when disabled, keep the spec so callers can introspect
# by name; just don't allocate.
self._slots[slot.name] = slot
return slot
shape = slot.shape_fn(self.max_bs, self.max_num_tokens)
device = slot.device if slot.device is not None else self.device
if bind is not None:
expected = tuple(shape)
if tuple(bind.shape) != expected:
raise ValueError(
f"bind tensor for slot {slot.name!r} has shape "
f"{tuple(bind.shape)}, expected {expected}."
)
if bind.dtype != slot.dtype:
raise ValueError(
f"bind tensor for slot {slot.name!r} has dtype {bind.dtype}, "
f"expected {slot.dtype}."
)
slot.buffer = bind
self._slots[slot.name] = slot
return slot
buffer = torch.zeros(shape, dtype=slot.dtype, device=device)
if self.share_pool:
# Coalesce with any same-named buffer (e.g. the legacy
# DecodeInputBuffers field) so capture and replay see one
# physical allocation with a stable data_ptr.
buffer = share_input_buffer(slot.name, buffer)
if (
slot.padding_policy
in (PaddingPolicy.FILL_SENTINEL, PaddingPolicy.FILL_ONCE)
and slot.pad_value is not None
):
buffer.fill_(slot.pad_value)
slot.buffer = buffer
self._slots[slot.name] = slot
return slot
def has_slot(self, name: str) -> bool:
return name in self._slots and self._slots[name].enabled
def get_slot(self, name: str) -> GraphSlot:
return self._slots[name]
def slot_names(self) -> List[str]:
return [name for name, s in self._slots.items() if s.enabled]
# ---- per-iter ----------------------------------------------------------
def fill_from(
self,
forward_batch: "ForwardBatch",
*,
raw_bs: int,
padded_bs: int,
raw_num_tokens: int,
padded_num_tokens: int,
pp_proxy_tensors: Optional[Any] = None,
) -> None:
"""Copy FB → registry buffers.
Phase 1 — reset the padded tail per slot ``padding_policy``.
Phase 2 — grouped D2D copy of all enabled slots from FB (or from a
slot's ``source_fn`` for structured / side-sourced fields).
Phase 3 — run ``post_fill`` hooks for slots that need
post-copy transforms.
``pp_proxy_tensors`` is the out-of-band pipeline-parallel input; it is
not an FB attribute, so it reaches ``source_fn`` slots via
``FillContext.pp_proxy_tensors``.
Slots whose FB attribute (or ``source_fn`` result) is ``None`` are
silently skipped (the FB doesn't carry that field for the current
request).
"""
ctx = FillContext(
raw_bs=raw_bs,
padded_bs=padded_bs,
raw_num_tokens=raw_num_tokens,
padded_num_tokens=padded_num_tokens,
pp_proxy_tensors=pp_proxy_tensors,
)
# Phase 1: reset padded regions where it matters.
for slot in self._slots.values():
if not slot.enabled or slot.buffer is None:
continue
raw_n = slot._raw_n(raw_bs, raw_num_tokens)
padded_n = slot._padded_n(padded_bs, padded_num_tokens)
slot.reset_padding(raw_n, padded_n)
# Phase 2: collect (dst, src) pairs and dispatch a grouped copy.
gpu_dsts: List[torch.Tensor] = []
gpu_srcs: List[torch.Tensor] = []
cpu_dsts: List[torch.Tensor] = []
cpu_srcs: List[torch.Tensor] = []
for slot in self._slots.values():
if not slot.enabled or slot.buffer is None or not slot.copy_from_fb:
continue
if slot.source_fn is not None:
# Structured / side-sourced slot: source comes from a nested FB
# dataclass or an out-of-band input, and the copy is sliced to
# the source's own length rather than a bs/tokens axis.
src = slot.source_fn(forward_batch, ctx)
if src is None:
continue
dst = slot.buffer[: src.shape[0]]
else:
src = getattr(forward_batch, slot.name, None)
if src is None:
continue
if not isinstance(src, torch.Tensor):
# Non-tensor FB fields (e.g. dicts, dataclasses) are not
# auto-copied — caller handles via source_fn or post_fill.
continue
raw_n = slot._raw_n(raw_bs, raw_num_tokens)
if slot.slice_fn is not None:
dst = slot.slice_fn(slot.buffer, raw_n)
elif slot.axis == "none":
dst = slot.buffer
else:
dst = slot.buffer[:raw_n]
# foreach_copy_ requires same-device tensors per call — bucket
# by device.
if dst.device.type == "cpu":
cpu_dsts.append(dst)
cpu_srcs.append(src)
else:
gpu_dsts.append(dst)
gpu_srcs.append(src)
if gpu_dsts:
_grouped_foreach_copy_(gpu_dsts, gpu_srcs)
for dst, src in zip(cpu_dsts, cpu_srcs):
dst.copy_(src)
# Phase 3: post-fill hooks (compute-then-write slots).
for slot in self._slots.values():
if not slot.enabled or slot.buffer is None or slot.post_fill is None:
continue
slot.post_fill(slot.buffer, forward_batch, ctx)
def extract_buffer(
self,
*,
padded_bs: int,
padded_num_tokens: int,
forward_batch_template: "ForwardBatch",
) -> "ForwardBatch":
"""Return a FB view backed by registry slot buffers.
``forward_batch_template`` provides the non-slot fields
(``forward_mode`` / ``spec_info`` / ``sampling_info`` /
``capture_hidden_mode`` / ``dp_*`` / ``lora_ids`` / ...). Slot
fields are replaced with views into the registry buffers via
``dataclasses.replace`` — the template itself is not mutated.
"""
import dataclasses
replace_kwargs: Dict[str, Any] = {"batch_size": padded_bs}
for slot in self._slots.values():
if not slot.enabled or slot.buffer is None:
continue
# Structured slots use dotted names ("<field>.<sub>") and are not
# top-level FB attributes — their data is consumed in place off the
# adopted backing object, not re-attached to the FB view here.
if "." in slot.name:
continue
replace_kwargs[slot.name] = slot.slice_for(padded_bs, padded_num_tokens)
return dataclasses.replace(forward_batch_template, **replace_kwargs)
def build_decode_registry(
*,
device: torch.device,
max_bs: int,
max_num_token: int,
seq_len_fill_value: int,
cache_loc_dtype: torch.dtype,
enable_mamba_track: bool = False,
is_encoder_decoder: bool = False,
encoder_len_fill_value: int = 0,
enable_num_token_non_padded: bool = False,
require_gathered_buffer: bool = False,
enable_prefill_cp: bool = False,
require_mlp_tp_gather: bool = False,
dp_size: int = 1,
share_pool: bool = True,
source: Optional[Any] = None,
) -> CudaGraphBufferRegistry:
"""Registry mirroring the always-on (+ mamba / mrope) FB-shared decode
buffers, with padding policies matching
``DecodeInputBuffers.populate_from_forward_batch``:
- ``seq_lens`` / ``seq_lens_cpu`` -> FILL_SENTINEL(seq_len_fill_value)
- ``req_pool_indices`` / ``out_cache_loc`` / ``mamba_track_*`` -> ZERO
- ``input_ids`` / ``positions`` / ``mrope_positions`` -> FOREACH_COPY
(head ``[:raw_n]`` is always overwritten by the copy; the old code's
full-buffer ``zero_()`` / ``fill_()`` on ``bs != raw_bs`` is therefore
equivalent to the tail-only reset the policies apply here).
``custom_mask`` / ``next_token_logits_buffer`` / ``input_embeds`` are not
registered here — they are not per-replay FB copies (allocated and written
elsewhere), so the runner keeps owning them.
When ``source`` is given, each slot adopts the same-named tensor off
``source`` (e.g. a ``DecodeInputBuffers``) instead of allocating, so the
registry shares one physical allocation with that object.
"""
reg = CudaGraphBufferRegistry(
device=device,
max_bs=max_bs,
max_num_tokens=max_num_token,
share_pool=share_pool,
)
def _tokens(_bs: int, mt: int) -> Tuple[int, ...]:
return (mt,)
def _bs(bs: int, _mt: int) -> Tuple[int, ...]:
return (bs,)
slots = [
GraphSlot("input_ids", _tokens, torch.int64, axis="tokens"),
GraphSlot("positions", _tokens, torch.int64, axis="tokens"),
GraphSlot(
"out_cache_loc",
_tokens,
cache_loc_dtype,
axis="tokens",
padding_policy=PaddingPolicy.ZERO,
),
GraphSlot(
"req_pool_indices",
_bs,
torch.int64,
axis="bs",
padding_policy=PaddingPolicy.ZERO,
),
GraphSlot(
"seq_lens",
_bs,
torch.int32,
axis="bs",
padding_policy=PaddingPolicy.FILL_SENTINEL,
pad_value=seq_len_fill_value,
),
GraphSlot(
"seq_lens_cpu",
_bs,
torch.int32,
axis="bs",
device=torch.device("cpu"),
padding_policy=PaddingPolicy.FILL_SENTINEL,
pad_value=seq_len_fill_value,
),
GraphSlot(
"mrope_positions",
lambda _bs2, mt: (3, mt),
torch.int64,
axis="tokens",
slice_fn=lambda buf, n: buf[:, :n],
),
]
if enable_mamba_track:
slots.append(
GraphSlot(
"mamba_track_indices",
_bs,
torch.int64,
axis="bs",
padding_policy=PaddingPolicy.ZERO,
)
)
slots.append(
GraphSlot(
"mamba_track_mask",
_bs,
torch.bool,
axis="bs",
padding_policy=PaddingPolicy.ZERO,
)
)
if is_encoder_decoder:
# Initialized once to encoder_len_fill_value, copied head-only, never
# reset per iter — matching the legacy DecodeInputBuffers behavior.
slots.append(
GraphSlot(
"encoder_lens",
_bs,
torch.int32,
axis="bs",
padding_policy=PaddingPolicy.FILL_ONCE,
pad_value=encoder_len_fill_value,
)
)
if enable_num_token_non_padded:
from sglang.srt.model_executor.forward_batch_info import (
compute_local_num_token_non_padded,
)
def _num_token_non_padded_post_fill(buf, fb, ctx):
# Gathered (DP) path overwrites the plain FB copy with this rank's
# local count; the non-gathered path keeps the copied value.
if require_gathered_buffer and not enable_prefill_cp:
buf.copy_(
compute_local_num_token_non_padded(
global_num_token_non_padded=fb.num_token_non_padded,
num_tokens_per_dp=ctx.padded_num_tokens,
)
)
slots.append(
GraphSlot(
"num_token_non_padded",
lambda _bs, _mt: (1,),
torch.int32,
axis="none",
post_fill=_num_token_non_padded_post_fill,
)
)
def _global_num_tokens_post_fill(buf, fb, ctx):
# Filled with the padded token count on the gathered (DP) path; left
# untouched otherwise. Not an FB copy (copy_from_fb=False).
if require_gathered_buffer:
buf.fill_(ctx.padded_num_tokens)
_global_shape = (
(lambda _bs, _mt: (dp_size,))
if require_mlp_tp_gather
else (lambda _bs, _mt: (1,))
)
for _global_name in ("global_num_tokens_gpu", "global_num_tokens_for_logprob_gpu"):
slots.append(
GraphSlot(
_global_name,
_global_shape,
torch.int32,
axis="none",
copy_from_fb=False,
post_fill=_global_num_tokens_post_fill,
)
)
for slot in slots:
bind = None
if source is not None:
bind = getattr(source, slot.name, None)
if bind is None:
raise ValueError(
f"source is missing buffer {slot.name!r} required by the "
"decode registry; cannot adopt."
)
reg.register_slot(slot, bind=bind)
# Structured slots whose backing storage still lives on the source object
# (adopt-only during migration): registered only when the source actually
# carries them. The per-replay copy source is a nested FB dataclass field,
# supplied via source_fn; head is copied (source-length slice), tail kept.
if source is not None:
ngram = getattr(source, "ngram_embedding_info", None)
if ngram is not None:
def _ngram_source(attr):
def _fn(fb, _ctx):
info = getattr(fb, "ngram_embedding_info", None)
return None if info is None else getattr(info, attr)
return _fn
for _attr in ("column_starts", "req_lens"):
backing = getattr(ngram, _attr)
reg.register_slot(
GraphSlot(
name=f"ngram_embedding_info.{_attr}",
shape_fn=lambda _bs, _mt, _s=tuple(backing.shape): _s,
dtype=backing.dtype,
axis="none",
padding_policy=PaddingPolicy.KEEP_PAD,
source_fn=_ngram_source(_attr),
),
bind=backing,
)
# Pipeline-parallel proxy tensors: a dict of per-key buffers, sourced
# from the out-of-band pp input on FillContext rather than the FB.
pp = getattr(source, "pp_proxy_tensors", None)
if pp is not None:
def _pp_source(key):
def _fn(_fb, ctx):
ppx = ctx.pp_proxy_tensors
return None if ppx is None else ppx.tensors[key]
return _fn
for _key, _backing in pp.items():
reg.register_slot(
GraphSlot(
name=f"pp_proxy_tensors.{_key}",
shape_fn=lambda _bs, _mt, _s=tuple(_backing.shape): _s,
dtype=_backing.dtype,
axis="none",
padding_policy=PaddingPolicy.KEEP_PAD,
source_fn=_pp_source(_key),
),
bind=_backing,
)
# KV-canary id buffers (off by default): plain bs-axis FB copies,
# adopt-only when the source carries them. Head [:raw_bs] is copied;
# the tail keeps its init (rids_int 0, bootstrap_room_ids_int -1).
for _cname in ("rids_int", "bootstrap_room_ids_int"):
canary = getattr(source, _cname, None)
if canary is not None:
reg.register_slot(
GraphSlot(
name=_cname,
shape_fn=lambda _bs, _mt, _s=tuple(canary.shape): _s,
dtype=canary.dtype,
axis="bs",
),
bind=canary,
)
return reg
def build_prefill_registry(
*,
device: torch.device,
max_bs: int,
max_num_token: int,
cache_loc_dtype: torch.dtype,
is_multimodal: bool = False,
hidden_size: int = 0,
embed_dtype: Optional[torch.dtype] = None,
enable_mamba_track: bool = False,
share_pool: bool = True,
source: Optional[Any] = None,
) -> CudaGraphBufferRegistry:
"""Registry mirroring the **token-axis** FB-shared buffers for the
piecewise / breakable (prefill) cuda-graph runners.
Padding policies match the inline copy/zero in
``PiecewiseCudaGraphRunner.replay_prepare``: ``input_ids`` / ``positions``
/ ``out_cache_loc`` / ``mrope_positions`` / ``input_embeds`` reset their
padded tail ``[raw_num_tokens:padded_num_tokens]`` to ``0`` (the padded
tokens *are* processed by the graph, so they must be benign), then the head
``[:raw_num_tokens]`` is copied from the FB. ``input_embeds`` is not an FB
copy — the model writes the embeds into it inside the graph — so it is
reset-only (``copy_from_fb=False``). ``mamba_track_*`` are bs-axis copies
with no padding reset (bs is not padded on this path).
When ``source`` is given, each slot adopts the same-named tensor off
``source`` (the ``PrefillInputBuffers``) instead of allocating, so the
registry shares one physical allocation (and ``data_ptr``) with it.
"""
reg = CudaGraphBufferRegistry(
device=device,
max_bs=max_bs,
max_num_tokens=max_num_token,
share_pool=share_pool,
)
def _tokens(_bs: int, mt: int) -> Tuple[int, ...]:
return (mt,)
def _bs(bs: int, _mt: int) -> Tuple[int, ...]:
return (bs,)
slots = [
GraphSlot(
"input_ids",
_tokens,
torch.int64,
axis="tokens",
padding_policy=PaddingPolicy.ZERO,
),
GraphSlot(
"positions",
_tokens,
torch.int64,
axis="tokens",
padding_policy=PaddingPolicy.ZERO,
),
GraphSlot(
"out_cache_loc",
_tokens,
cache_loc_dtype,
axis="tokens",
padding_policy=PaddingPolicy.ZERO,
),
]
if is_multimodal:
slots.append(
GraphSlot(
"mrope_positions",
lambda _bs2, mt: (3, mt),
torch.int64,
axis="tokens",
padding_policy=PaddingPolicy.ZERO,
slice_fn=lambda buf, n: buf[:, :n],
)
)
slots.append(
GraphSlot(
"input_embeds",
lambda _bs2, mt: (mt, hidden_size),
embed_dtype,
axis="tokens",
padding_policy=PaddingPolicy.ZERO,
copy_from_fb=False,
)
)
if enable_mamba_track:
slots.append(GraphSlot("mamba_track_indices", _bs, torch.int64, axis="bs"))
slots.append(GraphSlot("mamba_track_mask", _bs, torch.bool, axis="bs"))
slots.append(GraphSlot("mamba_track_seqlens", _bs, torch.int32, axis="bs"))
for slot in slots:
bind = None
if source is not None:
bind = getattr(source, slot.name, None)
if bind is None:
raise ValueError(
f"source is missing buffer {slot.name!r} required by the "
"prefill registry; cannot adopt."
)
reg.register_slot(slot, bind=bind)
return reg
@@ -25,7 +25,7 @@ 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 types import SimpleNamespace
from typing import TYPE_CHECKING, Callable, Dict, List, Optional, Tuple, Union from typing import TYPE_CHECKING, Callable, Dict, Optional, Union
import torch import torch
import tqdm import tqdm
@@ -58,6 +58,7 @@ from sglang.srt.layers.moe.token_dispatcher.deepep import DeepEPBuffer
from sglang.srt.layers.moe.utils import get_deepep_mode, get_moe_a2a_backend from sglang.srt.layers.moe.utils import get_deepep_mode, get_moe_a2a_backend
from sglang.srt.layers.utils import MultiPlatformOp from sglang.srt.layers.utils import MultiPlatformOp
from sglang.srt.layers.utils.cp_utils import is_mla_prefill_cp_enabled from sglang.srt.layers.utils.cp_utils import is_mla_prefill_cp_enabled
from sglang.srt.model_executor.cuda_graph_buffer_registry import build_decode_registry
from sglang.srt.model_executor.forward_batch_info import ( from sglang.srt.model_executor.forward_batch_info import (
CaptureHiddenMode, CaptureHiddenMode,
ForwardBatch, ForwardBatch,
@@ -105,8 +106,6 @@ logger = logging.getLogger(__name__)
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.model_executor.model_runner import ModelRunner from sglang.srt.model_executor.model_runner import ModelRunner
_has_foreach_copy = hasattr(torch, "_foreach_copy_")
def build_replay_fb_view( def build_replay_fb_view(
forward_batch: "ForwardBatch", forward_batch: "ForwardBatch",
@@ -159,27 +158,6 @@ def build_replay_fb_view(
) )
def _grouped_foreach_copy_(dsts: List[torch.Tensor], srcs: List[torch.Tensor]) -> None:
"""Call torch._foreach_copy_ grouped by (dst_dtype, src_dtype) pairs."""
def foreach_copy(dsts: List[torch.Tensor], srcs: List[torch.Tensor]) -> None:
if _has_foreach_copy:
torch._foreach_copy_(dsts, srcs)
else:
for dst, src in zip(dsts, srcs):
dst.copy_(src)
groups: Dict[Tuple[torch.dtype, torch.dtype], Tuple[List, List]] = {}
for dst, src in zip(dsts, srcs):
key = (dst.dtype, src.dtype)
if key not in groups:
groups[key] = ([], [])
groups[key][0].append(dst)
groups[key][1].append(src)
for group_dsts, group_srcs in groups.values():
foreach_copy(group_dsts, group_srcs)
@dataclass @dataclass
class DecodeInputBuffers(ForwardInputBuffers): class DecodeInputBuffers(ForwardInputBuffers):
@@ -344,110 +322,19 @@ class DecodeInputBuffers(ForwardInputBuffers):
num_tokens_per_bs: int, num_tokens_per_bs: int,
dsa_enable_prefill_cp: bool, dsa_enable_prefill_cp: bool,
enable_num_token_non_padded_flag: bool, enable_num_token_non_padded_flag: bool,
registry,
pp_proxy_tensors: Optional[PPProxyTensors] = None, pp_proxy_tensors: Optional[PPProxyTensors] = None,
): ):
if bs != raw_bs: # Reset padded tails + copy FB into the registry-adopted graph buffers
self.seq_lens.fill_(seq_len_fill_value) # (same storage the old per-field populate wrote).
self.out_cache_loc.zero_() registry.fill_from(
# Pair with seq_lens fill: padded rows must point at reserved forward_batch,
# req_pool slot 0 (req_to_token[0, :] is all zeros from init), raw_bs=raw_bs,
# so dummy attention reads land on slot 0 instead of a stale padded_bs=bs,
# req_to_token row left by an earlier replay. raw_num_tokens=raw_num_token,
self.req_pool_indices.zero_() padded_num_tokens=bs * num_tokens_per_bs,
if self.mamba_track_indices is not None: pp_proxy_tensors=pp_proxy_tensors,
self.mamba_track_indices.zero_() )
if self.mamba_track_mask is not None:
self.mamba_track_mask.fill_(False)
# Build batched copy lists for all GPU tensors.
dsts = [
self.input_ids[:raw_num_token],
self.req_pool_indices[:raw_bs],
self.seq_lens[:raw_bs],
self.out_cache_loc[:raw_num_token],
self.positions[:raw_num_token],
]
srcs = [
forward_batch.input_ids,
forward_batch.req_pool_indices,
forward_batch.seq_lens,
forward_batch.out_cache_loc,
forward_batch.positions,
]
if self.ngram_embedding_info is not None:
ngram_embedding_info = forward_batch.ngram_embedding_info
self.ngram_embedding_info.column_starts[:raw_bs].copy_(
ngram_embedding_info.column_starts
)
self.ngram_embedding_info.req_lens[:raw_bs].copy_(
ngram_embedding_info.req_lens
)
if (
self.mamba_track_indices is not None
and forward_batch.mamba_track_indices is not None
):
dsts.append(self.mamba_track_indices[:raw_bs])
srcs.append(forward_batch.mamba_track_indices)
if (
self.mamba_track_mask is not None
and forward_batch.mamba_track_mask is not None
):
dsts.append(self.mamba_track_mask[:raw_bs])
srcs.append(forward_batch.mamba_track_mask)
if self.encoder_lens is not None and forward_batch.encoder_lens is not None:
dsts.append(self.encoder_lens[:raw_bs])
srcs.append(forward_batch.encoder_lens)
if forward_batch.mrope_positions is not None:
dsts.append(self.mrope_positions[:, :raw_num_token])
srcs.append(forward_batch.mrope_positions)
if self.rids_int is not None and forward_batch.rids_int is not None:
dsts.append(self.rids_int[:raw_bs])
srcs.append(forward_batch.rids_int)
if (
self.bootstrap_room_ids_int is not None
and forward_batch.bootstrap_room_ids_int is not None
):
dsts.append(self.bootstrap_room_ids_int[:raw_bs])
srcs.append(forward_batch.bootstrap_room_ids_int)
if require_gathered_buffer:
self.global_num_tokens_gpu.fill_(bs * num_tokens_per_bs)
self.global_num_tokens_for_logprob_gpu.fill_(bs * num_tokens_per_bs)
if enable_num_token_non_padded_flag:
if require_gathered_buffer and not dsa_enable_prefill_cp:
num_tokens_per_dp = bs * num_tokens_per_bs
local = compute_local_num_token_non_padded(
global_num_token_non_padded=forward_batch.num_token_non_padded,
num_tokens_per_dp=num_tokens_per_dp,
)
dsts.append(self.num_token_non_padded)
srcs.append(local)
else:
dsts.append(self.num_token_non_padded)
srcs.append(forward_batch.num_token_non_padded)
# Pipeline-parallel proxy tensors.
if pp_proxy_tensors is not None and self.pp_proxy_tensors is not None:
for key, buf in self.pp_proxy_tensors.items():
src = pp_proxy_tensors.tensors[key]
dim = src.shape[0]
dsts.append(buf[:dim])
srcs.append(src)
# Batch all GPU copies, grouped by dtype pair.
_grouped_foreach_copy_(dsts, srcs)
# CPU tensor copy (cannot be batched with GPU tensors).
if forward_batch.seq_lens_cpu is not None:
if bs != raw_bs:
self.seq_lens_cpu.fill_(seq_len_fill_value)
self.seq_lens_cpu[:raw_bs].copy_(forward_batch.seq_lens_cpu)
# Detect whether the current forward pass is in capture mode # Detect whether the current forward pass is in capture mode
@@ -767,6 +654,26 @@ class CudaGraphRunner:
), ),
) )
self.buffers.share_buffers() self.buffers.share_buffers()
# FB-shared slot registry, adopting the DecodeInputBuffers storage so
# it mirrors the same physical buffers (stable data_ptr for capture vs
# replay). This is the unified fill/extract surface that eager /
# capture / replay migrate onto, replacing populate_from_forward_batch.
self.buffer_registry = build_decode_registry(
device=self.device,
max_bs=self.max_bs,
max_num_token=self.max_num_token,
seq_len_fill_value=self.seq_len_fill_value,
cache_loc_dtype=self._cache_loc_dtype(),
enable_mamba_track=enable_mamba_track,
is_encoder_decoder=self.is_encoder_decoder,
encoder_len_fill_value=self.encoder_len_fill_value,
enable_num_token_non_padded=enable_num_token_non_padded(),
require_gathered_buffer=self.require_gathered_buffer,
enable_prefill_cp=self.enable_prefill_cp,
require_mlp_tp_gather=self.require_mlp_tp_gather,
dp_size=self.dp_size,
source=self.buffers,
)
self.tbo_plugin = TboCudaGraphRunnerPlugin() self.tbo_plugin = TboCudaGraphRunnerPlugin()
@@ -995,18 +902,24 @@ class CudaGraphRunner:
stream = self.stream stream = self.stream
num_tokens = bs * self.num_tokens_per_bs num_tokens = bs * self.num_tokens_per_bs
# Graph inputs # Graph inputs. The registry-owned FB-shared slots come from the
input_ids = buffers.input_ids[:num_tokens] # registry (it adopted the DecodeInputBuffers storage, so these are the
req_pool_indices = buffers.req_pool_indices[:bs] # same physical tensors); the rest still come off `buffers` directly.
seq_lens = buffers.seq_lens[:bs] registry = self.buffer_registry
seq_lens_cpu = buffers.seq_lens_cpu[:bs]
out_cache_loc = buffers.out_cache_loc[:num_tokens] def _slot(name):
positions = buffers.positions[:num_tokens] return registry.get_slot(name).slice_for(bs, num_tokens)
if self.is_encoder_decoder:
encoder_lens = buffers.encoder_lens[:bs] input_ids = _slot("input_ids")
else: req_pool_indices = _slot("req_pool_indices")
encoder_lens = None seq_lens = _slot("seq_lens")
mrope_positions = buffers.mrope_positions[:, :num_tokens] seq_lens_cpu = _slot("seq_lens_cpu")
out_cache_loc = _slot("out_cache_loc")
positions = _slot("positions")
encoder_lens = (
_slot("encoder_lens") if registry.has_slot("encoder_lens") else None
)
mrope_positions = _slot("mrope_positions")
next_token_logits_buffer = buffers.next_token_logits_buffer[:num_tokens] next_token_logits_buffer = buffers.next_token_logits_buffer[:num_tokens]
rids_int = buffers.rids_int[:bs] if buffers.rids_int is not None else None rids_int = buffers.rids_int[:bs] if buffers.rids_int is not None else None
bootstrap_room_ids_int = ( bootstrap_room_ids_int = (
@@ -1065,16 +978,14 @@ class CudaGraphRunner:
else: else:
lora_ids = None lora_ids = None
# mamba state tracking # mamba state tracking (registry-owned when enabled)
mamba_track_indices = ( mamba_track_indices = (
buffers.mamba_track_indices[:bs] _slot("mamba_track_indices")
if buffers.mamba_track_indices is not None if registry.has_slot("mamba_track_indices")
else None else None
) )
mamba_track_mask = ( mamba_track_mask = (
buffers.mamba_track_mask[:bs] _slot("mamba_track_mask") if registry.has_slot("mamba_track_mask") else None
if buffers.mamba_track_mask is not None
else None
) )
if stream_idx is None: if stream_idx is None:
@@ -1275,6 +1186,7 @@ class CudaGraphRunner:
# "any prefill-CP flavor enabled" (DSA CP or MLA CP). # "any prefill-CP flavor enabled" (DSA CP or MLA CP).
dsa_enable_prefill_cp=self.enable_prefill_cp, dsa_enable_prefill_cp=self.enable_prefill_cp,
enable_num_token_non_padded_flag=enable_num_token_non_padded(), enable_num_token_non_padded_flag=enable_num_token_non_padded(),
registry=self.buffer_registry,
pp_proxy_tensors=pp_proxy_tensors, pp_proxy_tensors=pp_proxy_tensors,
) )
@@ -2,36 +2,42 @@ from __future__ import annotations
import dataclasses import dataclasses
from dataclasses import dataclass, fields from dataclasses import dataclass, fields
from typing import Dict from typing import Dict, Tuple
import torch import torch
from sglang.srt.utils import is_npu from sglang.srt.utils import is_npu
_forward_input_buffer_pool: Dict[str, torch.Tensor] = {} # Process-wide pool keyed by (name, numel, dtype, device); see share_input_buffer.
_PoolKey = Tuple[str, int, torch.dtype, torch.device]
_forward_input_buffer_pool: Dict[_PoolKey, torch.Tensor] = {}
def share_input_buffer(name: str, new_buffer: torch.Tensor) -> torch.Tensor:
"""Coalesce a buffer by ``(name, size, dtype, device)`` into the
process-wide input-buffer pool.
Distinct callers that request the same field ``name`` with the same
size/dtype/device share one physical allocation (and therefore one
``data_ptr``): the first registrant's buffer becomes canonical and every
later identical request is returned as a view aliased onto it. Requests
that differ in size get their own allocation — they never reuse or displace
an existing entry — so the sharing *structure* is independent of
registration order and no already-captured buffer is ever repointed.
"""
key: _PoolKey = (name, new_buffer.numel(), new_buffer.dtype, new_buffer.device)
canonical = _forward_input_buffer_pool.get(key, None)
if canonical is None:
_forward_input_buffer_pool[key] = new_buffer
canonical = new_buffer
return canonical.as_strided(new_buffer.size(), new_buffer.stride())
@dataclass @dataclass
class ForwardInputBuffers: class ForwardInputBuffers:
def _share_one_buffer(self, name: str, new_buffer: torch.Tensor) -> torch.Tensor: def _share_one_buffer(self, name: str, new_buffer: torch.Tensor) -> torch.Tensor:
return share_input_buffer(name, new_buffer)
buffer_size = new_buffer.size()
buffer_stride = new_buffer.stride()
old_buffer = _forward_input_buffer_pool.get(name, None)
if old_buffer is not None:
assert (
new_buffer.dtype == old_buffer.dtype
), f"Buffer {name} has different dtype than before."
assert (
new_buffer.device == old_buffer.device
), f"Buffer {name} has different device than before."
if old_buffer.numel() > new_buffer.numel():
new_buffer = old_buffer
_forward_input_buffer_pool[name] = new_buffer
return new_buffer.as_strided(buffer_size, buffer_stride)
def share_buffers(self): def share_buffers(self):
# disable share input buffer on npu due to accuracy issue # disable share input buffer on npu due to accuracy issue
@@ -52,6 +52,7 @@ from sglang.srt.layers.logits_processor import LogitsProcessorOutput
from sglang.srt.layers.moe.utils import get_moe_a2a_backend from sglang.srt.layers.moe.utils import get_moe_a2a_backend
from sglang.srt.layers.pooler import EmbeddingPoolerOutput from sglang.srt.layers.pooler import EmbeddingPoolerOutput
from sglang.srt.layers.utils import MultiPlatformOp from sglang.srt.layers.utils import MultiPlatformOp
from sglang.srt.model_executor.cuda_graph_buffer_registry import build_prefill_registry
from sglang.srt.model_executor.forward_batch_info import ( from sglang.srt.model_executor.forward_batch_info import (
CaptureHiddenMode, CaptureHiddenMode,
ForwardBatch, ForwardBatch,
@@ -295,6 +296,20 @@ class PiecewiseCudaGraphRunner:
) )
self.buffers.share_buffers() self.buffers.share_buffers()
# Token-axis FB-shared slot registry, adopting the PrefillInputBuffers
# storage (one data_ptr shared with capture + replay).
self.buffer_registry = build_prefill_registry(
device=self.device,
max_bs=self.max_bs,
max_num_token=self.max_num_tokens,
cache_loc_dtype=self._cache_loc_dtype(),
is_multimodal=self.is_multimodal,
hidden_size=self.model_runner.model_config.hidden_size,
embed_dtype=self.model_runner.dtype,
enable_mamba_track=self.mamba_track_enabled,
source=self.buffers,
)
self.attention_layers = self.model_runner.attention_layers self.attention_layers = self.model_runner.attention_layers
self.moe_layers = self.model_runner.moe_layers self.moe_layers = self.model_runner.moe_layers
self.moe_fusions = self.model_runner.moe_fusions self.moe_fusions = self.model_runner.moe_fusions
@@ -355,27 +370,32 @@ class PiecewiseCudaGraphRunner:
def warmup_compile(self, num_tokens: int): def warmup_compile(self, num_tokens: int):
"""Warmup the model with a simple forward pass before CUDA graph capture.""" """Warmup the model with a simple forward pass before CUDA graph capture."""
buffers = self.buffers registry = self.buffer_registry
input_ids = buffers.input_ids[:num_tokens] bs = 1
input_embeds = buffers.input_embeds[:num_tokens] if self.is_multimodal else None
positions = buffers.positions[:num_tokens] def _slot(name):
return registry.get_slot(name).slice_for(bs, num_tokens)
input_ids = _slot("input_ids")
positions = _slot("positions")
out_cache_loc = _slot("out_cache_loc")
input_embeds = (
_slot("input_embeds") if registry.has_slot("input_embeds") else None
)
mrope_positions = ( mrope_positions = (
buffers.mrope_positions[:, :num_tokens] if self.is_multimodal else None _slot("mrope_positions") if registry.has_slot("mrope_positions") else None
) )
out_cache_loc = buffers.out_cache_loc[:num_tokens]
mamba_track_indices = ( mamba_track_indices = (
buffers.mamba_track_indices[:1] _slot("mamba_track_indices")
if buffers.mamba_track_indices is not None if registry.has_slot("mamba_track_indices")
else None else None
) )
mamba_track_mask = ( mamba_track_mask = (
buffers.mamba_track_mask[:1] _slot("mamba_track_mask") if registry.has_slot("mamba_track_mask") else None
if buffers.mamba_track_mask is not None
else None
) )
mamba_track_seqlens = ( mamba_track_seqlens = (
buffers.mamba_track_seqlens[:1] _slot("mamba_track_seqlens")
if buffers.mamba_track_seqlens is not None if registry.has_slot("mamba_track_seqlens")
else None else None
) )
with torch.device(self.device): with torch.device(self.device):
@@ -508,33 +528,36 @@ class PiecewiseCudaGraphRunner:
self.capture_one_batch_size(num_tokens) self.capture_one_batch_size(num_tokens)
def capture_one_batch_size(self, num_tokens: int): def capture_one_batch_size(self, num_tokens: int):
buffers = self.buffers registry = self.buffer_registry
bs = 1 bs = 1
# Graph inputs # Graph inputs — views into the registry's (adopted) graph-resident
input_ids = buffers.input_ids[:num_tokens] # slots; capture burns these addresses into the graph.
input_embeds = buffers.input_embeds[:num_tokens] if self.is_multimodal else None def _slot(name):
return registry.get_slot(name).slice_for(bs, num_tokens)
out_cache_loc = buffers.out_cache_loc[:num_tokens] input_ids = _slot("input_ids")
positions = _slot("positions")
out_cache_loc = _slot("out_cache_loc")
input_embeds = (
_slot("input_embeds") if registry.has_slot("input_embeds") else None
)
mrope_positions = (
_slot("mrope_positions") if registry.has_slot("mrope_positions") else None
)
mamba_track_indices = ( mamba_track_indices = (
buffers.mamba_track_indices[:bs] _slot("mamba_track_indices")
if buffers.mamba_track_indices is not None if registry.has_slot("mamba_track_indices")
else None else None
) )
mamba_track_mask = ( mamba_track_mask = (
buffers.mamba_track_mask[:bs] _slot("mamba_track_mask") if registry.has_slot("mamba_track_mask") else None
if buffers.mamba_track_mask is not None
else None
) )
mamba_track_seqlens = ( mamba_track_seqlens = (
buffers.mamba_track_seqlens[:bs] _slot("mamba_track_seqlens")
if buffers.mamba_track_seqlens is not None if registry.has_slot("mamba_track_seqlens")
else None else None
) )
positions = buffers.positions[:num_tokens]
mrope_positions = (
buffers.mrope_positions[:, :num_tokens] if self.is_multimodal else None
)
global_dp_buffer_len = None global_dp_buffer_len = None
global_num_tokens_cpu = None global_num_tokens_cpu = None
@@ -649,72 +672,51 @@ class PiecewiseCudaGraphRunner:
forward_batch: ForwardBatch, forward_batch: ForwardBatch,
**kwargs, **kwargs,
): ):
buffers = self.buffers
num_tokens = len(forward_batch.input_ids) num_tokens = len(forward_batch.input_ids)
index = bisect.bisect_left(self.capture_num_tokens, num_tokens) index = bisect.bisect_left(self.capture_num_tokens, num_tokens)
static_num_tokens = self.capture_num_tokens[index] static_num_tokens = self.capture_num_tokens[index]
self.raw_num_tokens = num_tokens self.raw_num_tokens = num_tokens
if static_num_tokens != num_tokens:
buffers.out_cache_loc.zero_()
buffers.input_ids[num_tokens:static_num_tokens].zero_()
buffers.positions[num_tokens:static_num_tokens].zero_()
if self.is_multimodal:
buffers.input_embeds[num_tokens:static_num_tokens].zero_()
if forward_batch.mrope_positions is not None:
buffers.mrope_positions[:, num_tokens:static_num_tokens].zero_()
bs = forward_batch.batch_size bs = forward_batch.batch_size
registry = self.buffer_registry
# Reset the padded token tail (ZERO) + copy the [:num_tokens] head for
# every graph-resident slot in one grouped pass. input_embeds is
# reset-only (the model writes embeds into it inside the graph).
registry.fill_from(
forward_batch,
raw_bs=bs,
padded_bs=bs,
raw_num_tokens=num_tokens,
padded_num_tokens=static_num_tokens,
)
buffers.input_ids[:num_tokens].copy_(forward_batch.input_ids) def _slot(name):
buffers.positions[:num_tokens].copy_(forward_batch.positions) return registry.get_slot(name).slice_for(bs, static_num_tokens)
buffers.out_cache_loc[:num_tokens].copy_(forward_batch.out_cache_loc)
if (
buffers.mamba_track_indices is not None
and forward_batch.mamba_track_indices is not None
):
buffers.mamba_track_indices[:bs].copy_(forward_batch.mamba_track_indices)
if (
buffers.mamba_track_mask is not None
and forward_batch.mamba_track_mask is not None
):
buffers.mamba_track_mask[:bs].copy_(forward_batch.mamba_track_mask)
if (
buffers.mamba_track_seqlens is not None
and forward_batch.mamba_track_seqlens is not None
):
buffers.mamba_track_seqlens[:bs].copy_(forward_batch.mamba_track_seqlens)
input_ids = buffers.input_ids[:static_num_tokens]
positions = buffers.positions[:static_num_tokens]
out_cache_loc = buffers.out_cache_loc[:static_num_tokens]
input_ids = _slot("input_ids")
positions = _slot("positions")
out_cache_loc = _slot("out_cache_loc")
mamba_track_indices = ( mamba_track_indices = (
buffers.mamba_track_indices[:bs] _slot("mamba_track_indices")
if buffers.mamba_track_indices is not None if registry.has_slot("mamba_track_indices")
else None else None
) )
mamba_track_mask = ( mamba_track_mask = (
buffers.mamba_track_mask[:bs] _slot("mamba_track_mask") if registry.has_slot("mamba_track_mask") else None
if buffers.mamba_track_mask is not None
else None
) )
mamba_track_seqlens = ( mamba_track_seqlens = (
buffers.mamba_track_seqlens[:bs] _slot("mamba_track_seqlens")
if buffers.mamba_track_seqlens is not None if registry.has_slot("mamba_track_seqlens")
else None else None
) )
if forward_batch.mrope_positions is not None:
buffers.mrope_positions[:, :num_tokens].copy_(forward_batch.mrope_positions)
input_ids = buffers.input_ids[:static_num_tokens]
input_embeds = ( input_embeds = (
buffers.input_embeds[:static_num_tokens] if self.is_multimodal else None _slot("input_embeds") if registry.has_slot("input_embeds") else None
) )
mrope_positions = ( mrope_positions = (
buffers.mrope_positions[:, :static_num_tokens] _slot("mrope_positions")
if forward_batch.mrope_positions is not None if (
registry.has_slot("mrope_positions")
and forward_batch.mrope_positions is not None
)
else None else None
) )
File diff suppressed because it is too large Load Diff