[CUDA] cuDNN backward attention (#2762)

This commit is contained in:
Cheng
2025-11-19 08:13:50 +09:00
committed by GitHub
parent b167f0df1c
commit 6f35017d1b
8 changed files with 472 additions and 82 deletions
+33
View File
@@ -738,6 +738,39 @@ class TestSDPA(mlx_tests.MLXTestCase):
)
self.assertTrue(mx.allclose(out, expected, atol=1e-5))
def test_sdpa_grad(self):
B, N_kv, T, D = (2, 8, 128, 64)
scale = D**-0.5
f1 = lambda q, k, v: mlx_ref_attn(q, k, v, scale=scale)
f2 = lambda q, k, v: mx.fast.scaled_dot_product_attention(q, k, v, scale=scale)
f3 = lambda q, k, v: mlx_ref_attn(q, k, v, scale=scale).sum()
f4 = lambda q, k, v: (
mx.fast.scaled_dot_product_attention(q, k, v, scale=scale)
).sum()
# High tolerance due to cuDNN SDPA kernel requiring tf32.
tolerance = {"rtol": 1e-2, "atol": 1e-2}
for N_q in (8, 32):
q = mx.random.normal(shape=(B, N_q, T, D), dtype=mx.float16)
k = mx.random.normal(shape=(B, N_kv, T, D), dtype=mx.float16)
v = mx.random.normal(shape=(B, N_kv, T, D), dtype=mx.float16)
cotan = mx.ones_like(q)
o1, vjp1 = mx.vjp(f1, [q, k, v], [cotan])
o2, vjp2 = mx.vjp(f2, [q, k, v], [cotan])
self.assertTrue(mx.allclose(o1[0], o2[0], **tolerance))
for i in range(3):
self.assertTrue(mx.allclose(vjp1[i], vjp2[i], **tolerance))
g1 = mx.grad(f3)(q, k, v)
g2 = mx.grad(f4)(q, k, v)
self.assertTrue(mx.allclose(g1, g2, **tolerance))
if __name__ == "__main__":
mlx_tests.MLXTestRunner(failfast=True)