[AMD] [GLM-5.3-Flash Day 0] Enable zero-RoPE MHA prefill on ROCm (#39338)
Co-authored-by: Thomas Wang <thomawan@amd.com> Co-authored-by: Kevin Mi <45493463+kevin-mii@users.noreply.github.com>
This commit is contained in:
co-authored by
Thomas Wang
Kevin Mi
parent
582389cec5
commit
042b6a488f
+7
-1
@@ -279,8 +279,14 @@ class DeepseekMHARocmForwardMixin:
|
||||
def _concat_and_cast_mha_k_rocm(
|
||||
self: DeepseekV2AttentionMLA,
|
||||
k_nope: torch.Tensor,
|
||||
k_pe: torch.Tensor,
|
||||
k_pe: torch.Tensor | None,
|
||||
):
|
||||
if self.qk_rope_head_dim == 0:
|
||||
assert k_pe is None or k_pe.shape[-1] == 0
|
||||
# No RoPE tail to append, so k is k_nope as-is. The concat branch
|
||||
# below keeps k_nope's dtype, so no cast is needed here either.
|
||||
return k_nope.contiguous()
|
||||
|
||||
k_shape = (k_nope.shape[0], self.num_local_heads, self.qk_head_dim)
|
||||
k = k_nope.new_empty(*k_shape)
|
||||
if self.current_attention_backend == "aiter":
|
||||
|
||||
@@ -0,0 +1,83 @@
|
||||
import sys
|
||||
from types import SimpleNamespace
|
||||
from unittest import mock
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from sglang.srt.models.deepseek_common.attention_forward_methods import (
|
||||
forward_mha,
|
||||
forward_mha_rocm,
|
||||
)
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
|
||||
register_cpu_ci(
|
||||
est_time=10,
|
||||
suite="base-a-test-cpu",
|
||||
nightly=False,
|
||||
disabled=None,
|
||||
)
|
||||
|
||||
|
||||
def _run_nope_concat(backend: str, kv_cache_dtype: str, pool_dtype: torch.dtype):
|
||||
fake_self = SimpleNamespace(
|
||||
qk_rope_head_dim=0,
|
||||
current_attention_backend=backend,
|
||||
kv_cache_dtype=kv_cache_dtype,
|
||||
)
|
||||
fake_pool = SimpleNamespace(dtype=pool_dtype)
|
||||
with (
|
||||
mock.patch.object(forward_mha, "_is_cuda", True),
|
||||
mock.patch.object(forward_mha, "get_token_to_kv_pool", return_value=fake_pool),
|
||||
):
|
||||
return forward_mha.DeepseekMHAForwardMixin._concat_and_cast_mha_k(
|
||||
fake_self,
|
||||
torch.randn(4, 2, 128, dtype=torch.bfloat16),
|
||||
None,
|
||||
None,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"backend,kv_cache_dtype,pool_dtype,expected_dtype",
|
||||
[
|
||||
("fa3", "fp8_e4m3", torch.float8_e4m3fn, torch.float8_e4m3fn),
|
||||
("fa3", "auto", torch.float8_e4m3fn, torch.bfloat16),
|
||||
("trtllm_gen", "fp8_e4m3", torch.float8_e4m3fn, torch.bfloat16),
|
||||
],
|
||||
)
|
||||
def test_nope_mha_k_cast(backend, kv_cache_dtype, pool_dtype, expected_dtype):
|
||||
out = _run_nope_concat(backend, kv_cache_dtype, pool_dtype)
|
||||
assert out.dtype == expected_dtype
|
||||
|
||||
|
||||
def _run_nope_concat_rocm(backend: str, k_pe: torch.Tensor | None):
|
||||
# qk_head_dim / qk_nope_head_dim are the roped-model values so that the
|
||||
# concat branch would be entered (and fail on the zero-width tail) if the
|
||||
# qk_rope_head_dim == 0 guard were missing.
|
||||
fake_self = SimpleNamespace(
|
||||
qk_rope_head_dim=0,
|
||||
qk_nope_head_dim=128,
|
||||
qk_head_dim=128,
|
||||
num_local_heads=2,
|
||||
current_attention_backend=backend,
|
||||
)
|
||||
return forward_mha_rocm.DeepseekMHARocmForwardMixin._concat_and_cast_mha_k_rocm(
|
||||
fake_self,
|
||||
torch.randn(4, 2, 128, dtype=torch.bfloat16),
|
||||
k_pe,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("backend", ["aiter", "triton"])
|
||||
@pytest.mark.parametrize("zero_width_k_pe", [False, True])
|
||||
def test_nope_mha_k_cast_rocm(backend, zero_width_k_pe):
|
||||
k_pe = torch.randn(4, 1, 0, dtype=torch.bfloat16) if zero_width_k_pe else None
|
||||
out = _run_nope_concat_rocm(backend, k_pe)
|
||||
assert out.shape == (4, 2, 128)
|
||||
assert out.dtype == torch.bfloat16
|
||||
assert out.is_contiguous()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(pytest.main([__file__]))
|
||||
Reference in New Issue
Block a user