[unified-memory] Enable prefill cuda-graph capture (#37418)

This commit is contained in:
Cheng Wan
2026-09-13 19:17:19 -07:00
committed by GitHub
parent 2ec4bbcbd4
commit 6410800af9
5 changed files with 154 additions and 60 deletions
+26 -17
View File
@@ -262,25 +262,34 @@ def handle_unified_memory_pool(server_args: Any) -> None:
)
if cfg.dcp_size > 1:
_validate_unified_memory_dcp(server_args)
# Only monolithic decode cuda-graph capture is wired; piecewise prefill
# capture is not. Guard when the user opts into it.
# Prefill cuda-graph capture IS wired for the unified pool: the captured
# batch reads `out_cache_loc` out of the registry slot, which
# `populate_from_forward_batch` refills from the already-rebound (kernel-
# facing) loc before every replay, and the read tables are refilled
# out-of-graph from the live v2p.
#
# The FULL backend is the one exception, and not for a unified reason: its
# metadata path (`_init_full_cg_prefill_metadata`) exists only on the
# fa3/fa4 family. Any other backend lands in the decode-shaped
# `_apply_cuda_graph_metadata`, which has no EXTEND branch at all. Inkling
# declares FULL as a MODEL default, indistinguishable here from a flag the
# user typed, so warn and fall back rather than refuse to boot.
_cg_cfg = cfg.cuda_graph_config
if _cg_cfg is not None and _cg_cfg.prefill.backend != Backend.DISABLED:
if cfg.cuda_graph_backend_prefill is not None:
raise ValueError(
"--enable-unified-memory supports decode cuda-graph "
"capture only; prefill capture is not wired (the prefill "
"graph runner bypasses the unified virtual->physical loc "
"rebind). Got --cuda-graph-backend-prefill="
f"{cfg.cuda_graph_backend_prefill!r}; pass "
"--cuda-graph-backend-prefill=disabled."
if _cg_cfg is not None and _cg_cfg.prefill.backend == Backend.FULL:
full_cg_backends = {"fa3", "fa4"}
backends = set(attention_backends_of(resolved_view(server_args)))
backends.discard(None)
if not backends <= full_cg_backends:
_cg_cfg.prefill.backend = Backend.DISABLED
logger.warning(
"--enable-unified-memory: disabling the FULL prefill "
"cuda-graph backend. It builds its block table in "
"_init_full_cg_prefill_metadata, which only %s implement; the "
"resolved attention backends are %s. Decode capture and the "
"other prefill backends are unaffected.",
sorted(full_cg_backends),
sorted(backends),
)
_cg_cfg.prefill.backend = Backend.DISABLED
logger.warning(
"--enable-unified-memory: disabling prefill cuda-graph "
"capture (not wired for the unified pool's loc rebind); "
"decode capture is unaffected."
)
def _validate_unified_memory_dcp(server_args: Any) -> None:
@@ -26,6 +26,7 @@ from sglang.srt.layers.attention.verify_mask import VerifyMask, maybe_create_ver
from sglang.srt.layers.cp.base import CPAttentionBackendKind, get_cp_strategy
from sglang.srt.layers.cp.utils import is_cp_active
from sglang.srt.layers.radix_attention import AttentionType
from sglang.srt.mem_cache.kv_index_translator import KVReadTables
from sglang.srt.mem_cache.memory_pool import KVWriteLoc
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
@@ -673,8 +674,26 @@ class FlashAttentionBackend(AttentionBackend):
m.cu_seqlens_q[1:].copy_(
torch.cumsum(forward_batch.extend_seq_lens[:bs], dim=0)
)
translating = self.kv_index_translator.is_translating
max_seq_len_k = int(forward_batch.seq_lens_cpu[:bs].max().item())
if max_seq_len_k > 0:
if translating:
# Unified pool: the block table is a TRANSLATED page table, built
# straight into these capture-stable buffers from the LIVE v2p, so
# a page relocated by compaction since capture is picked up. Same
# substitution the eager extend branch makes in its `_unified_read`
# fixup; `build_index_table` emits page-granular kernel-facing ids
# directly, so there is no `// page_size` to undo.
self.kv_index_translator.build_index_table(
req_pool_indices=forward_batch.req_pool_indices[:bs],
seq_lens=forward_batch.seq_lens[:bs],
into=KVReadTables(
full=m.page_table,
sliding_window=(
m.swa_page_table if self.use_sliding_window_kv_pool else None
),
),
)
elif max_seq_len_k > 0:
# Build the block table like the eager extend branch: take every
# page_size-th token slot from req_to_token and divide by page_size.
# Identity for page_size == 1 (strided is 0..max_seq_len_k-1, //1).
@@ -700,11 +719,21 @@ class FlashAttentionBackend(AttentionBackend):
self.full_cg_prefill_swa_out_cache_loc.shape[0],
"full-CG prefill SWA write-location buffer",
)
self.full_cg_prefill_swa_out_cache_loc[:num_out].copy_(
self.token_to_kv_pool.translate_loc_from_full_to_swa(
# Under the unified pool `out_cache_loc` was rebound to FULL-side
# KERNEL-FACING ids at ForwardBatch construction, so the full->swa
# map cannot be re-run on it -- those values index far past the swa
# v2p table (a device-side "index out of bounds" assert). Phase 2 of
# the write contract derives the swa loc from them instead.
swa_write_loc = (
self.kv_index_translator.sliding_window_write_loc_for(
forward_batch.out_cache_loc
)
if translating
else self.token_to_kv_pool.translate_loc_from_full_to_swa(
forward_batch.out_cache_loc
)
)
self.full_cg_prefill_swa_out_cache_loc[:num_out].copy_(swa_write_loc)
# Captured kernels read the full bucket. Route its inactive tail to
# SWA's zero dummy slot to prevent stale writes into live slots.
self.full_cg_prefill_swa_out_cache_loc[num_out:].zero_()
@@ -51,7 +51,7 @@ _MODEL_PATH = os.environ.get("INKLING_TEST_MODEL_PATH", "thinkingmachines/Inklin
_MODEL_REVISION = os.environ.get("INKLING_TEST_MODEL_REVISION", "test")
def _unified_args():
def _unified_args(*, attention_backend="triton", prefill_cuda_graph=False):
"""Server args for the tri-pool boot. Mirrors test_inkling.py's fixture
minus the multimodal/parser surface (KV-path focus), plus the unified
flags. The ratios still feed boot sizing until the byte configurator
@@ -59,17 +59,12 @@ def _unified_args():
args = [
"--trust-remote-code",
"--enable-unified-memory",
# Unified requires the Triton strided page-major read/write paths.
"--attention-backend",
"triton",
attention_backend,
"--page-size",
"128",
"--mamba-radix-cache-strategy",
"extra_buffer",
# Inkling defaults to a FULL prefill graph, which unified rejects at
# boot: the prefill graph runner bypasses the virtual->physical rebind.
"--cuda-graph-backend-prefill",
"disabled",
"--swa-full-tokens-ratio",
"0.1",
"--mamba-full-memory-ratio",
@@ -77,6 +72,11 @@ def _unified_args():
"--mem-fraction-static",
"0.5",
]
if not prefill_cuda_graph:
# Inkling declares a FULL prefill graph as a model default; the Triton
# cells cannot serve it (the cuda-graph metadata path has no EXTEND
# branch), so pin it off rather than lean on the auto-fallback.
args += ["--cuda-graph-backend-prefill", "disabled"]
if _MODEL_REVISION:
args += ["--revision", _MODEL_REVISION]
return args
@@ -91,8 +91,8 @@ def _static_args():
"128",
"--mamba-radix-cache-strategy",
"extra_buffer",
# Inkling defaults to a FULL prefill graph, which unified rejects at
# boot: the prefill graph runner bypasses the virtual->physical rebind.
# Match the unified cell's Triton pin, which cannot serve Inkling's
# default FULL prefill graph.
"--cuda-graph-backend-prefill",
"disabled",
"--swa-full-tokens-ratio",
@@ -128,6 +128,10 @@ def _greedy_generate(base_url, text, max_new_tokens=32, logprobs=False):
class TestInklingUnifiedTriPool(CustomTestCase):
@classmethod
def server_args(cls):
return _unified_args()
@classmethod
def setUpClass(cls):
cls.model = _MODEL_PATH
@@ -136,7 +140,7 @@ class TestInklingUnifiedTriPool(CustomTestCase):
cls.model,
cls.base_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
other_args=_unified_args(),
other_args=cls.server_args(),
env={**os.environ, "SGLANG_ENABLE_UNIFIED_RADIX_TREE": "1"},
)
@@ -194,6 +198,28 @@ class TestInklingUnifiedTriPool(CustomTestCase):
self.assertGreater(len(data["text"].strip()), 0, data)
class TestInklingUnifiedFullPrefillGraph(TestInklingUnifiedTriPool):
"""The same tri-pool guards with Inkling's OWN default FULL prefill cuda
graph left on, over fa4 -- the backend family whose
`_init_full_cg_prefill_metadata` implements that path.
The pairing used to be a hard boot failure: unified disabled prefill
capture outright. With capture on, both the captured block table and the
SWA write loc have to come from the translator. Re-running the full->swa
map on `out_cache_loc` does not work here -- it is already FULL-side
kernel-facing by then, and indexes far past the swa v2p table (a
device-side "index out of bounds" assert).
`test_input_output_logprobs_match` is the sharp guard: a wrong-slot SWA
write moves logprobs at once, and
`test_long_decode_slides_past_swa_window` keeps compaction running
underneath a replaying graph.
"""
@classmethod
def server_args(cls):
return _unified_args(attention_backend="fa4", prefill_cuda_graph=True)
@unittest.skipUnless(
os.environ.get("INKLING_UNIFIED_PARITY") == "1",
"eval-host lane: set INKLING_UNIFIED_PARITY=1 (two sequential server boots)",
@@ -24,6 +24,17 @@ class TestFlashAttentionGraphMetadata(CustomTestCase):
backend.token_to_kv_pool = SimpleNamespace(
translate_loc_from_full_to_swa=lambda locations: locations
)
# The metadata builder reads `is_translating` to choose between the
# translated block table and the strided one this test covers, so the
# source has to be real; the stub pools disable translation, which is
# the static-pool view the assertions below are written against.
backend.kv_index_translator = KVIndexTranslator(
req_to_token=backend.req_to_token,
token_to_kv_pool_allocator=SimpleNamespace(),
token_to_kv_pool=SimpleNamespace(),
page_size=backend.page_size,
device="cpu",
)
forward_batch = SimpleNamespace(
batch_size=1,
seq_lens=torch.zeros(1, dtype=torch.int64),
@@ -11,26 +11,31 @@
# See the License for the specific language governing permissions and
# limitations under the License.
# ==============================================================================
"""`--enable-unified-memory` disables PREFILL cuda-graph capture.
"""`--enable-unified-memory` and PREFILL cuda-graph capture.
BUG REGRESSION. Only decode capture is wired: the prefill graph runner builds
its ForwardBatch directly, so it never runs the unified pool's write-loc
rebind (rebind_write_loc) and the captured batch holds VIRTUAL ids -- the
captured store would silently write wrong slots.
Capture is wired: the captured batch reads `out_cache_loc` out of the registry
slot, refilled before each replay from the already-rebound kernel-facing loc,
and the read tables are refilled out-of-graph from the live v2p. So BREAKABLE
(the CUDA default) and TC_PIECEWISE must be left alone -- an earlier gate
disabled every prefill backend outright, which cost every unified run its
prefill graph.
The old gate only rejected `TC_PIECEWISE`, but the generic prefill default is
`BREAKABLE` -- so the DEFAULT unified invocation was broken; it only ever
worked when `--cuda-graph-backend-prefill=disabled` happened to be passed.
The FULL backend is the exception, and for a reason that is not about unified
memory: its metadata path (`_init_full_cg_prefill_metadata`) is implemented
only by the fa3/fa4 family. Anything else lands in the decode-shaped
`_apply_cuda_graph_metadata`, which has no EXTEND branch.
Pinned: the default is auto-disabled with a warning (unified boots out of the
box), an EXPLICIT prefill backend still raises (never silently override a
user's stated intent), and decode capture is untouched either way.
Pinned here: FULL survives on fa3/fa4, FULL is disabled with a warning on any
other backend (Inkling declares FULL as a MODEL default, so refusing to boot
would fail on a flag the user never typed), and decode capture is never
touched.
python -m pytest test/registered/unit/server_args/test_unified_prefill_cuda_graph_gate.py -v
"""
import unittest
from types import SimpleNamespace
from unittest.mock import patch
import msgspec
@@ -42,7 +47,7 @@ from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=8, suite="base-a-test-cpu")
def _run_handler(*, prefill_backend, explicit):
def _run_handler(*, prefill_backend, attention_backends):
"""Run just `handle_unified_memory_pool` over a minimal stand-in."""
sa = ServerArgs(model_path="dummy")
cg = SimpleNamespace(
@@ -59,33 +64,47 @@ def _run_handler(*, prefill_backend, explicit):
"enable_two_batch_overlap": False,
"dcp_size": 1,
"cuda_graph_config": cg,
"cuda_graph_backend_prefill": prefill_backend if explicit else None,
"cuda_graph_backend_prefill": prefill_backend,
}.items():
msgspec.Struct.__setattr__(sa, name, value)
handle_unified_memory_pool(sa)
with patch(
"sglang.srt.arg_groups.kv_cache_hook.attention_backends_of",
return_value=attention_backends,
):
handle_unified_memory_pool(sa)
return cg
class TestUnifiedPrefillCudaGraphGate(unittest.TestCase):
def test_default_prefill_capture_is_auto_disabled(self):
"""The generic default (BREAKABLE) must be turned off, not crash the
server 30 seconds later inside graph capture."""
for backend in (Backend.BREAKABLE, Backend.FULL, Backend.TC_PIECEWISE):
cg = _run_handler(prefill_backend=backend, explicit=False)
self.assertEqual(cg.prefill.backend, Backend.DISABLED)
# Decode capture is the wired path and must survive untouched.
self.assertEqual(cg.decode.backend, Backend.FULL)
def test_non_full_prefill_backends_are_left_enabled(self):
"""BUG REGRESSION. Unified used to disable prefill capture outright, so
the default BREAKABLE graph silently never ran."""
for backend in (Backend.BREAKABLE, Backend.TC_PIECEWISE):
for attn in (("fa4", "fa4"), ("triton", "triton")):
with self.subTest(prefill=backend, attn=attn):
cg = _run_handler(prefill_backend=backend, attention_backends=attn)
self.assertEqual(cg.prefill.backend, backend)
self.assertEqual(cg.decode.backend, Backend.FULL)
def test_explicit_prefill_backend_is_refused(self):
"""A user who explicitly asked for prefill graphs gets a clear error,
not a silent override of their stated intent."""
for backend in (Backend.BREAKABLE, Backend.FULL, Backend.TC_PIECEWISE):
with self.assertRaises(ValueError) as ctx:
_run_handler(prefill_backend=backend, explicit=True)
self.assertIn("prefill capture is not wired", str(ctx.exception))
def test_full_prefill_survives_on_the_fa_family(self):
for attn in (("fa3", "fa3"), ("fa4", "fa4")):
with self.subTest(attn=attn):
cg = _run_handler(prefill_backend=Backend.FULL, attention_backends=attn)
self.assertEqual(cg.prefill.backend, Backend.FULL)
def test_full_prefill_is_disabled_on_other_backends(self):
"""Warn and fall back rather than raise: Inkling declares FULL as a
model default, indistinguishable at this point from a user flag."""
for attn in (("triton", "triton"), ("flashinfer", "flashinfer")):
with self.subTest(attn=attn):
cg = _run_handler(prefill_backend=Backend.FULL, attention_backends=attn)
self.assertEqual(cg.prefill.backend, Backend.DISABLED)
self.assertEqual(cg.decode.backend, Backend.FULL)
def test_already_disabled_is_a_no_op(self):
cg = _run_handler(prefill_backend=Backend.DISABLED, explicit=True)
cg = _run_handler(
prefill_backend=Backend.DISABLED, attention_backends=("triton", "triton")
)
self.assertEqual(cg.prefill.backend, Backend.DISABLED)
self.assertEqual(cg.decode.backend, Backend.FULL)