[ROCm] Take the fused DSA metadata kernels and drop redundant work from the absorb path (#37124)

Co-authored-by: yanyuan.qin <yanyuan.qin@amd.com>
Co-authored-by: Zhang, Jiejing <jiejing.zhang@amd.com>
Co-authored-by: Thomas Wang <thomawan@amd.com>
Co-authored-by: HAI <hixiao@gmail.com>
This commit is contained in:
xiaobochen-amd
2026-09-06 17:39:37 -07:00
committed by GitHub
co-authored by yanyuan.qin Zhang, Jiejing Thomas Wang HAI
parent 30705c004c
commit 15aa2fb843
3 changed files with 53 additions and 20 deletions
@@ -122,7 +122,7 @@ class DeepseekSparseAttnBackendMTPPrecomputeMixin:
"""Precompute metadata for normal decode mode.""" """Precompute metadata for normal decode mode."""
max_len = self.decode_cuda_graph_metadata[bs].page_table_1.shape[1] max_len = self.decode_cuda_graph_metadata[bs].page_table_1.shape[1]
if _is_cuda and not _is_hip and self.dsa_index_kpool <= 1: if (_is_cuda or _is_hip) and self.dsa_index_kpool <= 1:
from sglang.kernels.ops.attention.dsa_metadata import ( from sglang.kernels.ops.attention.dsa_metadata import (
fused_dsa_decode_metadata, fused_dsa_decode_metadata,
) )
@@ -245,7 +245,7 @@ class DeepseekSparseAttnBackendMTPPrecomputeMixin:
max_seqlen_k = self.decode_cuda_graph_metadata[bs].page_table_1.shape[1] max_seqlen_k = self.decode_cuda_graph_metadata[bs].page_table_1.shape[1]
seqlens_expanded_size = bs * self.speculative_num_draft_tokens seqlens_expanded_size = bs * self.speculative_num_draft_tokens
if _is_cuda and not _is_hip and self.dsa_index_kpool <= 1: if (_is_cuda or _is_hip) and self.dsa_index_kpool <= 1:
from sglang.kernels.ops.attention.dsa_metadata import ( from sglang.kernels.ops.attention.dsa_metadata import (
fused_dsa_target_verify_metadata, fused_dsa_target_verify_metadata,
) )
@@ -1548,7 +1548,7 @@ class DeepseekSparseAttnBackend(
# Normal Decode # Normal Decode
max_len = self._graph_page_table_width(metadata) max_len = self._graph_page_table_width(metadata)
if is_cuda() and not _is_hip and self.dsa_index_kpool <= 1: if (is_cuda() or _is_hip) and self.dsa_index_kpool <= 1:
fused_dsa_decode_metadata( fused_dsa_decode_metadata(
seq_lens=seq_lens, seq_lens=seq_lens,
req_pool_indices=req_pool_indices, req_pool_indices=req_pool_indices,
@@ -1588,7 +1588,7 @@ class DeepseekSparseAttnBackend(
elif forward_mode.is_target_verify(): elif forward_mode.is_target_verify():
max_seqlen_k = self._graph_page_table_width(metadata) max_seqlen_k = self._graph_page_table_width(metadata)
if is_cuda() and not _is_hip and self.dsa_index_kpool <= 1: if (is_cuda() or _is_hip) and self.dsa_index_kpool <= 1:
paged_mqa_ctx_lens_2d = None paged_mqa_ctx_lens_2d = None
if ( if (
self.speculative_num_draft_tokens >= 2 self.speculative_num_draft_tokens >= 2
@@ -1684,7 +1684,7 @@ class DeepseekSparseAttnBackend(
device=self.device, device=self.device,
) )
if is_cuda() and not _is_hip and self.dsa_index_kpool <= 1: if (is_cuda() or _is_hip) and self.dsa_index_kpool <= 1:
fused_dsa_draft_extend_metadata( fused_dsa_draft_extend_metadata(
seq_lens=seq_lens, seq_lens=seq_lens,
extend_seq_lens=extend_seq_lens, extend_seq_lens=extend_seq_lens,
@@ -2077,7 +2077,6 @@ class DeepseekSparseAttnBackend(
) )
# Do absorbed multi-latent attention (MLA path) # Do absorbed multi-latent attention (MLA path)
assert q_rope is not None
kv_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id) kv_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id)
if q_rope is not None: if q_rope is not None:
@@ -2087,6 +2086,7 @@ class DeepseekSparseAttnBackend(
layer.tp_q_head_num, layer.tp_q_head_num,
layer.head_dim - layer.v_head_dim, layer.head_dim - layer.v_head_dim,
) )
q_all = None
else: else:
q_all = q.contiguous().view(-1, layer.tp_q_head_num, layer.head_dim) q_all = q.contiguous().view(-1, layer.tp_q_head_num, layer.head_dim)
q_nope = q_all[:, :, : layer.v_head_dim] q_nope = q_all[:, :, : layer.v_head_dim]
@@ -2167,7 +2167,11 @@ class DeepseekSparseAttnBackend(
sm_scale=layer.scaling, sm_scale=layer.scaling,
d_v=layer.v_head_dim, d_v=layer.v_head_dim,
) )
q_all = concat_mla_absorb_q_general(q_nope, q_rope) # Cat-skip, as in forward_decode: q_rope=None means the caller
# already handed us the concatenated form and q_all is a
# zero-copy view of it. `not _is_hip` keeps CUDA byte-identical.
if q_all is None or not _is_hip:
q_all = concat_mla_absorb_q_general(q_nope, q_rope)
return self._forward_tilelang( return self._forward_tilelang(
q_all=q_all, q_all=q_all,
kv_cache=kv_cache, kv_cache=kv_cache,
@@ -47,6 +47,9 @@ from sglang.srt.lora.deepseek_mla_correction import (
) )
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_executor.forward_context import get_token_to_kv_pool from sglang.srt.model_executor.forward_context import get_token_to_kv_pool
from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import (
is_in_tc_piecewise_cuda_graph,
)
from sglang.srt.models.deepseek_common.attention_forward_methods.forward_mla import ( from sglang.srt.models.deepseek_common.attention_forward_methods.forward_mla import (
_select_local_dcp_heads_for_autotune, _select_local_dcp_heads_for_autotune,
is_dcp_mla_decode_phase, is_dcp_mla_decode_phase,
@@ -141,6 +144,17 @@ if _use_aiter_gfx95:
from sglang.srt.layers.rocm_linear_utils import fused_qk_rope_cat_and_cache_mla from sglang.srt.layers.rocm_linear_utils import fused_qk_rope_cat_and_cache_mla
def _absorb_weight_bf16(w: torch.Tensor, w_scale) -> torch.Tensor:
"""Dequantize an absorbed MLA weight, skipping the pass when it is a no-op."""
if (
w.dtype == torch.bfloat16
and isinstance(w_scale, (int, float))
and w_scale == 1.0
):
return w
return w.to(torch.bfloat16) * w_scale
def rocm_absorb_q_bmm( def rocm_absorb_q_bmm(
attn: DeepseekV2AttentionMLA, attn: DeepseekV2AttentionMLA,
q_nope: torch.Tensor, q_nope: torch.Tensor,
@@ -186,7 +200,7 @@ def rocm_absorb_q_bmm(
else: else:
q_nope_out = torch.bmm( q_nope_out = torch.bmm(
q_nope.to(torch.bfloat16).transpose(0, 1), q_nope.to(torch.bfloat16).transpose(0, 1),
attn.w_kc.to(torch.bfloat16) * attn.w_scale, _absorb_weight_bf16(attn.w_kc, attn.w_scale),
) )
return q_nope_out return q_nope_out
@@ -240,10 +254,26 @@ def rocm_absorb_v_bmm(
transpose_bm_in=True, transpose_bm_in=True,
dtype=torch.bfloat16, dtype=torch.bfloat16,
) )
elif not is_in_tc_piecewise_cuda_graph():
# Same (batch, heads, dim) layout as the quantized paths above, so the
# post-GEMM flatten is a view. Skipped under piecewise: torch dynamo
# rejects out= with a non-contiguous output tensor.
_bmm_buf = torch.empty(
attn_output.shape[0],
attn.num_local_heads,
attn.w_vc.shape[2],
device=attn_output.device,
dtype=torch.bfloat16,
)
torch.bmm(
attn_output.to(torch.bfloat16).transpose(0, 1),
_absorb_weight_bf16(attn.w_vc, attn.w_scale),
out=_bmm_buf.transpose(0, 1),
)
else: else:
attn_bmm_output = torch.bmm( attn_bmm_output = torch.bmm(
attn_output.to(torch.bfloat16).transpose(0, 1), attn_output.to(torch.bfloat16).transpose(0, 1),
attn.w_vc.to(torch.bfloat16) * attn.w_scale, _absorb_weight_bf16(attn.w_vc, attn.w_scale),
) )
if _bmm_buf is not None: if _bmm_buf is not None:
@@ -677,17 +707,16 @@ class DeepseekMLARocmForwardMixin:
forward_batch.out_cache_loc, forward_batch.out_cache_loc,
) )
save_kv_cache = False save_kv_cache = False
# On decode, pass q_cat directly to attn_mqa with q_rope=None so # Pass q_cat straight to attn_mqa with q_rope=None so the backend
# dsa_backend.forward_decode reuses q_cat as a zero-copy view # reuses it as a zero-copy view instead of rebuilding a tensor
# (`q.contiguous().view(...)` fast-path) instead of running the # byte-identical to it -- one `CatArrayBatchedCopy` per layer per
# redundant `concat_mla_absorb_q_general(q_nope_fused, q_pe_fused)` # step. Target-verify is the same absorbed shape as decode, just
# that would otherwise rebuild a tensor byte-identical to q_cat. # more rows; real prefill keeps the split form because the Triton
# On ROCm tilelang decode, this eliminates the # sparse-MLA kernel reads q_nope/q_rope separately.
# `CatArrayBatchedCopy<OpaqueType<1u>, ...>` kernel that used to if (
# fire once per layer per decode step (~2.6 us / layer saved). forward_batch.forward_mode.is_decode_or_idle()
# Prefill keeps the split form because dsa_backend.forward_extend or forward_batch.forward_mode.is_target_verify()
# asserts `q_rope is not None`. ):
if forward_batch.forward_mode.is_decode_or_idle():
if llama_4_scaling is not None: if llama_4_scaling is not None:
# llama_4_scaling applies only to the q_nope portion; # llama_4_scaling applies only to the q_nope portion;
# mutate in place via the slice view of q_cat. # mutate in place via the slice view of q_cat.