[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:
co-authored by
HAI
Thomas Wang
parent
a688682f4b
commit
44c90c6282
@@ -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)
|
||||
|
||||
# ===================================================================
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user