diff --git a/test/registered/attention/test_verify_shared_kv.py b/test/registered/attention/test_verify_shared_kv.py index 6a4549742..7dd469d4c 100644 --- a/test/registered/attention/test_verify_shared_kv.py +++ b/test/registered/attention/test_verify_shared_kv.py @@ -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(