From 15aa2fb8433d31c1f59d2e6fa74ae3f22b2b377c Mon Sep 17 00:00:00 2001 From: xiaobochen-amd Date: Mon, 7 Sep 2026 08:39:37 +0800 Subject: [PATCH] [ROCm] Take the fused DSA metadata kernels and drop redundant work from the absorb path (#37124) Co-authored-by: yanyuan.qin Co-authored-by: Zhang, Jiejing Co-authored-by: Thomas Wang Co-authored-by: HAI --- .../dsa/dsa_backend_mtp_precompute.py | 4 +- .../srt/layers/attention/dsa_backend.py | 14 +++-- .../forward_mla_rocm.py | 55 ++++++++++++++----- 3 files changed, 53 insertions(+), 20 deletions(-) diff --git a/python/sglang/srt/layers/attention/dsa/dsa_backend_mtp_precompute.py b/python/sglang/srt/layers/attention/dsa/dsa_backend_mtp_precompute.py index 6950e4a34..b6f1291a2 100644 --- a/python/sglang/srt/layers/attention/dsa/dsa_backend_mtp_precompute.py +++ b/python/sglang/srt/layers/attention/dsa/dsa_backend_mtp_precompute.py @@ -122,7 +122,7 @@ class DeepseekSparseAttnBackendMTPPrecomputeMixin: """Precompute metadata for normal decode mode.""" 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 ( fused_dsa_decode_metadata, ) @@ -245,7 +245,7 @@ class DeepseekSparseAttnBackendMTPPrecomputeMixin: max_seqlen_k = self.decode_cuda_graph_metadata[bs].page_table_1.shape[1] 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 ( fused_dsa_target_verify_metadata, ) diff --git a/python/sglang/srt/layers/attention/dsa_backend.py b/python/sglang/srt/layers/attention/dsa_backend.py index 80ea57740..a2c35aeae 100644 --- a/python/sglang/srt/layers/attention/dsa_backend.py +++ b/python/sglang/srt/layers/attention/dsa_backend.py @@ -1548,7 +1548,7 @@ class DeepseekSparseAttnBackend( # Normal Decode 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( seq_lens=seq_lens, req_pool_indices=req_pool_indices, @@ -1588,7 +1588,7 @@ class DeepseekSparseAttnBackend( elif forward_mode.is_target_verify(): 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 if ( self.speculative_num_draft_tokens >= 2 @@ -1684,7 +1684,7 @@ class DeepseekSparseAttnBackend( 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( seq_lens=seq_lens, extend_seq_lens=extend_seq_lens, @@ -2077,7 +2077,6 @@ class DeepseekSparseAttnBackend( ) # 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) if q_rope is not None: @@ -2087,6 +2086,7 @@ class DeepseekSparseAttnBackend( layer.tp_q_head_num, layer.head_dim - layer.v_head_dim, ) + q_all = None else: q_all = q.contiguous().view(-1, layer.tp_q_head_num, layer.head_dim) q_nope = q_all[:, :, : layer.v_head_dim] @@ -2167,7 +2167,11 @@ class DeepseekSparseAttnBackend( sm_scale=layer.scaling, 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( q_all=q_all, kv_cache=kv_cache, diff --git a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla_rocm.py b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla_rocm.py index 667d54856..27b13f232 100644 --- a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla_rocm.py +++ b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla_rocm.py @@ -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_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 ( _select_local_dcp_heads_for_autotune, 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 +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( attn: DeepseekV2AttentionMLA, q_nope: torch.Tensor, @@ -186,7 +200,7 @@ def rocm_absorb_q_bmm( else: q_nope_out = torch.bmm( 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 @@ -240,10 +254,26 @@ def rocm_absorb_v_bmm( transpose_bm_in=True, 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: attn_bmm_output = torch.bmm( 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: @@ -677,17 +707,16 @@ class DeepseekMLARocmForwardMixin: forward_batch.out_cache_loc, ) save_kv_cache = False - # On decode, pass q_cat directly to attn_mqa with q_rope=None so - # dsa_backend.forward_decode reuses q_cat as a zero-copy view - # (`q.contiguous().view(...)` fast-path) instead of running the - # redundant `concat_mla_absorb_q_general(q_nope_fused, q_pe_fused)` - # that would otherwise rebuild a tensor byte-identical to q_cat. - # On ROCm tilelang decode, this eliminates the - # `CatArrayBatchedCopy, ...>` kernel that used to - # fire once per layer per decode step (~2.6 us / layer saved). - # Prefill keeps the split form because dsa_backend.forward_extend - # asserts `q_rope is not None`. - if forward_batch.forward_mode.is_decode_or_idle(): + # Pass q_cat straight to attn_mqa with q_rope=None so the backend + # reuses it as a zero-copy view instead of rebuilding a tensor + # byte-identical to it -- one `CatArrayBatchedCopy` per layer per + # step. Target-verify is the same absorbed shape as decode, just + # more rows; real prefill keeps the split form because the Triton + # sparse-MLA kernel reads q_nope/q_rope separately. + if ( + forward_batch.forward_mode.is_decode_or_idle() + or forward_batch.forward_mode.is_target_verify() + ): if llama_4_scaling is not None: # llama_4_scaling applies only to the q_nope portion; # mutate in place via the slice view of q_cat.