[CUDA] cuDNN backward attention (#2762)
This commit is contained in:
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user