From 7b03ec21beaff66424aa177008098f608137bfe1 Mon Sep 17 00:00:00 2001 From: Yanzhao Wang Date: Tue, 8 Sep 2026 00:38:22 -0700 Subject: [PATCH] Add D256 to GQA-8 two-pass vector attention --- .../scaled_dot_product_attention.metal | 1 + .../metal/scaled_dot_product_attention.cpp | 3 ++- python/tests/test_fast_sdpa.py | 25 +++++++++++++++++++ 3 files changed, 28 insertions(+), 1 deletion(-) diff --git a/mlx/backend/metal/kernels/scaled_dot_product_attention.metal b/mlx/backend/metal/kernels/scaled_dot_product_attention.metal index db86a0edea..bb3325e230 100644 --- a/mlx/backend/metal/kernels/scaled_dot_product_attention.metal +++ b/mlx/backend/metal/kernels/scaled_dot_product_attention.metal @@ -47,6 +47,7 @@ using namespace metal; instantiate_sdpa_vector(type, 256, 256) \ instantiate_sdpa_vector_gqa(type, 64, 64, 8, 8) \ instantiate_sdpa_vector_gqa(type, 128, 128, 8, 4) \ + instantiate_sdpa_vector_gqa(type, 256, 256, 8, 2) \ instantiate_sdpa_vector_gqa(type, 128, 128, 12, 4) \ instantiate_sdpa_vector_gqa(type, 128, 128, 16, 2) \ instantiate_sdpa_vector_aggregation(type, 64) \ diff --git a/mlx/backend/metal/scaled_dot_product_attention.cpp b/mlx/backend/metal/scaled_dot_product_attention.cpp index 6d9ade1ccc..48494eca02 100644 --- a/mlx/backend/metal/scaled_dot_product_attention.cpp +++ b/mlx/backend/metal/scaled_dot_product_attention.cpp @@ -538,7 +538,8 @@ void sdpa_vector_2pass( kname += "sdpa_vector_2pass_1"; int gqa_factor = q.shape(1) / k.shape(1); bool gqa_dims = - (gqa_factor == 8 && (q.shape(-1) == 64 || q.shape(-1) == 128)) || + (gqa_factor == 8 && + (q.shape(-1) == 64 || q.shape(-1) == 128 || q.shape(-1) == 256)) || ((gqa_factor == 12 || gqa_factor == 16) && q.shape(-1) == 128); if (!mask && !sinks && q.shape(2) == 1 && q.shape(1) == gqa_factor * k.shape(1) && q.shape(-1) == v.shape(-1) && diff --git a/python/tests/test_fast_sdpa.py b/python/tests/test_fast_sdpa.py index 6b185ecf1f..45eff6fda8 100644 --- a/python/tests/test_fast_sdpa.py +++ b/python/tests/test_fast_sdpa.py @@ -418,6 +418,31 @@ def test_sdpa_vector_gqa_long(self): out = mx.fast.scaled_dot_product_attention(q, k, v, scale=scale) self.assertTrue(mx.allclose(ref, out, atol=1e-4, rtol=1e-4)) + @unittest.skipUnless(mx.metal.is_available(), "Metal is not available") + def test_sdpa_vector_gqa_d256(self): + if mx.default_device() != mx.gpu: + self.skipTest("requires GPU") + mx.random.seed(0) + for dtype, atol in [ + (mx.float32, 1e-4), + (mx.float16, 1e-3), + (mx.bfloat16, 5e-3), + ]: + for B, length in [(1, 8192), (1, 8201), (2, 8192), (2, 32768)]: + with self.subTest(dtype=dtype, B=B, length=length): + q = mx.random.normal((B, 16, 1, 256)).astype(dtype) + k = mx.random.normal((B, 2, length + 32, 256)).astype(dtype) + v = mx.random.normal((B, 2, length + 32, 256)).astype(dtype) + k, v = k[:, :, :length], v[:, :, :length] + ref = mlx_primitives_sdpa( + q.astype(mx.float32), + mx.repeat(k.astype(mx.float32), 8, axis=1), + mx.repeat(v.astype(mx.float32), 8, axis=1), + 256**-0.5, + ) + out = mx.fast.scaled_dot_product_attention(q, k, v, scale=256**-0.5) + self.assertTrue(mx.allclose(ref, out, atol=atol, rtol=1e-3)) + def test_sdpa_fully_masked(self): Lkv = 8 mask = mx.array(False)