From 3b8b48279969253e2d0fad43109328b94e766f6d Mon Sep 17 00:00:00 2001 From: Baizhou Zhang Date: Thu, 10 Sep 2026 00:59:18 -0700 Subject: [PATCH] Capture opt-in DSA verify metadata refresh inside CUDA graphs Co-authored-by: Xinyuan Tong Co-authored-by: zRzRzRzRzRzRzR Co-authored-by: Shijin Zhang <75300765+Dovis01@users.noreply.github.com> Co-authored-by: zanes-ops --- .../sglang/kernels/ops/attention/dsv4/topk.py | 19 +- python/sglang/srt/environ.py | 1 + .../attention/dsa/dsa_metadata_ingraph.py | 236 ++++++++++++++++++ .../srt/layers/attention/dsa_backend.py | 13 + python/sglang/test/kits/dsa_metadata_kit.py | 4 +- .../attention/test_dsa_ingraph_metadata.py | 125 ++++++++++ 6 files changed, 392 insertions(+), 6 deletions(-) create mode 100644 python/sglang/srt/layers/attention/dsa/dsa_metadata_ingraph.py create mode 100644 test/registered/kernel/attention/test_dsa_ingraph_metadata.py diff --git a/python/sglang/kernels/ops/attention/dsv4/topk.py b/python/sglang/kernels/ops/attention/dsv4/topk.py index e5f0d0fd1889..de46f84b8c72 100644 --- a/python/sglang/kernels/ops/attention/dsv4/topk.py +++ b/python/sglang/kernels/ops/attention/dsv4/topk.py @@ -75,7 +75,11 @@ def topk_transform_paged( _PLAN_METADATA_INTS_PER_BATCH = 2 -def plan_topk_v2(seq_lens: torch.Tensor, static_threshold: int = 0) -> torch.Tensor: +def plan_topk_v2( + seq_lens: torch.Tensor, + static_threshold: int = 0, + out: Optional[torch.Tensor] = None, +) -> torch.Tensor: """Preprocess the per-batch routing plan for :func:`topk_transform_paged_v2`. IMPORTANT: every entry of ``seq_lens`` must be NON-NEGATIVE. The device @@ -84,12 +88,19 @@ def plan_topk_v2(seq_lens: torch.Tensor, static_threshold: int = 0) -> torch.Ten the plan, and drives the transform kernel into an illegal memory access. Producers of padded rows must clamp their lengths to 0 (0 selects the trivial all-(-1) output path, which is safe). + + Pass a preallocated ``out`` plan to refresh CUDA-graph metadata in place. """ module = _jit_topk_v2_module() bs = seq_lens.shape[0] - metadata = seq_lens.new_empty(bs + 1, _PLAN_METADATA_INTS_PER_BATCH) - module.topk_plan(seq_lens, metadata, static_threshold) - return metadata + if out is None: + out = seq_lens.new_empty(bs + 1, _PLAN_METADATA_INTS_PER_BATCH) + else: + assert out.shape == (bs + 1, _PLAN_METADATA_INTS_PER_BATCH) + assert out.dtype == seq_lens.dtype and out.device == seq_lens.device + assert out.is_contiguous() + module.topk_plan(seq_lens, out, static_threshold) + return out def topk_transform_ragged_v2( diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index db80e167d337..98d182d24634 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -1546,6 +1546,7 @@ class Envs: ) # Enabled for supported CUDA KPool geometry; set to 0 to use ordinary metadata. SGLANG_EXPERIMENTAL_DSA_KPOOL_METADATA_FUSION = EnvBool(True) + SGLANG_EXPERIMENTAL_DSA_INGRAPH_VERIFY_METADATA = EnvBool(False) SGLANG_DSA_TOPK_FLASHINFER_DETERMINISTIC = EnvBool(False) SGLANG_DSA_TOPK_FLASHINFER_TIE_BREAK = EnvStr(None) SGLANG_DSA_PREFILL_DENSE_ATTN_KV_LEN_THRESHOLD = EnvIntWithAlias( diff --git a/python/sglang/srt/layers/attention/dsa/dsa_metadata_ingraph.py b/python/sglang/srt/layers/attention/dsa/dsa_metadata_ingraph.py new file mode 100644 index 000000000000..2254320d0e66 --- /dev/null +++ b/python/sglang/srt/layers/attention/dsa/dsa_metadata_ingraph.py @@ -0,0 +1,236 @@ +"""Capture DSA target-verify metadata refreshes into the model CUDA graph.""" + +from __future__ import annotations + +import logging +from functools import cache +from typing import TYPE_CHECKING, Tuple + +import torch + +from sglang.srt.environ import envs +from sglang.srt.model_executor.forward_batch_info import ForwardMode +from sglang.srt.runtime_context import get_parallel, get_platform +from sglang.srt.utils import is_cuda, is_hip + +if TYPE_CHECKING: + from sglang.srt.layers.attention.dsa_backend import DSAMetadata + from sglang.srt.model_executor.forward_batch_info import ForwardBatch + +if is_cuda(): + import deep_gemm + +_is_hip = is_hip() +logger = logging.getLogger(__name__) + + +@cache +def get_plan_topk_v2(): + # Resolving the JIT module at srt import time would initialize CUDA too early. + from sglang.kernels.ops.attention.dsv4.topk import plan_topk_v2 + + return plan_topk_v2 + + +class _DSAInGraphVerifyMetadataState: + """Captured refreshes require fixed input-buffer identities at replay. + Mismatches must fail because captured nodes would overwrite fallback metadata.""" + + __slots__ = ( + "prefill_impl_state", + "_seq_lens_ptr", + "_seq_lens_stride", + "_req_pool_indices_ptr", + "_req_pool_indices_stride", + ) + + def __init__( + self, + *, + prefill_impl_state: Tuple[bool, str], + seq_lens: torch.Tensor, + req_pool_indices: torch.Tensor, + ): + self.prefill_impl_state = prefill_impl_state + self._seq_lens_ptr = seq_lens.data_ptr() + self._seq_lens_stride = seq_lens.stride(0) + self._req_pool_indices_ptr = req_pool_indices.data_ptr() + self._req_pool_indices_stride = req_pool_indices.stride(0) + + def matches(self, seq_lens: torch.Tensor, req_pool_indices: torch.Tensor) -> bool: + # A [:bs] slice keeps the base data_ptr/stride, so the raw incoming + # graph-runner buffers compare equal to the sliced views recorded at + # capture time. + return ( + seq_lens.data_ptr() == self._seq_lens_ptr + and req_pool_indices.data_ptr() == self._req_pool_indices_ptr + and seq_lens.stride(0) == self._seq_lens_stride + and req_pool_indices.stride(0) == self._req_pool_indices_stride + ) + + +class DSAInGraphVerifyMetadataMixin: + ingraph_verify_metadata_enabled = False + + def _init_ingraph_verify_metadata(self): + self.ingraph_verify_metadata_enabled = ( + envs.SGLANG_EXPERIMENTAL_DSA_INGRAPH_VERIFY_METADATA.get() + ) + + def _replay_ingraph_verify_metadata( + self, metadata, seq_lens, req_pool_indices, forward_mode + ): + if not forward_mode.is_target_verify(): + return False + state = getattr(metadata, "_ingraph_verify_metadata", None) + if state is None: + return False + if not state.matches(seq_lens, req_pool_indices): + raise RuntimeError( + "DSA in-graph verify metadata replay requires the static buffers " + "for seq_lens and req_pool_indices used during capture" + ) + self.use_mha, self.dsa_prefill_impl = state.prefill_impl_state + self.forward_metadata = metadata + return True + + def _ingraph_verify_metadata_eligible(self) -> bool: + """Config-level eligibility for the in-graph TARGET_VERIFY refresh. + + The recorded nodes reproduce the fused verify sequence. Configurations + that would take the unfused fallback + (kpool > 1 without the fusion gate), DCP, HIP, or flashmla_kv (whose + per-replay metadata recompute is host-side) stay out-of-graph. + """ + return ( + self.ingraph_verify_metadata_enabled + and is_cuda() + and not _is_hip + and not get_parallel().dcp_enabled + and (self.dsa_index_kpool <= 1 or self.experimental_kpool_metadata_fusion) + and self.dsa_decode_impl != "flashmla_kv" + ) + + def init_forward_metadata_in_graph(self, forward_batch: ForwardBatch): + """Record target-verify refreshes only for capture-safe configurations. + Replay reads the same static input buffers with updated contents.""" + if not forward_batch.forward_mode.is_target_verify(): + return + if not self._ingraph_verify_metadata_eligible(): + return + bs = forward_batch.batch_size + metadata = self.decode_cuda_graph_metadata.get(bs) + if metadata is None: + return + self._record_ingraph_verify_metadata( + metadata, + bs, + forward_batch.seq_lens[:bs], + forward_batch.req_pool_indices[:bs], + ) + + def _record_ingraph_verify_metadata( + self, + metadata: DSAMetadata, + bs: int, + seq_lens: torch.Tensor, + req_pool_indices: torch.Tensor, + ) -> None: + """Publish replay state only after every structural check passes. + Otherwise the generic out-of-graph path remains responsible.""" + if metadata.paged_mqa_schedule_metadata is None: + return + next_n = self.speculative_num_draft_tokens + if not next_n: + return + expanded_size = bs * next_n + max_seqlen_k = self._graph_page_table_width(metadata) + + paged_mqa_ctx_lens_2d = None + if ( + next_n >= 2 + and get_platform().is_sm100 + and metadata.paged_mqa_ctx_lens_2d is not None + and metadata.paged_mqa_ctx_lens_2d.dim() == 2 + and metadata.paged_mqa_ctx_lens_2d.size(0) == bs + and metadata.paged_mqa_ctx_lens_2d.size(1) == next_n + ): + paged_mqa_ctx_lens_2d = metadata.paged_mqa_ctx_lens_2d + ctx_lens_written = paged_mqa_ctx_lens_2d is not None + + if ctx_lens_written: + # DG-native layout: the fused kernel writes ctx lens straight + # into the captured (bs, next_n) buffer; the schedule reads it. + schedule_src_2d = metadata.paged_mqa_ctx_lens_2d + ctx_lens_copy_src = None + else: + if ( + next_n >= 2 and get_platform().is_sm100 + ): # Degenerate capture (DG-native ctx-lens layout expected but + # missing); keep the whole refresh out-of-graph. + return + seqlens_view = metadata.dsa_seqlens_expanded[:expanded_size] + if not seqlens_view.is_contiguous(): + return + schedule_src_2d = seqlens_view.view(-1, 1) + if metadata.paged_mqa_ctx_lens_2d is None: + # Verify capture always materializes the 2D ctx-lens buffer; + # if missing, leave the object.__setattr__ publication to + # the generic path. + return + ctx_lens_copy_src = schedule_src_2d + + self._fused_verify_metadata( + seq_lens=seq_lens, + req_pool_indices=req_pool_indices, + req_to_token=self.req_to_token, + cache_seqlens=metadata.cache_seqlens_int32, + cu_seqlens_k=metadata.cu_seqlens_k, + page_table_1=metadata.page_table_1, + seqlens_expanded=metadata.dsa_seqlens_expanded, + dsa_cache_seqlens=metadata.dsa_cache_seqlens_int32, + dsa_cu_seqlens_k=metadata.dsa_cu_seqlens_k, + real_page_table=metadata.real_page_table, + bs=bs, + max_seqlen_k=max_seqlen_k, + dsa_index_topk=self.dsa_index_topk, + real_page_size=self.real_page_size, + next_n=next_n, + paged_mqa_ctx_lens_2d=paged_mqa_ctx_lens_2d, + ) + + metadata.paged_mqa_schedule_metadata.copy_( + deep_gemm.get_paged_mqa_logits_metadata( + schedule_src_2d, 64, deep_gemm.get_num_sms() + ) + ) + if ctx_lens_copy_src is not None: + metadata.paged_mqa_ctx_lens_2d.copy_(ctx_lens_copy_src) + + if metadata.topk_v2_plan is not None: + # A preallocated output avoids allocation during capture. + get_plan_topk_v2()(metadata.dsa_seqlens_expanded, out=metadata.topk_v2_plan) + + self._update_kpool_metadata_replay( + metadata, + seq_lens, + req_pool_indices, + ForwardMode.TARGET_VERIFY, + ) + + self.set_dsa_prefill_impl(forward_batch=None) + state = _DSAInGraphVerifyMetadataState( + prefill_impl_state=(self.use_mha, self.dsa_prefill_impl), + seq_lens=seq_lens, + req_pool_indices=req_pool_indices, + ) + # Undeclared attribute on the frozen dataclass: dropped by both + # recapture (fresh DSAMetadata) and dataclasses.replace(). + object.__setattr__(metadata, "_ingraph_verify_metadata", state) + logger.info( + "DSA in-graph verify metadata recorded: bs=%d next_n=%d " + "ctx_lens_written=%s", + bs, + next_n, + ctx_lens_written, + ) diff --git a/python/sglang/srt/layers/attention/dsa_backend.py b/python/sglang/srt/layers/attention/dsa_backend.py index 8fdd7ba12dda..d75b1cf8f83c 100644 --- a/python/sglang/srt/layers/attention/dsa_backend.py +++ b/python/sglang/srt/layers/attention/dsa_backend.py @@ -58,6 +58,9 @@ compute_cu_seqlens, ) from sglang.srt.layers.attention.dsa.dsa_indexer_metadata import DSAIndexerMetadata +from sglang.srt.layers.attention.dsa.dsa_metadata_ingraph import ( + DSAInGraphVerifyMetadataMixin, +) from sglang.srt.layers.attention.dsa.dsa_metadata_manager import ( DSAMetadataManagementMixin, ) @@ -297,6 +300,7 @@ def _cat(tensors: list[torch.Tensor], dim: int = -1) -> torch.Tensor: class DeepseekSparseAttnBackend( DSAMetadataManagementMixin, + DSAInGraphVerifyMetadataMixin, DeepseekSparseAttnBackendKPoolMixin, DeepseekSparseAttnBackendMTPPrecomputeMixin, AttentionBackend, @@ -336,6 +340,7 @@ def __init__( self.dsa_index_kpool = get_dsa_index_kpool(hf_config) self.needs_cpu_seq_lens = self.dsa_index_kpool > 1 self._init_kpool_metadata_fusion() + self._init_ingraph_verify_metadata() self.max_context_len = model_runner.model_config.context_len self.num_q_heads = ( model_runner.model_config.num_attention_heads // get_parallel().attn_tp_size @@ -814,6 +819,10 @@ def init_forward_metadata_out_graph( forward_batch: ForwardBatch, in_capture: bool = False, ): + if in_capture: + metadata = self.decode_cuda_graph_metadata.get(forward_batch.batch_size) + if metadata is not None and hasattr(metadata, "_ingraph_verify_metadata"): + object.__delattr__(metadata, "_ingraph_verify_metadata") seq_lens_cpu = ( forward_batch.seq_lens.cpu() if in_capture else forward_batch.seq_lens_cpu ) @@ -1507,6 +1516,10 @@ def _apply_cuda_graph_metadata( metadata: DSAMetadata = self.decode_cuda_graph_metadata[bs] + if self._replay_ingraph_verify_metadata( + metadata, seq_lens, req_pool_indices, forward_mode + ): + return self.set_dsa_prefill_impl(forward_batch=None) seq_lens = seq_lens[:bs] diff --git a/python/sglang/test/kits/dsa_metadata_kit.py b/python/sglang/test/kits/dsa_metadata_kit.py index a3e5c36293bd..fc59254b0763 100644 --- a/python/sglang/test/kits/dsa_metadata_kit.py +++ b/python/sglang/test/kits/dsa_metadata_kit.py @@ -26,14 +26,14 @@ def inputs(lengths, requests): ) -def make_backend(mode, seq, req, *, fusion=True): +def make_backend(mode, seq, req, *, fusion=True, pool_size=POOL): backend = object.__new__(DeepseekSparseAttnBackend) backend.device = torch.device("cuda") backend.device_sm_major = torch.cuda.get_device_capability()[0] backend.num_q_heads = 64 backend.real_page_size = 64 backend.dsa_index_topk = TOPK - backend.dsa_index_kpool = POOL + backend.dsa_index_kpool = pool_size backend.speculative_num_draft_tokens = NEXT_N backend.dsa_drop_wide_page_table = False backend.dsa_decode_impl = "fa3" diff --git a/test/registered/kernel/attention/test_dsa_ingraph_metadata.py b/test/registered/kernel/attention/test_dsa_ingraph_metadata.py new file mode 100644 index 000000000000..a5dba849a87d --- /dev/null +++ b/test/registered/kernel/attention/test_dsa_ingraph_metadata.py @@ -0,0 +1,125 @@ +"""Captured TARGET_VERIFY metadata follows the graph runner's static inputs.""" + +import unittest +from types import SimpleNamespace + +import torch + +from sglang.srt.model_executor.forward_batch_info import ForwardMode +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.kits.dsa_metadata_kit import ( + BS, + POOL, + ROUNDS, + addresses, + apply_metadata, + assert_metadata_equal, + capture_verify_metadata, + inputs, + make_backend, +) +from sglang.test.test_utils import CustomTestCase + +register_cuda_ci(est_time=30, stage="base-b-kernel-unit", runner_config="1-gpu-large") + + +class TestDSAInGraphMetadata(CustomTestCase): + def test_captured_verify_refreshes_all_derived_buffers(self): + self._check_captured_verify_refresh(pool_size=POOL) + + def test_captured_verify_without_kpool(self): + self._check_captured_verify_refresh(pool_size=1) + + def _check_captured_verify_refresh(self, pool_size): + mode = ForwardMode.TARGET_VERIFY + seq, req = inputs(*ROUNDS[0]) + backend = make_backend( + mode, seq, req, fusion=pool_size > 1, pool_size=pool_size + ) + ordinary = make_backend(mode, seq, req, fusion=False, pool_size=pool_size) + if pool_size == 1: + self.assertFalse(backend.experimental_kpool_metadata_fusion) + self.assertIsNone(backend.forward_metadata.kpool_write_plan) + pointers = addresses(backend.forward_metadata) + graph = capture_verify_metadata(backend, seq, req) + self.assertIsNotNone( + getattr(backend.forward_metadata, "_ingraph_verify_metadata", None), + "the verify hook must install captured replay state", + ) + for lengths, requests in ROUNDS[1:]: + seq.copy_(torch.tensor(lengths, device="cuda")) + req.copy_(torch.tensor(requests, device="cuda")) + # Eager replay must leave captured metadata refresh to the graph. + before = backend.forward_metadata.dsa_seqlens_expanded.clone() + apply_metadata(backend, mode, seq, req) + torch.testing.assert_close( + backend.forward_metadata.dsa_seqlens_expanded, before + ) + graph.replay() + apply_metadata(ordinary, mode, seq, req) + assert_metadata_equal( + self, backend.forward_metadata, ordinary.forward_metadata + ) + self.assertEqual(pointers, addresses(backend.forward_metadata)) + + def test_recapture_replaces_static_input_identities(self): + mode = ForwardMode.TARGET_VERIFY + old_seq, old_req = inputs(*ROUNDS[0]) + backend = make_backend(mode, old_seq, old_req) + old_graph = capture_verify_metadata(backend, old_seq, old_req) + old_graph.replay() + old_state = getattr(backend.forward_metadata, "_ingraph_verify_metadata", None) + self.assertIsNotNone(old_state) + + seq, req = inputs(*ROUNDS[1]) + batch = SimpleNamespace( + batch_size=BS, + forward_mode=mode, + seq_lens=seq, + req_pool_indices=req, + spec_info=None, + out_cache_loc=None, + ) + # The public capture preparation must retire the previous alias guard + # before refreshing metadata from a new graph runner's buffers. + backend.init_forward_metadata_out_graph(batch, in_capture=True) + self.assertIsNone( + getattr(backend.forward_metadata, "_ingraph_verify_metadata", None) + ) + graph = capture_verify_metadata(backend, seq, req) + state = getattr(backend.forward_metadata, "_ingraph_verify_metadata", None) + self.assertIsNotNone(state) + self.assertIsNot(state, old_state) + pointers = addresses(backend.forward_metadata) + ordinary = make_backend(mode, seq, req, fusion=False) + + lengths, requests = ROUNDS[2] + seq.copy_(torch.tensor(lengths, device="cuda")) + req.copy_(torch.tensor(requests, device="cuda")) + apply_metadata(backend, mode, seq, req) + graph.replay() + apply_metadata(ordinary, mode, seq, req) + assert_metadata_equal(self, backend.forward_metadata, ordinary.forward_metadata) + self.assertEqual(pointers, addresses(backend.forward_metadata)) + with self.assertRaisesRegex(RuntimeError, "alias|static buffers"): + apply_metadata(backend, mode, old_seq, old_req) + + def test_replay_rejects_replaced_input_buffers(self): + mode = ForwardMode.TARGET_VERIFY + seq, req = inputs(*ROUNDS[0]) + backend = make_backend(mode, seq, req) + graph = capture_verify_metadata(backend, seq, req) + self.assertIsNotNone( + getattr(backend.forward_metadata, "_ingraph_verify_metadata", None) + ) + for bad_seq, bad_req in ((seq.clone(), req), (seq, req.clone())): + with self.subTest(seq_replaced=bad_seq is not seq): + with self.assertRaisesRegex(RuntimeError, "alias|static buffers"): + apply_metadata(backend, mode, bad_seq, bad_req) + # Views with the same pointer, shape and stride remain valid. + apply_metadata(backend, mode, seq[:], req[:]) + graph.replay() + + +if __name__ == "__main__": + unittest.main()