[AMD] DSv4: fuse the qk-norm-rope pair on the MTP target-verify path (#34973)

Co-authored-by: HAI <hixiao@gmail.com>
Co-authored-by: Thomas Wang <thomawan@amd.com>
This commit is contained in:
karverma-amd
2026-08-20 23:08:22 -07:00
committed by GitHub
co-authored by HAI Thomas Wang
parent a688682f4b
commit 44c90c6282
2 changed files with 49 additions and 8 deletions
+4
View File
@@ -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)
# ===================================================================
+45 -8
View File
@@ -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,