[FullCG] Support chunked cached-prefix prefill (#30825)

Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
paulzhang-tm
2026-07-28 11:20:25 -07:00
committed by GitHub
co-authored by Claude Fable 5
parent 1eee8fbdcc
commit 4e5a05148a
10 changed files with 746 additions and 37 deletions
@@ -105,6 +105,27 @@ class AttentionBackend(ABC):
# object during capture, and refresh its dynamic fields before each replay.
use_captured_forward_metadata_for_breakable_cuda_graph: bool = False
# Chunked-prefix FullCG capture has a second model topology and stable
# prefix buffers. Backends must opt in explicitly so the runner does not
# assume that generic ForwardBatch metadata is sufficient for every
# attention implementation.
supports_full_cuda_graph_chunked_prefix: bool = False
def prepare_full_cuda_graph_chunked_prefix(
self,
forward_batch: ForwardBatch,
*,
in_capture: bool,
) -> None:
"""Prepare backend-private metadata for chunked-prefix FullCG.
Only called for backends that set
``supports_full_cuda_graph_chunked_prefix``; the runner validates the
flag up front. The runner owns and refreshes the shared ForwardBatch
prefix buffers. Backends that need wrappers or other derived metadata
should override this hook for both capture and replay.
"""
def init_cuda_graph_state(self, max_bs: int, max_num_tokens: int):
"""Init the global shared states for cuda graph."""
raise NotImplementedError()
@@ -136,6 +136,10 @@ class FlashAttentionBackend(AttentionBackend):
needs_cpu_seq_lens: bool = False
supports_ragged_verify_graph: bool = True
# Chunked-prefix attention reads the stable ForwardBatch cu-seqlens and
# KV-index buffers directly, so it needs no backend-private replay state.
supports_full_cuda_graph_chunked_prefix = True
def __init__(
self,
model_runner: ModelRunner,
@@ -63,6 +63,20 @@ class HybridAttnBackend(AttentionBackend):
else:
return self.prefill_backend
@property
def supports_full_cuda_graph_chunked_prefix(self) -> bool:
return self.prefill_backend.supports_full_cuda_graph_chunked_prefix
def prepare_full_cuda_graph_chunked_prefix(
self,
forward_batch: ForwardBatch,
*,
in_capture: bool,
) -> None:
self.prefill_backend.prepare_full_cuda_graph_chunked_prefix(
forward_batch, in_capture=in_capture
)
def init_forward_metadata_out_graph(
self,
forward_batch: ForwardBatch,
@@ -66,12 +66,20 @@ ALLOWED_BACKENDS_PER_PHASE = {
# Per-phase settings schema. Keys other than backend are runner-level
# (read by any backend in that phase); tc_compiler is the lone
# backend-specific knob (only meaningful when backend == tc_piecewise).
# For prefill, bs carries the captured shape size (token count for
# For prefill, bs carries the captured shape size (token count for Full and
# tc_piecewise, request count for breakable) — one shape knob per phase.
# full_prefill_max_req is prefill-only and only meaningful when backend == full.
# full_prefill_max_req and full_prefill_prefix_chunk_tokens are prefill-only and
# only meaningful when backend == full.
ALLOWED_KEYS_PER_PHASE = {
Phase.DECODE: ("backend", "max_bs", "bs", "tc_compiler"),
Phase.PREFILL: ("backend", "max_bs", "bs", "tc_compiler", "full_prefill_max_req"),
Phase.PREFILL: (
"backend",
"max_bs",
"bs",
"tc_compiler",
"full_prefill_max_req",
"full_prefill_prefix_chunk_tokens",
),
}
@@ -90,6 +98,12 @@ class PhaseConfig:
# batches fall back to eager. Ignored by BCG (bs=1 only) and TC_PIECEWISE
# (bs-invariant via torch.compile). None auto-derives chunked_prefill_size // 512.
full_prefill_max_req: Optional[int] = None
# Only meaningful for Full prefill CUDA graphs that capture a distinct
# cached-prefix topology: aggregate cached-prefix tokens represented by one
# fixed-capacity chunk across all request slots. FullCG captures 1/2/4/8/16
# chunk variants and chooses the smallest one covering a batch. None uses
# the scheduler's aggregate chunked_prefill_size token budget.
full_prefill_prefix_chunk_tokens: Optional[int] = None
def default_prefill_backend() -> str:
@@ -41,11 +41,15 @@ import copy
import inspect
import logging
from contextlib import contextmanager
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any, Dict, Optional, Union
import torch
import tqdm
from sglang.kernels.ops.kvcache.kv_indices import (
create_chunked_prefix_cache_kv_indices,
)
from sglang.srt.distributed.parallel_state import graph_capture
from sglang.srt.layers.attention.dsa.utils import is_dsa_enable_prefill_cp
from sglang.srt.layers.dp_attention import (
@@ -118,6 +122,17 @@ logger = logging.getLogger(__name__)
# lists can otherwise turn the lower launch overhead into substantially more
# model work than an exact-shape eager forward.
_MAX_PREFILL_CUDA_GRAPH_PADDING_FACTOR = 2
# Prefix attention adds one loop body per chunk to the captured topology, so
# capture a small geometric set and round each replay up to the nearest one.
_CHUNKED_PREFIX_VARIANTS = (1, 2, 4, 8, 16)
def _chunked_prefix_variant(num_chunks: int) -> str:
return f"chunked_prefix:{num_chunks}"
def _ceil_div(a: int, b: int) -> int:
return -(-a // b)
def _slice_output_rows(output: Any, num_tokens: int) -> Any:
@@ -139,6 +154,22 @@ def _slice_output_rows(output: Any, num_tokens: int) -> Any:
raise TypeError(f"Unsupported full prefill CUDA graph output: {type(output)}")
@dataclass(frozen=True)
class _ChunkedPrefixCaptureBuffers:
"""Runner-owned tensors shared by every prefix and token-bucket variant.
``starts`` / ``starts_cpu`` are constant (chunk ``c`` always starts at
``c * prefix_chunk_len``) and filled once at allocation.
"""
starts: torch.Tensor # (max_chunks, req_slots)
seq_lens: torch.Tensor # (max_chunks, req_slots)
cu_seq_lens: torch.Tensor # (max_chunks, req_slots + 1)
starts_cpu: torch.Tensor
seq_lens_cpu: torch.Tensor
kv_indices: torch.Tensor # (max_chunks, prefix_chunk_capacity)
def prefill_failure_msg(backend_name: str) -> str:
"""Render PREFILL_CUDA_GRAPH_CAPTURE_FAILED_MSG with a backend-specific
numbered suggestion list. The runner is only constructed for BREAKABLE
@@ -308,7 +339,6 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
# Auto: scale request slots with the chunked prefill size.
max_req = max(model_runner.server_args.chunked_prefill_size // 512, 1)
self._capture_req_slots = min(max_req, self.max_bs)
# BCG/Full record LoRA kernels, so the metadata they read must live in
# static buffers refreshed in place per batch; unsupported LoRA
# configs were already routed to the eager runner.
@@ -331,24 +361,62 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
)
self._capture_req_slots = lora_max_bs
if self._is_full_backend:
self._full_cg_seq_lens_cpu = torch.zeros(
(self._capture_req_slots,), dtype=torch.int64, device="cpu"
self._full_cg_seq_lens_cpu = (
torch.zeros((self._capture_req_slots,), dtype=torch.int64, device="cpu")
if self._is_full_backend
else None
)
# This flag controls whether the model dispatches through the distinct
# chunked-prefix topology; backend capability is validated separately.
self._capture_chunked_prefix = (
self._is_full_backend
and not model_runner.server_args.disable_chunked_prefix_cache
)
self._prefix_chunk_len = 0
self._prefix_chunk_capacity = 0
self._prefix_max_len = 0
self._prefix_capture_variants: tuple[int, ...] = ()
self._prefix_capture_batches: Dict[ShapeKey, ForwardBatch] = {}
self._prefix_capture_buffers: Optional[_ChunkedPrefixCaptureBuffers] = None
if self._capture_chunked_prefix:
attn_backend = model_runner.attn_backend
assert attn_backend.supports_full_cuda_graph_chunked_prefix, (
f"{type(attn_backend).__name__} does not support chunked-prefix "
"Full prefill CUDA graphs"
)
prefix_config = model_runner.server_args.cuda_graph_config.prefill
(
self._prefix_chunk_len,
self._prefix_chunk_capacity,
) = self._resolve_prefix_chunk_shape(model_runner, self._capture_req_slots)
self._prefix_max_len = self._max_addressable_prefix_len(model_runner)
max_real_chunks = _ceil_div(self._prefix_max_len, self._prefix_chunk_len)
# Variants are geometric, so n // 2 is the next-smaller variant; keep
# n only if the smaller variant does not already cover the max prefix.
self._prefix_capture_variants = tuple(
n for n in _CHUNKED_PREFIX_VARIANTS if n // 2 < max_real_chunks
)
self._prefix_capture_buffers = self._create_chunked_prefix_buffers()
logger.info(
"Full prefill CUDA graph cached-prefix chunks: "
"%d aggregate tokens/chunk (%d/request x %d slots), "
"variants=%s (%s)",
self._prefix_chunk_capacity,
self._prefix_chunk_len,
self._capture_req_slots,
self._prefix_capture_variants,
(
"configured"
if prefix_config.full_prefill_prefix_chunk_tokens is not None
else "auto from chunked_prefill_size"
),
)
if isinstance(self.backend, (BreakableCudaGraphBackend, FullCudaGraphBackend)):
with torch.device(self.device):
self._prefill_static_buffers = {
name: torch.zeros((self.max_bs,), dtype=torch.int64)
for name in _PREFILL_STATIC_FIELDS
}
elif isinstance(self.backend, BreakableCudaGraphBackend):
self._full_cg_seq_lens_cpu = None
with torch.device(self.device):
self._prefill_static_buffers = {
name: torch.zeros((self.max_bs,), dtype=torch.int64)
for name in _PREFILL_STATIC_FIELDS
}
else:
self._full_cg_seq_lens_cpu = None
# Static hidden_states buffer giving the captured graph a stable
# address; load_batch refreshes it from live spec_info at replay.
@@ -649,6 +717,197 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
int(num_tokens) == 0 for num_tokens in global_num_tokens
)
@staticmethod
def _has_prefix_hit(forward_batch: ForwardBatch) -> bool:
prefix_lens = forward_batch.extend_prefix_lens_cpu
return prefix_lens is not None and any(
int(length) > 0 for length in prefix_lens
)
@staticmethod
def _max_addressable_prefix_len(model_runner) -> int:
table_width = model_runner.req_to_token_pool.req_to_token.shape[1]
configured_context = model_runner.server_args.context_length
return (
min(table_width, configured_context)
if configured_context is not None and configured_context > 0
else table_width
)
@staticmethod
def _resolve_prefix_chunk_shape(
model_runner, capture_req_slots: int
) -> tuple[int, int]:
"""Resolve per-request length and aggregate capacity of one chunk."""
prefix_config = model_runner.server_args.cuda_graph_config.prefill
requested_capacity = prefix_config.full_prefill_prefix_chunk_tokens
if requested_capacity is None:
requested_capacity = model_runner.server_args.chunked_prefill_size
if requested_capacity is None or requested_capacity <= 0:
requested_capacity = prefix_config.max_bs
if requested_capacity is None or requested_capacity <= 0:
raise ValueError("full_prefill_prefix_chunk_tokens must be positive")
max_addressable_prefix_tokens = (
PrefillCudaGraphRunner._max_addressable_prefix_len(model_runner)
* capture_req_slots
)
requested_capacity = min(requested_capacity, max_addressable_prefix_tokens)
# Divide before allocating so increasing the number of request slots
# does not multiply the aggregate workspace represented by one chunk.
prefix_chunk_len = max(requested_capacity // capture_req_slots, 1)
return prefix_chunk_len, prefix_chunk_len * capture_req_slots
def _select_prefix_capture_chunks(
self, forward_batch: ForwardBatch
) -> Optional[int]:
"""Smallest captured variant covering the batch's max prefix, or None."""
max_prefix_len = max(
int(length) for length in forward_batch.extend_prefix_lens_cpu
)
real_n = _ceil_div(max_prefix_len, self._prefix_chunk_len)
return next((n for n in self._prefix_capture_variants if n >= real_n), None)
def _shape_key(self, num_tokens: int, forward_batch: ForwardBatch) -> ShapeKey:
variant = None
if self._capture_chunked_prefix and self._has_prefix_hit(forward_batch):
captured_n = self._select_prefix_capture_chunks(forward_batch)
assert captured_n is not None, "prefix batch has no captured FullCG variant"
variant = _chunked_prefix_variant(captured_n)
return ShapeKey(size=num_tokens, variant_label=variant)
def _create_chunked_prefix_buffers(self) -> _ChunkedPrefixCaptureBuffers:
"""Allocate the stable chunk-metadata tensors shared by all variants."""
max_chunks = max(self._prefix_capture_variants)
bs = self._capture_req_slots
starts_cpu = (
(torch.arange(max_chunks, dtype=torch.int32) * self._prefix_chunk_len)
.unsqueeze(1)
.repeat(1, bs)
)
return _ChunkedPrefixCaptureBuffers(
starts=starts_cpu.to(self.device, copy=True),
seq_lens=torch.zeros(
(max_chunks, bs), dtype=torch.int32, device=self.device
),
cu_seq_lens=torch.zeros(
(max_chunks, bs + 1), dtype=torch.int32, device=self.device
),
starts_cpu=starts_cpu,
seq_lens_cpu=torch.zeros((max_chunks, bs), dtype=torch.int32),
kv_indices=torch.zeros(
(max_chunks, self._prefix_chunk_capacity),
dtype=torch.int32,
device=self.device,
),
)
def _prepare_chunked_prefix_capture(
self,
forward_batch: ForwardBatch,
shape_key: ShapeKey,
captured_n: int,
) -> None:
"""Attach fixed-capacity, runner-owned prefix buffers for capture."""
assert self._prefill_static_buffers is not None
buffers = self._prefix_capture_buffers
assert buffers is not None
bs = self._capture_req_slots
prefix_len = min(self._prefix_chunk_len * captured_n, self._prefix_max_len)
prefix_lens_cpu = [prefix_len] * bs
self._prefill_static_buffers["extend_prefix_lens"][:bs].fill_(prefix_len)
self._populate_chunked_prefix_buffers(
captured_n=captured_n, prefix_lens_cpu=prefix_lens_cpu
)
forward_batch.extend_prefix_lens_cpu = prefix_lens_cpu
forward_batch.prefix_chunk_len = self._prefix_chunk_len
forward_batch.num_prefix_chunks = captured_n
forward_batch.prefix_chunk_idx = -1
forward_batch.prefix_chunk_starts = buffers.starts
forward_batch.prefix_chunk_starts_cpu = buffers.starts_cpu
forward_batch.prefix_chunk_seq_lens = buffers.seq_lens
forward_batch.prefix_chunk_seq_lens_cpu = buffers.seq_lens_cpu
forward_batch.prefix_chunk_cu_seq_lens = buffers.cu_seq_lens
forward_batch.prefix_chunk_max_seq_lens = [self._prefix_chunk_len] * captured_n
# Replay may pad request slots with zero-length prefixes. Capture the
# conservative branch needed by backends that fix up zero-KV rows.
forward_batch.prefix_chunk_has_zero_kv = [True] * captured_n
forward_batch.prefix_chunk_num_tokens = [
self._prefix_chunk_capacity
] * captured_n
forward_batch.prefix_chunk_kv_indices = list(
buffers.kv_indices[:captured_n].unbind(0)
)
self._prefix_capture_batches[shape_key] = forward_batch
self.model_runner.attn_backend.prepare_full_cuda_graph_chunked_prefix(
forward_batch, in_capture=True
)
def _populate_chunked_prefix_buffers(
self,
*,
captured_n: int,
prefix_lens_cpu: list[int],
) -> None:
"""Refresh shared metadata/KV indices and zero rounded-up chunks."""
assert self._prefill_static_buffers is not None
buffers = self._prefix_capture_buffers
assert buffers is not None
bs = self._capture_req_slots
assert len(prefix_lens_cpu) == bs
real_n = _ceil_div(max(prefix_lens_cpu), self._prefix_chunk_len)
assert real_n <= captured_n
# Chunk starts are constant; only the lengths change. Compute once on
# CPU and copy into the stable CPU/GPU buffers the graph reads.
prefix_lens = torch.tensor(prefix_lens_cpu, dtype=torch.int32)
seq_lens = (prefix_lens - buffers.starts_cpu[:captured_n]).clamp(
min=0, max=self._prefix_chunk_len
)
cu_seq_lens = torch.zeros((captured_n, bs + 1), dtype=torch.int32)
cu_seq_lens[:, 1:] = seq_lens.cumsum(dim=1)
buffers.seq_lens_cpu[:captured_n].copy_(seq_lens)
buffers.seq_lens[:captured_n].copy_(seq_lens)
buffers.cu_seq_lens[:captured_n].copy_(cu_seq_lens)
req_to_token = self.model_runner.req_to_token_pool.req_to_token
# The kernel reads all request slots, so use the slot-padded static
# buffer (arange at capture, live indices + zeroed tail at replay).
req_pool_indices = self._prefill_static_buffers["req_pool_indices"][:bs]
buffers.kv_indices[:captured_n].zero_()
for chunk_idx in range(real_n):
create_chunked_prefix_cache_kv_indices[(bs,)](
req_to_token,
req_pool_indices,
buffers.starts[chunk_idx],
buffers.seq_lens[chunk_idx],
buffers.cu_seq_lens[chunk_idx],
buffers.kv_indices[chunk_idx],
req_to_token.shape[1],
)
def _prepare_chunked_prefix_replay(
self, shape_key: ShapeKey, forward_batch: ForwardBatch
) -> None:
"""Refresh stable buffers for the selected rounded-up chunk variant."""
capture_batch = self._prefix_capture_batches[shape_key]
raw_bs = forward_batch.batch_size
prefix_lens_cpu = [
int(length) for length in forward_batch.extend_prefix_lens_cpu[:raw_bs]
] + [0] * (self._capture_req_slots - raw_bs)
self._populate_chunked_prefix_buffers(
captured_n=capture_batch.num_prefix_chunks,
prefix_lens_cpu=prefix_lens_cpu,
)
# Kept in sync for backend diagnostics; replay kernels read the stable
# GPU/CPU chunk tensors above, not this Python list.
capture_batch.extend_prefix_lens_cpu = prefix_lens_cpu
self.model_runner.attn_backend.prepare_full_cuda_graph_chunked_prefix(
capture_batch, in_capture=False
)
def _init_forward_metadata_for_capture(
self, forward_batch: ForwardBatch, num_tokens: int
) -> None:
@@ -765,9 +1024,17 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
padded_num_tokens = self._pad_to_bucket(num_tokens, self.capture_num_tokens)
if padded_num_tokens > num_tokens * _MAX_PREFILL_CUDA_GRAPH_PADDING_FACTOR:
return False
# No exact-shape check here: load_batch bucket-pads to the nearest
# captured shape. The factor above only rejects replays whose padded
# model work is disproportionate to the useful token count.
# Other backends and non-MLA FullCG keep using their normal graph with
# replay-refreshed metadata; only this extra topology has a prefix cap.
if (
self._capture_chunked_prefix
and self._has_prefix_hit(forward_batch)
and self._select_prefix_capture_chunks(forward_batch) is None
):
return False
# load_batch bucket-pads to the nearest captured shape. The factor
# above rejects replays whose padded model work is disproportionate
# to the useful token count.
#
# Multi-req replay is supported by body-capture backends via the
# layer_model.forward monkey-patch in replay(): the captured graph runs
@@ -944,8 +1211,11 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
f"Capturing num tokens ({num_tokens=} {avail_mem=:.2f} GB)"
)
self.capture_one_shape(num_tokens)
if self._capture_chunked_prefix:
for captured_n in self._prefix_capture_variants:
self.capture_one_shape(num_tokens, prefix_num_chunks=captured_n)
def capture_one_shape(self, size: int) -> None:
def capture_one_shape(self, size: int, *, prefix_num_chunks: int = 0) -> None:
"""Per-shape capture: build dummy ForwardBatch + run_once,
delegate to backend. size is the prefill token count.
"""
@@ -961,8 +1231,27 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
"limits; the graph would read stale LoRA metadata at replay."
)
lora_manager.prepare_lora_batch(forward_batch)
shape_key = ShapeKey(
size=num_tokens,
variant_label=(
_chunked_prefix_variant(prefix_num_chunks)
if prefix_num_chunks
else None
),
)
if prefix_num_chunks:
self._prepare_chunked_prefix_capture(
forward_batch, shape_key, prefix_num_chunks
)
if self._is_full_backend:
attn_backend.init_forward_metadata_out_graph(forward_batch, in_capture=True)
if not prefix_num_chunks:
attn_backend.init_forward_metadata_out_graph(
forward_batch, in_capture=True
)
# The prefix variant intentionally reuses the capture-stable
# metadata object initialized by the suffix-only variant above.
# Reaching into a backend-specific metadata cache here would make
# this path incompatible with the OSS FlashAttention backend.
else:
self._init_forward_metadata_for_capture(forward_batch, num_tokens)
@@ -984,7 +1273,7 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
else:
post_warmup_hook = getattr(attn_backend, "on_after_cuda_graph_warmup", None)
self.backend.capture_one(
ShapeKey(size=num_tokens),
shape_key,
run_once,
# DP padding can install capture-only tensors on this dummy batch;
# BCG retains it so their recorded addresses remain valid.
@@ -1172,7 +1461,6 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
forward_batch, static_forward_batch, static_num_tokens
)
self._static_num_tokens = static_num_tokens
return static_forward_batch
def _execute_body_capture(
@@ -1181,13 +1469,12 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
static_forward_batch: ForwardBatch,
static_num_tokens: int,
raw_num_tokens: int,
shape_key: ShapeKey,
**kwargs,
):
# BCG / Full: replay the captured body, run the LM head +
# logits_processor eagerly.
shape_key = ShapeKey(size=self._static_num_tokens)
full_path = self._is_full_backend
static_n = self._static_num_tokens
ie_idx = self._input_embeds_arg_idx
def replay_layer_forward(*args, **layer_kwargs):
@@ -1205,7 +1492,7 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
ie = args[ie_idx]
if ie is not None:
self.buffer_registry.get_slot("input_embeds").slice_for(
1, static_n
1, static_num_tokens
)[: ie.shape[0]].copy_(ie)
hs = self.backend.replay(shape_key, static_forward_batch, **kwargs)
return _slice_output_rows(hs, raw_num_tokens) if full_path else hs
@@ -1244,7 +1531,7 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
raw_num_tokens=raw_num_tokens,
):
return self.backend.replay(
ShapeKey(size=self._static_num_tokens),
ShapeKey(size=static_num_tokens),
static_forward_batch,
**kwargs,
)
@@ -1298,6 +1585,10 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
static_forward_batch = self.load_batch(forward_batch, **kwargs)
static_num_tokens = len(static_forward_batch.input_ids)
raw_num_tokens = self.raw_num_tokens
shape_key = self._shape_key(static_num_tokens, forward_batch)
# The only variants this runner records are chunked-prefix ones.
if shape_key.variant_label is not None:
self._prepare_chunked_prefix_replay(shape_key, forward_batch)
if self._uses_eager_prefill_tail():
output = self._execute_body_capture(
@@ -1305,6 +1596,7 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
static_forward_batch,
static_num_tokens,
raw_num_tokens,
shape_key,
**kwargs,
)
else:
@@ -27,8 +27,9 @@ class ShapeKey:
- prefill: num_tokens
- decode: bs
stream_idx: pdmux stream index, or None for single-stream runners.
variant_label: LoRA-variant label ("lora" / "nolora"), or None
for runners that don't record per-variant graphs.
variant_label: optional execution variant (for example, "lora",
"nolora", or "chunked_prefix"), or None for runners that don't
record per-variant graphs.
"""
size: int
@@ -1,19 +1,22 @@
"""Integration test for the full prefill CUDA graph backend.
"""Integration tests for the full prefill CUDA graph backend.
Spins up Qwen3-8B with --cuda-graph-backend-prefill=full and checks
mgsm_en accuracy, mirroring the breakable-CG integration test.
The Qwen3-8B test checks end-to-end accuracy with FlashInfer. The smaller
DeepSeek-Coder-V2-Lite test checks that an MLA radix-prefix hit selects the
OSS FA4 cached-prefix graph variant and matches an eager cold request.
The attention backend is pinned to flashinfer: plain EXTEND under full
CUDA graph requires the backend's init_forward_metadata_out_graph to
support extend (capture-stable plan state). flashinfer and the
FlashAttention backend (fa4; fa3 untested — needs SM90 hardware)
implement it; flashinfer is pinned here for CI-hardware portability
(fa4 requires Blackwell).
implement it; the FA4 case is restricted to Blackwell.
"""
import re
import unittest
from sglang.srt.utils import kill_process_tree
import requests
from sglang.srt.utils import get_device_sm, kill_process_tree
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.run_eval import run_eval
from sglang.test.test_utils import (
@@ -24,8 +27,8 @@ from sglang.test.test_utils import (
popen_launch_server,
)
# CI Registration — large suite to fit the integration test's server startup.
register_cuda_ci(est_time=79, stage="base-b", runner_config="1-gpu-large")
# OSS FA4 coverage requires Blackwell. Each test still uses only one GPU.
register_cuda_ci(est_time=170, stage="base-b", runner_config="4-gpu-b200")
class TestFullCudaGraphPrefill(CustomTestCase):
@@ -65,5 +68,97 @@ class TestFullCudaGraphPrefill(CustomTestCase):
self.assertGreaterEqual(score, 0.80)
@unittest.skipIf(get_device_sm() < 100, "Test requires CUDA SM 100 or higher")
class TestFullCudaGraphChunkedPrefix(unittest.TestCase):
"""A radix-cache hit replays the OSS FA4 FullCG prefix variant."""
@classmethod
def setUpClass(cls):
cls.model = "deepseek-ai/DeepSeek-Coder-V2-Lite-Instruct"
cls.base_url = DEFAULT_URL_FOR_TEST
cls.process = popen_launch_server(
cls.model,
cls.base_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
other_args=[
"--trust-remote-code",
"--prefill-attention-backend=fa4",
"--decode-attention-backend=flashinfer",
"--disable-flashinfer-autotune",
"--context-length=256",
"--max-total-tokens=512",
"--max-running-requests=1",
"--chunked-prefill-size=256",
"--skip-server-warmup",
"--enable-metrics",
"--cuda-graph-config",
'{"decode":{"backend":"disabled"},'
'"prefill":{"backend":"full","bs":[32],"max_bs":32,'
'"full_prefill_max_req":1,'
'"full_prefill_prefix_chunk_tokens":64}}',
],
)
@classmethod
def tearDownClass(cls):
kill_process_tree(cls.process.pid)
def _generate(self, input_ids):
response = requests.post(
self.base_url + "/generate",
json={
"input_ids": input_ids,
"sampling_params": {"max_new_tokens": 4, "temperature": 0},
},
timeout=120,
)
response.raise_for_status()
return response.json()
def _prefill_graph_count(self):
metrics = requests.get(self.base_url + "/metrics", timeout=30).text
match = re.search(
r'^sglang:cuda_graph_passes_total\{[^}]*mode="prefill_cuda_graph"[^}]*\}'
r"\s+([0-9.eE+-]+)$",
metrics,
re.MULTILINE,
)
return float(match.group(1)) if match else 0.0
def test_cached_prefix_replays_full_cuda_graph(self):
prefix = list(range(1000, 1048))
prompt = prefix + list(range(2000, 2032))
requests.post(self.base_url + "/flush_cache", timeout=30).raise_for_status()
cold = self._generate(prompt)
requests.post(self.base_url + "/flush_cache", timeout=30).raise_for_status()
self._generate(prefix)
graph_count = self._prefill_graph_count()
cached = self._generate(prompt)
self.assertEqual(cached["meta_info"]["cached_tokens"], len(prefix))
self.assertEqual(cached["output_ids"], cold["output_ids"])
self.assertEqual(self._prefill_graph_count(), graph_count + 1)
def test_cached_prefix_replays_two_64_token_full_cuda_graph_chunks(self):
prefix = list(range(1000, 1128))
prompt = prefix + list(range(2000, 2032))
# The 160-token cold reference fits in one scheduler prefill chunk, so
# only the 32-token cache-hit suffix is eligible for this FullCG bucket.
requests.post(self.base_url + "/flush_cache", timeout=30).raise_for_status()
cold = self._generate(prompt)
requests.post(self.base_url + "/flush_cache", timeout=30).raise_for_status()
self._generate(prefix)
graph_count = self._prefill_graph_count()
cached = self._generate(prompt)
self.assertEqual(cached["meta_info"]["cached_tokens"], len(prefix))
self.assertEqual(cached["output_ids"], cold["output_ids"])
self.assertEqual(self._prefill_graph_count(), graph_count + 1)
if __name__ == "__main__":
unittest.main()
@@ -31,6 +31,7 @@ class TestMultimodalPiecewiseCudaGraph(CustomTestCase):
runner = PrefillCudaGraphRunner.__new__(PrefillCudaGraphRunner)
runner._is_full_backend = False
runner.enable_lora = False
runner._capture_chunked_prefix = False
runner.prefill_backend_name = backend
runner.has_mha_companion_layers = backend == Backend.BREAKABLE
runner.capture_hidden_mode = CaptureHiddenMode.NULL
@@ -20,6 +20,7 @@ class TestPrefillCudaGraphPadding(CustomTestCase):
runner = PrefillCudaGraphRunner.__new__(PrefillCudaGraphRunner)
runner._is_full_backend = False
runner.enable_lora = False
runner._capture_chunked_prefix = False
runner.prefill_backend_name = Backend.TC_PIECEWISE
runner.has_mha_companion_layers = False
runner.capture_hidden_mode = CaptureHiddenMode.NULL
@@ -0,0 +1,266 @@
"""CPU coverage for chunked-prefix Full prefill CUDA-graph state."""
import unittest
from types import SimpleNamespace
from unittest.mock import patch
import torch
import sglang.srt.model_executor.runner.prefill_cuda_graph_runner as runner_module
from sglang.srt.model_executor.cuda_graph_config import Backend
from sglang.srt.model_executor.runner.prefill_cuda_graph_runner import (
PrefillCudaGraphRunner,
)
from sglang.srt.model_executor.runner.shape_key import ShapeKey
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=2, suite="base-a-test-cpu")
class _FakeAttentionBackend:
supports_full_cuda_graph_chunked_prefix = True
def __init__(self):
self.calls = []
def prepare_full_cuda_graph_chunked_prefix(self, forward_batch, *, in_capture):
self.calls.append((forward_batch, in_capture))
class _FakeKVIndexKernel:
def __getitem__(self, grid):
del grid
def run(
req_to_token,
req_pool_indices,
starts,
seq_lens,
cu_seq_lens,
output,
req_to_token_stride,
):
del cu_seq_lens, req_to_token_stride
cursor = 0
for row in range(seq_lens.numel()):
seq_len = int(seq_lens[row])
start = int(starts[row])
req = int(req_pool_indices[row])
output[cursor : cursor + seq_len].copy_(
req_to_token[req, start : start + seq_len]
)
cursor += seq_len
return run
class TestPrefillCudaGraphRunnerChunkedPrefix(CustomTestCase):
def test_prefix_chunk_capacity_is_aggregate_and_can_be_overridden(self):
model_runner = SimpleNamespace(
server_args=SimpleNamespace(
chunked_prefill_size=16,
context_length=None,
cuda_graph_config=SimpleNamespace(
prefill=SimpleNamespace(
full_prefill_prefix_chunk_tokens=None, max_bs=8
)
),
),
req_to_token_pool=SimpleNamespace(
req_to_token=torch.empty((1, 32), dtype=torch.int32)
),
)
self.assertEqual(
PrefillCudaGraphRunner._resolve_prefix_chunk_shape(model_runner, 4),
(4, 16),
)
model_runner.server_args.chunked_prefill_size = -1
self.assertEqual(
PrefillCudaGraphRunner._resolve_prefix_chunk_shape(model_runner, 4),
(2, 8),
)
model_runner.server_args.chunked_prefill_size = 16
model_runner.server_args.cuda_graph_config.prefill.full_prefill_prefix_chunk_tokens = (
24
)
self.assertEqual(
PrefillCudaGraphRunner._resolve_prefix_chunk_shape(model_runner, 4),
(6, 24),
)
model_runner.server_args.cuda_graph_config.prefill.full_prefill_prefix_chunk_tokens = (
256
)
self.assertEqual(
PrefillCudaGraphRunner._resolve_prefix_chunk_shape(model_runner, 4),
(32, 128),
)
# At least one token is reserved per request lane even if the requested
# aggregate capacity is smaller than the fixed request-slot count.
model_runner.server_args.cuda_graph_config.prefill.full_prefill_prefix_chunk_tokens = (
2
)
self.assertEqual(
PrefillCudaGraphRunner._resolve_prefix_chunk_shape(model_runner, 4),
(1, 4),
)
model_runner.server_args.cuda_graph_config.prefill.full_prefill_prefix_chunk_tokens = (
0
)
with self.assertRaisesRegex(ValueError, "must be positive"):
PrefillCudaGraphRunner._resolve_prefix_chunk_shape(model_runner, 4)
def test_buffers_are_shared_across_token_buckets(self):
backend = _FakeAttentionBackend()
runner = PrefillCudaGraphRunner.__new__(PrefillCudaGraphRunner)
runner._capture_req_slots = 3
runner._prefix_chunk_len = 2
runner._prefix_chunk_capacity = 6
runner._prefix_max_len = 8
runner._prefix_capture_variants = (1, 2, 4)
runner.device = torch.device("cpu")
runner._prefill_static_buffers = {
"extend_prefix_lens": torch.zeros(3, dtype=torch.int64),
"req_pool_indices": torch.tensor([2, 0, 1], dtype=torch.int64),
}
runner._prefix_capture_batches = {}
runner._prefix_capture_buffers = runner._create_chunked_prefix_buffers()
runner.model_runner = SimpleNamespace(
attn_backend=backend,
req_to_token_pool=SimpleNamespace(
req_to_token=torch.arange(24, dtype=torch.int32).view(3, 8)
),
)
first = SimpleNamespace()
second = SimpleNamespace()
first_key = ShapeKey(size=8, variant_label="chunked_prefix:4")
second_key = ShapeKey(size=16, variant_label="chunked_prefix:4")
with patch.object(
runner_module,
"create_chunked_prefix_cache_kv_indices",
_FakeKVIndexKernel(),
):
runner._prepare_chunked_prefix_capture(first, first_key, 4)
runner._prepare_chunked_prefix_capture(second, second_key, 4)
buffers = runner._prefix_capture_buffers
self.assertIsNotNone(buffers)
# Chunk starts are constant and prefilled at allocation.
self.assertEqual(
buffers.starts_cpu.tolist(),
[[0, 0, 0], [2, 2, 2], [4, 4, 4], [6, 6, 6]],
)
self.assertEqual(first.extend_prefix_lens_cpu, [8, 8, 8])
self.assertEqual(first.prefix_chunk_num_tokens, [6, 6, 6, 6])
self.assertIs(first.prefix_chunk_starts, buffers.starts)
self.assertIs(first.prefix_chunk_seq_lens, buffers.seq_lens)
self.assertIs(first.prefix_chunk_cu_seq_lens, buffers.cu_seq_lens)
self.assertIs(first.prefix_chunk_starts, second.prefix_chunk_starts)
self.assertIs(first.prefix_chunk_seq_lens, second.prefix_chunk_seq_lens)
self.assertIs(
first.prefix_chunk_cu_seq_lens,
second.prefix_chunk_cu_seq_lens,
)
# Per-chunk KV indices are views of one shared 2-D buffer; what
# capture bakes into the graph is the address, so compare pointers.
for kv_chunk_idx in (0, 3):
self.assertEqual(
first.prefix_chunk_kv_indices[kv_chunk_idx].data_ptr(),
buffers.kv_indices[kv_chunk_idx].data_ptr(),
)
self.assertEqual(
first.prefix_chunk_kv_indices[kv_chunk_idx].data_ptr(),
second.prefix_chunk_kv_indices[kv_chunk_idx].data_ptr(),
)
runner._prepare_chunked_prefix_replay(
second_key,
SimpleNamespace(batch_size=2, extend_prefix_lens_cpu=[5, 1]),
)
self.assertEqual(
second.prefix_chunk_seq_lens.tolist(),
[[2, 1, 0], [2, 0, 0], [1, 0, 0], [0, 0, 0]],
)
self.assertEqual(
second.prefix_chunk_kv_indices[0].tolist(),
[16, 17, 0, 0, 0, 0],
)
self.assertEqual(
second.prefix_chunk_kv_indices[1].tolist(),
[18, 19, 0, 0, 0, 0],
)
self.assertEqual(
second.prefix_chunk_kv_indices[2].tolist(),
[20, 0, 0, 0, 0, 0],
)
self.assertEqual(second.prefix_chunk_kv_indices[3].tolist(), [0] * 6)
self.assertEqual(
backend.calls,
[(first, True), (second, True), (second, False)],
)
def test_prefix_gate_only_applies_to_chunked_prefix_variant(self):
runner = PrefillCudaGraphRunner.__new__(PrefillCudaGraphRunner)
runner._capture_req_slots = 4
runner.enable_lora = False
runner.capture_hidden_mode = None
runner.max_num_tokens = 32
runner.capture_num_tokens = [4]
runner.backend = SimpleNamespace()
runner.prefill_backend_name = Backend.FULL
runner.has_mha_companion_layers = False
runner._prefix_chunk_len = 2
runner._prefix_capture_variants = (1, 2, 4)
forward_batch = SimpleNamespace(
batch_size=1,
input_ids=torch.zeros(4, dtype=torch.int64),
input_embeds=None,
replace_embeds=None,
forward_mode=SimpleNamespace(is_target_verify=lambda: False),
capture_hidden_mode=None,
global_num_tokens_cpu=None,
return_logprob=False,
extend_prefix_lens_cpu=[8],
)
# Prefix hits in BCG/TC-piecewise and ordinary non-MLA FullCG use the
# normal graph topology and must retain their existing eligibility.
runner._capture_chunked_prefix = False
for is_full_backend in (False, True):
with self.subTest(is_full_backend=is_full_backend):
runner._is_full_backend = is_full_backend
self.assertTrue(runner.can_run_graph(forward_batch))
# The dedicated chunked-prefix topology has a fixed captured capacity.
runner._is_full_backend = True
runner._capture_chunked_prefix = True
self.assertTrue(runner.can_run_graph(forward_batch))
self.assertEqual(
runner._shape_key(4, forward_batch).variant_label,
"chunked_prefix:4",
)
forward_batch.batch_size = 2
# Capacity is per request, not a sum: three real chunks round up to the
# four-chunk graph even though the aggregate prefix has eight tokens.
forward_batch.extend_prefix_lens_cpu = [5, 3]
self.assertTrue(runner.can_run_graph(forward_batch))
self.assertEqual(
runner._shape_key(4, forward_batch).variant_label,
"chunked_prefix:4",
)
forward_batch.extend_prefix_lens_cpu = [9, 1]
self.assertFalse(runner.can_run_graph(forward_batch))
if __name__ == "__main__":
unittest.main()