[AMD][CI] Fix shared-KV verify tests for multi-head GQA (#36296)

This commit is contained in:
Bingxu Chen
2026-08-25 19:46:27 -07:00
committed by GitHub
parent 4ae30dc736
commit 07a9de25b4
@@ -29,6 +29,7 @@ def _build_inputs(
h_q,
head_dim,
v_head_dim,
h_kv=1,
cache_dtype=torch.bfloat16,
):
device = "cuda"
@@ -43,10 +44,10 @@ def _build_inputs(
return torch.randn(*shape, dtype=dtype, device=device, generator=generator)
q = randn(num_extend_tokens, h_q, head_dim)
k = randn(num_extend_tokens, 1, head_dim)
v = randn(num_extend_tokens, 1, v_head_dim)
k_buffer = randn(total_prefix, 1, head_dim).to(cache_dtype)
v_buffer = randn(total_prefix, 1, v_head_dim).to(cache_dtype)
k = randn(num_extend_tokens, h_kv, head_dim)
v = randn(num_extend_tokens, h_kv, v_head_dim)
k_buffer = randn(total_prefix, h_kv, head_dim).to(cache_dtype)
v_buffer = randn(total_prefix, h_kv, v_head_dim).to(cache_dtype)
qo_indptr = torch.arange(
0, num_extend_tokens + 1, l_ext, dtype=torch.int32, device=device
)
@@ -63,6 +64,8 @@ class TestVerifySharedKV(CustomTestCase):
head_dim,
v_head_dim,
h_q=4,
h_kv=1,
is_causal=True,
cache_dtype=torch.bfloat16,
k_scale=1.0,
v_scale=1.0,
@@ -74,6 +77,7 @@ class TestVerifySharedKV(CustomTestCase):
prefix_lens=[512, 2048],
l_ext=l_ext,
h_q=h_q,
h_kv=h_kv,
head_dim=head_dim,
v_head_dim=v_head_dim,
cache_dtype=cache_dtype,
@@ -95,7 +99,7 @@ class TestVerifySharedKV(CustomTestCase):
kv_indptr,
kv_indices,
None,
True,
is_causal,
None,
l_ext,
k_scale,
@@ -113,7 +117,7 @@ class TestVerifySharedKV(CustomTestCase):
kv_indptr,
kv_indices,
None,
True,
is_causal,
None,
l_ext,
k_scale,
@@ -158,18 +162,25 @@ class TestVerifySharedKV(CustomTestCase):
def test_kimi_k3_absorbed_mla_shape(self):
self._run_parity(head_dim=576, v_head_dim=512)
def test_rejects_multiple_local_kv_heads(self):
inputs = list(
_build_inputs(
prefix_lens=[512],
l_ext=4,
h_q=4,
head_dim=256,
v_head_dim=256,
)
def test_kimi_k3_dspark_gqa_shape(self):
self._run_parity(
head_dim=64,
v_head_dim=64,
h_q=8,
h_kv=2,
is_causal=False,
l_ext=7,
)
def test_rejects_non_power_of_two_kv_group(self):
inputs = _build_inputs(
prefix_lens=[512],
l_ext=4,
h_q=6,
h_kv=2,
head_dim=256,
v_head_dim=256,
)
for index in (1, 2, 3, 4):
inputs[index] = inputs[index].expand(-1, 2, -1).contiguous()
q, k, v, k_buffer, v_buffer, qo_indptr, kv_indptr, kv_indices = inputs
output = torch.empty_like(q)
self.assertFalse(