[AMD] Add _skip_rope_for_aiter_fused_mla method and check to avoid double rotating with gfx950 and Aiter backend (#24148)
This commit is contained in:
@@ -352,10 +352,12 @@ class DeepseekMLAForwardMixin:
|
|||||||
q_nope_out = q_nope_out.transpose(0, 1)
|
q_nope_out = q_nope_out.transpose(0, 1)
|
||||||
|
|
||||||
skip_rope_for_nsa_tilelang_fused = self._skip_rope_for_nsa_tilelang_fused()
|
skip_rope_for_nsa_tilelang_fused = self._skip_rope_for_nsa_tilelang_fused()
|
||||||
|
skip_rope_for_aiter_fused_mla = self._skip_rope_for_aiter_fused_mla()
|
||||||
if (
|
if (
|
||||||
self.rotary_emb is not None
|
self.rotary_emb is not None
|
||||||
and (not self._fuse_rope_for_trtllm_mla(forward_batch))
|
and (not self._fuse_rope_for_trtllm_mla(forward_batch))
|
||||||
and (not skip_rope_for_nsa_tilelang_fused)
|
and (not skip_rope_for_nsa_tilelang_fused)
|
||||||
|
and (not skip_rope_for_aiter_fused_mla)
|
||||||
and (not _use_aiter or not _is_gfx95_supported or self.use_nsa)
|
and (not _use_aiter or not _is_gfx95_supported or self.use_nsa)
|
||||||
):
|
):
|
||||||
q_pe, k_pe = self.rotary_emb(positions, q_pe, k_pe)
|
q_pe, k_pe = self.rotary_emb(positions, q_pe, k_pe)
|
||||||
@@ -664,3 +666,15 @@ class DeepseekMLAForwardMixin:
|
|||||||
or server_args.nsa_prefill_backend == "tilelang"
|
or server_args.nsa_prefill_backend == "tilelang"
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def _skip_rope_for_aiter_fused_mla(self: DeepseekV2AttentionMLA) -> bool:
|
||||||
|
"""
|
||||||
|
Skip rope in prepare and let the fused kernel in forward_absorb_core handle it,
|
||||||
|
when running aiter-backend MLA on gfx95 (i.e., the `else` branch in forward_absorb_core
|
||||||
|
that calls fused_qk_rope_cat_and_cache_mla).
|
||||||
|
"""
|
||||||
|
return (
|
||||||
|
_use_aiter_gfx95
|
||||||
|
and self.current_attention_backend
|
||||||
|
not in FORWARD_ABSORB_CORE_ATTENTION_BACKENDS
|
||||||
|
)
|
||||||
|
|||||||
Reference in New Issue
Block a user