[AMD] Skip redundant CatArrayBatchedCopy in GLM-5 NSA TileLang decode (#24125)

This commit is contained in:
Jacob0226
2026-05-13 02:55:28 -07:00
committed by GitHub
parent a9359707c1
commit fc20f5b114
2 changed files with 60 additions and 20 deletions
@@ -1594,7 +1594,13 @@ class NativeSparseAttnBackend(
q_rope = q_rope.view( q_rope = q_rope.view(
-1, layer.tp_q_head_num, layer.head_dim - layer.v_head_dim -1, layer.tp_q_head_num, layer.head_dim - layer.v_head_dim
) )
# Caller passed split q_nope / q_rope; we'll need to concat below if
# the chosen impl wants q_all.
q_all = None
else: else:
# Caller passed already-concatenated q (q_all = q). Reuse it directly
# via a zero-copy view; the impl-specific blocks below will skip the
# otherwise redundant concat_mla_absorb_q_general call.
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]
q_rope = q_all[:, :, layer.v_head_dim :] q_rope = q_all[:, :, layer.v_head_dim :]
@@ -1643,7 +1649,11 @@ class NativeSparseAttnBackend(
page_table_1=page_table_1, page_table_1=page_table_1,
) )
elif self.nsa_decode_impl == "tilelang": elif self.nsa_decode_impl == "tilelang":
if q_rope is not None: # Cat-skip (HIP-only): when caller passes q_rope=None on HIP, q_all
# has already been set to a zero-copy view of q in the else branch
# above and we can reuse it directly. The `not _is_hip` clause keeps
# CUDA / MUSA paths byte-identical to pre-patch by always re-cat.
if q_all is None or not _is_hip:
q_all = concat_mla_absorb_q_general(q_nope, q_rope) 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,
@@ -1668,7 +1678,7 @@ class NativeSparseAttnBackend(
page_size=1, page_size=1,
) )
elif self.nsa_decode_impl == "aiter": elif self.nsa_decode_impl == "aiter":
if q_rope is not None: if q_all is None or not _is_hip:
q_all = torch.cat([q_nope, q_rope], dim=-1) q_all = torch.cat([q_nope, q_rope], dim=-1)
return self._forward_aiter( return self._forward_aiter(
q_all=q_all, q_all=q_all,
@@ -417,25 +417,55 @@ class DeepseekMLAForwardMixin:
self.rotary_emb.is_neox_style, self.rotary_emb.is_neox_style,
q_out_dtype=kv_cache_dtype, q_out_dtype=kv_cache_dtype,
) )
q_nope_fused = q_cat[..., : self.kv_lora_rank]
q_pe_fused = q_cat[..., self.kv_lora_rank :]
save_kv_cache = False save_kv_cache = False
if llama_4_scaling is not None: # On decode, pass q_cat directly to attn_mqa with q_rope=None so
q_nope_fused *= llama_4_scaling # nsa_backend.forward_decode reuses q_cat as a zero-copy view
attn_output = self.attn_mqa( # (`q.contiguous().view(...)` fast-path) instead of running the
q_nope_fused, # redundant `concat_mla_absorb_q_general(q_nope_fused, q_pe_fused)`
None, # that would otherwise rebuild a tensor byte-identical to q_cat.
None, # On ROCm tilelang decode, this eliminates the
forward_batch, # `CatArrayBatchedCopy<OpaqueType<1u>, ...>` kernel that used to
q_rope=q_pe_fused, # fire once per layer per decode step (~2.6 us / layer saved).
k_rope=k_pe_fused, # Prefill keeps the split form because nsa_backend.forward_extend
save_kv_cache=save_kv_cache, # asserts `q_rope is not None`.
**( if forward_batch.forward_mode.is_decode_or_idle():
dict(topk_indices=topk_indices) if llama_4_scaling is not None:
if topk_indices is not None # llama_4_scaling applies only to the q_nope portion;
else {} # mutate in place via the slice view of q_cat.
), q_cat[..., : self.kv_lora_rank] *= llama_4_scaling
) attn_output = self.attn_mqa(
q_cat,
None,
None,
forward_batch,
q_rope=None,
k_rope=k_pe_fused,
save_kv_cache=save_kv_cache,
**(
dict(topk_indices=topk_indices)
if topk_indices is not None
else {}
),
)
else:
q_nope_fused = q_cat[..., : self.kv_lora_rank]
q_pe_fused = q_cat[..., self.kv_lora_rank :]
if llama_4_scaling is not None:
q_nope_fused *= llama_4_scaling
attn_output = self.attn_mqa(
q_nope_fused,
None,
None,
forward_batch,
q_rope=q_pe_fused,
k_rope=k_pe_fused,
save_kv_cache=save_kv_cache,
**(
dict(topk_indices=topk_indices)
if topk_indices is not None
else {}
),
)
else: else:
extra_args = {} extra_args = {}
if self._fuse_rope_for_trtllm_mla(forward_batch): if self._fuse_rope_for_trtllm_mla(forward_batch):