[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) return x if self.use_dsa_indexer_fusion else rotate_activation(x)
def _should_skip_logits_computation(self, forward_batch: ForwardBatch) -> bool: def _should_skip_logits_computation(self, forward_batch: ForwardBatch) -> bool:
if ( # When kv_len <= index_topk the top-k selects ALL valid positions, so the
forward_batch.forward_mode.is_extend_without_speculative() # indexer's logits GEMM + paged_mqa_logits + top-k are wasted work: a plain
and forward_batch.seq_lens_cpu is not None # topk_transform(dummy_logits) already yields the correct "select-all"
): # (physical page-slot) indices. Skipping the logits path is safe here.
max_kv_len = forward_batch.seq_lens_cpu.max().item() #
# 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 max_kv_len <= self.index_topk
return False return False
def _get_q_k_bf16( 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 # - topk_result: pre-allocated padded buffer to fill in place (a downstream
# captured graph reads it at a fixed address). None => return a fresh, # captured graph reads it at a fixed address). None => return a fresh,
# naturally-sized tensor. # 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 x_meta = x[0] if isinstance(x, tuple) else x
# Fast path: only compute and store k cache, skip all q and weights ops. # 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, DecodeInputBuffers,
) )
from sglang.srt.model_executor.runner_utils.capture_mode import ( from sglang.srt.model_executor.runner_utils.capture_mode import (
_set_capture_dsa_variant,
_set_capture_lora_variant, _set_capture_lora_variant,
model_capture_mode, model_capture_mode,
) )
@@ -103,6 +104,7 @@ from sglang.srt.speculative.ragged_verify import resolve_ragged_verify_layout
from sglang.srt.utils import ( from sglang.srt.utils import (
empty_context, empty_context,
get_available_gpu_memory, get_available_gpu_memory,
is_hip,
require_attn_tp_gather, require_attn_tp_gather,
require_mlp_tp_gather, require_mlp_tp_gather,
) )
@@ -246,6 +248,37 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
model_runner.server_args.enable_profile_cuda_graph 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_size = get_parallel().attn_tp_size
self.attn_tp_rank = get_parallel().attn_tp_rank self.attn_tp_rank = get_parallel().attn_tp_rank
# True if a DSACPLayerCommunicator-style prefill-CP flavor is active # True if a DSACPLayerCommunicator-style prefill-CP flavor is active
@@ -511,16 +544,39 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
def _cache_loc_dtype(self): def _cache_loc_dtype(self):
return torch.int64 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( return ShapeKey(
size=size, size=size,
stream_idx=stream_idx, stream_idx=stream_idx,
variant_label=variant_label, variant_label=variant_label,
dsa_variant=dsa_variant,
) )
def _capture_graph_size(self, *, bs: int, num_tokens: int) -> int: def _capture_graph_size(self, *, bs: int, num_tokens: int) -> int:
return num_tokens if self.ragged_verify_mode else bs 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): def _resolve_lora_variant(self, forward_batch: ForwardBatch):
if not getattr(self, "record_nolora_graph", False): if not getattr(self, "record_nolora_graph", False):
return None return None
@@ -1025,6 +1081,17 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
if getattr(self, "record_nolora_graph", False) if getattr(self, "record_nolora_graph", False)
else [(None, None)] 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: for bs in capture_range:
if get_parallel().tp_rank == 0: if get_parallel().tp_rank == 0:
avail_mem = get_available_gpu_memory( avail_mem = get_available_gpu_memory(
@@ -1038,13 +1105,23 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
for variant_label, _variant_has_lora in lora_variants: for variant_label, _variant_has_lora in lora_variants:
_set_capture_lora_variant(variant_label) _set_capture_lora_variant(variant_label)
for dsa_variant in dsa_variants:
_set_capture_dsa_variant(dsa_variant)
with torch_compile_decoration.patch_model( with torch_compile_decoration.patch_model(
self.model_runner.model, self.model_runner.model,
bs in self.compile_bs, bs in self.compile_bs,
num_tokens=bs * self.captured_req_width, num_tokens=bs * self.captured_req_width,
tp_group=self.model_runner.tp_group, tp_group=self.model_runner.tp_group,
) as forward: ) as forward:
self.capture_one_shape(bs, forward, stream_idx, variant_label) 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( def capture_one_shape(
self, self,
@@ -1052,6 +1129,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
forward: Callable, forward: Callable,
stream_idx: Optional[int] = None, stream_idx: Optional[int] = None,
variant_label: Optional[str] = None, variant_label: Optional[str] = None,
dsa_variant: Optional[str] = None,
): ):
num_tokens = size * self.captured_req_width num_tokens = size * self.captured_req_width
bs = self._ragged_capture_slots(num_tokens) if self.ragged_verify_mode else size 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), self._capture_graph_size(bs=bs, num_tokens=num_tokens),
stream_idx, stream_idx,
variant_label, variant_label,
dsa_variant,
) )
post_warmup_hook = getattr( post_warmup_hook = getattr(
self.model_runner.attn_backend, self.model_runner.attn_backend,
@@ -1208,9 +1287,10 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
forward_batch.input_embeds forward_batch.input_embeds
) )
variant_label = self._resolve_lora_variant(forward_batch) 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 stream_idx = get_current_stream_idx() if self.enable_pdmux else None
self._replay_graph_key = self._make_graph_key( 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 return
@@ -1300,9 +1380,10 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
self.model_runner.hisparse_coordinator.num_real_reqs.fill_(raw_bs) self.model_runner.hisparse_coordinator.num_real_reqs.fill_(raw_bs)
variant_label = self._resolve_lora_variant(forward_batch) 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 stream_idx = get_current_stream_idx() if self.enable_pdmux else None
self._replay_graph_key = self._make_graph_key( 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: 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", variant_label: optional execution variant (for example, "lora",
"nolora", or "chunked_prefix"), or None for runners that don't "nolora", or "chunked_prefix"), or None for runners that don't
record per-variant graphs. 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 size: int
stream_idx: Optional[int] = None stream_idx: Optional[int] = None
variant_label: Optional[str] = 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. # None = not dual, "lora" = capturing lora variant, "nolora" = capturing nolora variant.
_capture_lora_variant: Optional[str] = None _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: def get_is_capture_mode() -> bool:
return is_capture_mode or is_in_breakable_cuda_graph() 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 _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 @contextmanager
def model_capture_mode(): def model_capture_mode():
global is_capture_mode global is_capture_mode