[AMD] [GLM5] Skip DSA decode indexer when kv_len <= index_topk (dense k-only fast path) (#31324)

Co-authored-by: Thomas Wang <thomawan@amd.com>
Co-authored-by: HaiShaw <hixiao@gmail.com>
This commit is contained in:
Jacob0226
2026-08-16 16:30:54 -07:00
committed by GitHub
co-authored by Thomas Wang HaiShaw
parent 5e73c89b34
commit f7cb328eb7
4 changed files with 180 additions and 16 deletions
@@ -389,12 +389,70 @@ class Indexer(DSANPUIndexerMixin, BaseFusedOp):
return x if self.use_dsa_indexer_fusion else rotate_activation(x)
def _should_skip_logits_computation(self, forward_batch: ForwardBatch) -> bool:
if (
forward_batch.forward_mode.is_extend_without_speculative()
and forward_batch.seq_lens_cpu is not None
):
max_kv_len = forward_batch.seq_lens_cpu.max().item()
# When kv_len <= index_topk the top-k selects ALL valid positions, so the
# indexer's logits GEMM + paged_mqa_logits + top-k are wasted work: a plain
# topk_transform(dummy_logits) already yields the correct "select-all"
# (physical page-slot) indices. Skipping the logits path is safe here.
#
# Prefill/extend: original fast path, all platforms.
# Decode: new here, and ROCm-only for now (see the _is_hip gate below).
# Under a captured decode cuda graph the chosen branch is frozen at
# capture time and would replay incorrectly for kv_len > index_topk, so
# the decode skip is not decided per-step during capture; it is driven by
# which graph variant is being captured instead.
fb = forward_batch
# Prefill/extend: original per-step gate (host sync on seq_lens_cpu is fine).
if fb.forward_mode.is_extend_without_speculative():
if fb.seq_lens_cpu is None or fb.seq_lens_cpu.numel() == 0:
return False
return int(fb.seq_lens_cpu.max().item()) <= self.index_topk
# Decode/idle.
if fb.forward_mode.is_decode_or_idle():
# Decode k-only skip (both the captured dual-graph "dense" variant
# and the eager per-step skip below) is currently HIP-only. On CUDA
# this common code keeps the original behavior (decode never skips
# the indexer, i.e. always runs the full logits path) because the
# decode k-only path has not been validated on CUDA yet. Mirrors the
# is_hip() gate on dsa_dual_graph in decode_cuda_graph_runner, which
# already prevents the CUDA capture path from setting a "dense"
# variant.
if not _is_hip:
return False
if get_is_capture_mode():
# Under a captured decode cuda graph the taken branch is frozen at
# capture time, so we must NOT branch on a runtime seq_len (also a
# host sync would break capture). The chosen branch is instead
# driven by which graph variant is being captured.
#
# The decode runner captures a "dense" (k-only) and a "sparse"
# (full indexer) graph per bs bucket and dispatches on max_kv_len
# at replay. The capture-variant signal tells us which one to
# bake in.
from sglang.srt.model_executor.runner_utils.capture_mode import (
get_capture_dsa_variant,
)
variant = get_capture_dsa_variant()
if variant == "dense":
return True
if variant == "sparse":
return False
# No dual-variant capture signal: default to the correct-for-all
# full-indexer (sparse) path.
return False
# Eager decode: safe to check per-step (host sync OK); correct for both
# kv_len<=index_topk (k-only) and kv_len>index_topk (falls through).
if fb.seq_lens_cpu is not None and fb.seq_lens_cpu.numel() > 0:
max_kv_len = int(fb.seq_lens_cpu.max().item())
elif fb.seq_lens is not None and fb.seq_lens.numel() > 0:
max_kv_len = int(fb.seq_lens.max().item())
else:
return False
return max_kv_len <= self.index_topk
return False
def _get_q_k_bf16(
@@ -1175,7 +1233,10 @@ class Indexer(DSANPUIndexerMixin, BaseFusedOp):
# - topk_result: pre-allocated padded buffer to fill in place (a downstream
# captured graph reads it at a fixed address). None => return a fresh,
# naturally-sized tensor.
assert forward_batch.forward_mode.is_extend_without_speculative()
assert (
forward_batch.forward_mode.is_extend_without_speculative()
or forward_batch.forward_mode.is_decode_or_idle()
)
x_meta = x[0] if isinstance(x, tuple) else x
# Fast path: only compute and store k cache, skip all q and weights ops.
@@ -90,6 +90,7 @@ from sglang.srt.model_executor.runner_utils.buffers import (
DecodeInputBuffers,
)
from sglang.srt.model_executor.runner_utils.capture_mode import (
_set_capture_dsa_variant,
_set_capture_lora_variant,
model_capture_mode,
)
@@ -103,6 +104,7 @@ from sglang.srt.speculative.ragged_verify import resolve_ragged_verify_layout
from sglang.srt.utils import (
empty_context,
get_available_gpu_memory,
is_hip,
require_attn_tp_gather,
require_mlp_tp_gather,
)
@@ -246,6 +248,37 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
model_runner.server_args.enable_profile_cuda_graph
)
# --- DSA dense-decode dual-graph -------------------------------
# Capture a "dense" (k-only, skip-indexer) and a "sparse" (full indexer)
# decode graph per bs bucket, and dispatch on max_kv_len vs index_topk at
# replay. Auto-enabled for DSA models (index_topk present in the HF
# config) — correct for mixed lengths since any request with
# kv_len > index_topk falls back to the sparse graph. Adds ~52 graphs and
# ~2x capture time.
#
# Scoped to HIP (AMD): the k-only dense-decode fast path has only been
# validated on MI355X. This is common (non-hardware-gated) code, so on
# CUDA we deliberately keep the original behavior (no dual-graph) to
# avoid silently changing the CUDA decode path for DSA models (e.g.
# DeepSeek-V3.2). CUDA can opt in later once validated there.
self.dsa_dual_graph = False
self.dsa_index_topk: Optional[int] = None
from sglang.srt.configs.model_config import (
get_dsa_index_topk,
is_deepseek_dsa,
)
hf_config = model_runner.model_config.hf_config
if is_hip() and is_deepseek_dsa(hf_config):
self.dsa_index_topk = get_dsa_index_topk(hf_config)
self.dsa_dual_graph = True
logger.info(
"[dense-decode] DSA dual-graph enabled: capturing "
"dense (k-only) + sparse (full indexer) decode graphs; "
"dispatch on max_kv_len vs index_topk=%d.",
self.dsa_index_topk,
)
self.attn_tp_size = get_parallel().attn_tp_size
self.attn_tp_rank = get_parallel().attn_tp_rank
# True if a DSACPLayerCommunicator-style prefill-CP flavor is active
@@ -511,16 +544,39 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
def _cache_loc_dtype(self):
return torch.int64
def _make_graph_key(self, size, stream_idx=None, variant_label=None):
def _make_graph_key(
self, size, stream_idx=None, variant_label=None, dsa_variant=None
):
return ShapeKey(
size=size,
stream_idx=stream_idx,
variant_label=variant_label,
dsa_variant=dsa_variant,
)
def _capture_graph_size(self, *, bs: int, num_tokens: int) -> int:
return num_tokens if self.ragged_verify_mode else bs
def _resolve_dsa_variant(self, forward_batch: ForwardBatch) -> Optional[str]:
"""Host dispatch: pick which pre-captured DSA decode graph to replay
from the batch-max kv_len. If any request has kv_len > index_topk
the dense (k-only) graph would be wrong for it, so the whole batch uses
the sparse (full indexer) graph. Returns None when dual-graph is off."""
if not getattr(self, "dsa_dual_graph", False):
return None
seq_lens_cpu = getattr(forward_batch, "seq_lens_cpu", None)
if seq_lens_cpu is not None and seq_lens_cpu.numel() > 0:
# Host-side mirror (maintained incrementally for plain decode) — no
# d2h sync needed.
max_kv_len = int(seq_lens_cpu.max().item())
elif forward_batch.seq_lens is not None and forward_batch.seq_lens.numel() > 0:
# Fallback: a single scalar reduction d2h (cheap, per-step).
max_kv_len = int(forward_batch.seq_lens.max().item())
else:
# No length info: be safe and use the correct-for-all sparse graph.
return "sparse"
return "dense" if max_kv_len <= self.dsa_index_topk else "sparse"
def _resolve_lora_variant(self, forward_batch: ForwardBatch):
if not getattr(self, "record_nolora_graph", False):
return None
@@ -1025,6 +1081,17 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
if getattr(self, "record_nolora_graph", False)
else [(None, None)]
)
# DSA: capture a dense (k-only) and a sparse (full indexer) graph
# per bs bucket. Order: dense first so its (smaller) capture-time peak
# runs while the shared pool is fresh; sparse's peak subsumes it.
# getattr default: subclasses like EAGLEDraftCudaGraphRunner reuse this
# capture() but don't run DecodeCudaGraphRunner.__init__ (so they never
# set dsa_dual_graph) and override capture_one_shape with a signature that
# has no dsa_variant. Default to no dual-graph and, for the None variant,
# call capture_one_shape without the extra arg so those overrides work.
dsa_variants = (
["dense", "sparse"] if getattr(self, "dsa_dual_graph", False) else [None]
)
for bs in capture_range:
if get_parallel().tp_rank == 0:
avail_mem = get_available_gpu_memory(
@@ -1038,13 +1105,23 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
for variant_label, _variant_has_lora in lora_variants:
_set_capture_lora_variant(variant_label)
with torch_compile_decoration.patch_model(
self.model_runner.model,
bs in self.compile_bs,
num_tokens=bs * self.captured_req_width,
tp_group=self.model_runner.tp_group,
) as forward:
self.capture_one_shape(bs, forward, stream_idx, variant_label)
for dsa_variant in dsa_variants:
_set_capture_dsa_variant(dsa_variant)
with torch_compile_decoration.patch_model(
self.model_runner.model,
bs in self.compile_bs,
num_tokens=bs * self.captured_req_width,
tp_group=self.model_runner.tp_group,
) as forward:
if dsa_variant is None:
self.capture_one_shape(
bs, forward, stream_idx, variant_label
)
else:
self.capture_one_shape(
bs, forward, stream_idx, variant_label, dsa_variant
)
_set_capture_dsa_variant(None)
def capture_one_shape(
self,
@@ -1052,6 +1129,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
forward: Callable,
stream_idx: Optional[int] = None,
variant_label: Optional[str] = None,
dsa_variant: Optional[str] = None,
):
num_tokens = size * self.captured_req_width
bs = self._ragged_capture_slots(num_tokens) if self.ragged_verify_mode else size
@@ -1140,6 +1218,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
self._capture_graph_size(bs=bs, num_tokens=num_tokens),
stream_idx,
variant_label,
dsa_variant,
)
post_warmup_hook = getattr(
self.model_runner.attn_backend,
@@ -1208,9 +1287,10 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
forward_batch.input_embeds
)
variant_label = self._resolve_lora_variant(forward_batch)
dsa_variant = self._resolve_dsa_variant(forward_batch)
stream_idx = get_current_stream_idx() if self.enable_pdmux else None
self._replay_graph_key = self._make_graph_key(
graph_size_key, stream_idx, variant_label
graph_size_key, stream_idx, variant_label, dsa_variant
)
return
@@ -1300,9 +1380,10 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
self.model_runner.hisparse_coordinator.num_real_reqs.fill_(raw_bs)
variant_label = self._resolve_lora_variant(forward_batch)
dsa_variant = self._resolve_dsa_variant(forward_batch)
stream_idx = get_current_stream_idx() if self.enable_pdmux else None
self._replay_graph_key = self._make_graph_key(
graph_size_key, stream_idx, variant_label
graph_size_key, stream_idx, variant_label, dsa_variant
)
def _ragged_graph_num_tokens(self, total_verify_tokens: int) -> int:
@@ -30,8 +30,12 @@ class ShapeKey:
variant_label: optional execution variant (for example, "lora",
"nolora", or "chunked_prefix"), or None for runners that don't
record per-variant graphs.
dsa_variant: DSA decode dual-graph variant ("dense" / "sparse"), or None
when DSA dual-graph capture is not enabled. Composes with variant_label
so LoRA and DSA variants can be captured independently.
"""
size: int
stream_idx: Optional[int] = None
variant_label: Optional[str] = None
dsa_variant: Optional[str] = None
@@ -35,6 +35,13 @@ is_capture_mode = False
# None = not dual, "lora" = capturing lora variant, "nolora" = capturing nolora variant.
_capture_lora_variant: Optional[str] = None
# When capturing dual DSA decode graphs (dense/sparse), tracks which variant is
# being captured. Read by the DSA indexer's capture-time skip-logits branch to
# force k-only ("dense") vs full indexer ("sparse").
# None = not dual-capturing; the indexer then bakes in the full-indexer path,
# which is correct for any kv_len.
_capture_dsa_variant: Optional[str] = None
def get_is_capture_mode() -> bool:
return is_capture_mode or is_in_breakable_cuda_graph()
@@ -63,6 +70,17 @@ def _set_capture_lora_variant(variant: Optional[str]) -> None:
_capture_lora_variant = variant
def get_capture_dsa_variant() -> Optional[str]:
"""Return the DSA decode variant being captured ("dense"/"sparse"), or None
when dual-variant capture is not active."""
return _capture_dsa_variant
def _set_capture_dsa_variant(variant: Optional[str]) -> None:
global _capture_dsa_variant
_capture_dsa_variant = variant
@contextmanager
def model_capture_mode():
global is_capture_mode