From 44c90c628264f46173867d4e43d8496e49e0e864 Mon Sep 17 00:00:00 2001 From: karverma-amd Date: Fri, 21 Aug 2026 01:08:22 -0500 Subject: [PATCH] [AMD] DSv4: fuse the qk-norm-rope pair on the MTP target-verify path (#34973) Co-authored-by: HAI Co-authored-by: Thomas Wang --- python/sglang/srt/environ.py | 4 ++ python/sglang/srt/models/deepseek_v4.py | 53 +++++++++++++++++++++---- 2 files changed, 49 insertions(+), 8 deletions(-) diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index 497bb6750..af5f70c7b 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -815,6 +815,10 @@ class Envs: # DSV4 Aiter flags SGLANG_OPT_USE_AITER_SILU_MUL = EnvBool(False) SGLANG_OPT_USE_FUSED_QK_NORM_ROPE = EnvBool(True) + # Unified KV wired the fused qk-norm-rope kernel to decode only, so MTP + # target-verify kept running the norm+RoPE as separate kernels. Set to 0 to + # go back to the unfused chain on the verify path. + SGLANG_OPT_FUSED_QK_NORM_ROPE_VERIFY = EnvBool(True) SGLANG_OPT_USE_AITER_INDEXER = EnvBool(False) # =================================================================== diff --git a/python/sglang/srt/models/deepseek_v4.py b/python/sglang/srt/models/deepseek_v4.py index bc2f39761..e088f6b05 100644 --- a/python/sglang/srt/models/deepseek_v4.py +++ b/python/sglang/srt/models/deepseek_v4.py @@ -1295,11 +1295,19 @@ class MQALayer(MqaAttentionBase): unified = is_unified_kv_triton() is_decode = forward_batch.forward_mode.is_decode_or_idle() - do_fused_store = (unified and is_decode) or ( + # The kernel is token-indexed (q, kv and positions are all length M), so + # a verify batch carrying several draft tokens per request is a shape it + # already handles. Only the cache store differs between decode and + # verify, and that half is left off below. + fuse_verify = ( + envs.SGLANG_OPT_FUSED_QK_NORM_ROPE_VERIFY.get() + and forward_batch.forward_mode.is_target_verify() + ) + do_fused_qk_norm_rope = (unified and (is_decode or fuse_verify)) or ( not unified and self.use_fused_qk_norm_rope ) - if do_fused_store: + if do_fused_qk_norm_rope: if _is_gfx95_supported: q_for_wqb, q_lora = _fused_rmsnorm_fp8_quant( q_lora, @@ -1318,7 +1326,27 @@ class MQALayer(MqaAttentionBase): ) token_to_kv_pool = get_token_to_kv_pool() - if unified: + if unified and fuse_verify: + # Target-verify runs through the unified_kv decode path. The + # backend writes the current chunk's KV into the ring *before* + # attention (save_kv_cache=True -> store_swa_into_unified ahead + # of runtime.decode), and per-token causal index streams -- built + # once per step in the backend metadata -- keep each draft query + # attending only to positions up to itself. Causal masking among + # the draft tokens comes from those index streams, not from store + # timing. So this path skips only the fused kernel's *own* store + # and returns kv, letting that existing causally-indexed backend + # store run unchanged; we fuse just the norm+RoPE. swa_loc is not + # computed -- it only addresses the kernel store this path drops. + # + # kv is a strided slice of qkv_a and the ring store requires a + # contiguous buffer, so materialise it before the kernel norms + # it in place. The unfused path pays the same copy inside + # _compute_kv_bf16. + kv = kv.contiguous() + swa_cache, swa_loc = None, None + swa_page_size, bf16_store = 1, True + elif unified: swa_cache = token_to_kv_pool.get_unified_kv(self.layer_id) # swa_loc is layer-independent; computed once per forward by the # backend and cached on the metadata (read here by every layer). @@ -1354,7 +1382,13 @@ class MQALayer(MqaAttentionBase): dtype=x.dtype, bf16_store=bf16_store, ) - kv = None + # On the verify path the kernel normed + RoPE'd kv in place and wrote + # nothing, so hand it back: the caller feeds it to attention as the + # current chunk (attn_k = kv) and save_kv_cache = kv is not None lets + # the backend do its normal causally-indexed store into the ring + # before the decode kernel runs -- exactly as the unfused path did. + if not (unified and fuse_verify): + kv = None if not unified and use_cp: # DSA CP: keep bf16 kv around for the cross-rank all-gather, then @@ -1563,10 +1597,13 @@ class MQALayer(MqaAttentionBase): x_quant=x_quant, ) - # The cache write is always fused / already done by _forward_prepare* -- - # tell the backend to skip its own store_cache. When `kv is None` - # (no DSA-CP), pass `q` as a sentinel for the `k is v` assert; the - # attention path doesn't read it once `save_kv_cache=False`. + # save_kv_cache = kv is not None selects who writes the ring. When kv is + # None the store was already fused into _forward_prepare* (decode) or + # done inline, so the backend skips its own store_cache; pass `q` as a + # sentinel for the `k is v` assert (attention won't read it once + # save_kv_cache=False). When kv is not None (target-verify, or DSA-CP), + # _forward_prepare* deliberately left the store off and the backend does + # its normal causally-indexed store from attn_k = kv. attn_k = kv if kv is not None else q from sglang.kernels.ops.attention.dsv4.unified_kv_kernels.env_gate import ( is_unified_kv_triton,