From 72e94c81e1685c90679ef03532c4b8897010abf9 Mon Sep 17 00:00:00 2001 From: Cheng Date: Wed, 11 Feb 2026 16:46:39 +0900 Subject: [PATCH] [CUDA] Attention sinks in cuDNN SDPA (#3118) --- mlx/backend/cuda/cudnn_utils.h | 14 +++ .../cuda/scaled_dot_product_attention.cpp | 91 +++++++++++++++++-- python/tests/test_fast_sdpa.py | 29 +++--- 3 files changed, 113 insertions(+), 21 deletions(-) diff --git a/mlx/backend/cuda/cudnn_utils.h b/mlx/backend/cuda/cudnn_utils.h index ef930213..5e8235f1 100644 --- a/mlx/backend/cuda/cudnn_utils.h +++ b/mlx/backend/cuda/cudnn_utils.h @@ -123,6 +123,20 @@ class DnnGraph : public fe::graph::Graph { return attrs; } + // Create a 4D cuDNN tensor from 1D array, with |axis| being contiguous dim. + auto tensor_4d(const char* name, int64_t uid, const array& x, int axis) { + assert(x.ndim() == 1); + auto attrs = Graph::tensor(fe::graph::Tensor_attributes().set_name(name)); + std::vector shape(4, 1); + std::vector strides(4, 1); + shape.at(axis) = x.size(); + if (axis > 0) { + strides.at(axis - 1) = x.size(); + } + set_tensor_attrs(attrs, uid, x, shape, strides); + return attrs; + } + // Create a cuDNN tensor for scalar. auto scalar(const char* name, int64_t uid, Dtype dtype) { return Graph::tensor( diff --git a/mlx/backend/cuda/scaled_dot_product_attention.cpp b/mlx/backend/cuda/scaled_dot_product_attention.cpp index 1cec0aa9..7f0e1f70 100644 --- a/mlx/backend/cuda/scaled_dot_product_attention.cpp +++ b/mlx/backend/cuda/scaled_dot_product_attention.cpp @@ -25,6 +25,18 @@ array prepare_sdpa_input(const array& x, Stream s) { return x; } +array prepare_sdpa_sinks(const array& sinks, Stream s) { + // cuDNN requires sinks to be float32. + if (sinks.dtype() == float32) { + return sinks; + } + array sinks_f32(sinks.shape(), float32, nullptr, {}); + copy_gpu(sinks, sinks_f32, CopyType::Vector, s); + auto& encoder = cu::get_command_encoder(s); + encoder.add_temporary(sinks_f32); + return sinks_f32; +} + void malloc_with_same_layout( cu::CommandEncoder& encoder, array& o, @@ -123,6 +135,7 @@ struct SDPACacheKey { bool do_causal; std::array mask_shape; std::array mask_strides; + bool has_sinks; bool output_logsumexp; }; @@ -133,6 +146,7 @@ inline BytesKey build_sdpa_cache_key( const array& v, bool do_causal, const std::optional& mask_arr, + const std::optional& sinks, bool decoding = false, bool output_logsumexp = false) { BytesKey cache_key; @@ -148,6 +162,7 @@ inline BytesKey build_sdpa_cache_key( do_causal, {}, {}, + sinks.has_value(), output_logsumexp, }; if (mask_arr) { @@ -182,6 +197,7 @@ enum UIDS { V, SCALE, BIAS, + SINKS, SEQ_LEN_Q, SEQ_LEN_KV, O, @@ -200,6 +216,7 @@ DnnGraph build_sdpa_graph( const array& v, bool do_causal, const std::optional& mask_arr, + const std::optional& sinks, const std::optional& seq_len_q, const std::optional& seq_len_kv, bool output_logsumexp, @@ -221,6 +238,9 @@ DnnGraph build_sdpa_graph( if (mask_arr) { options.set_bias(graph.tensor("BIAS", BIAS, *mask_arr)); } + if (sinks) { + options.set_sink_token(graph.tensor_4d("SINKS", SINKS, *sinks, 1)); + } if (seq_len_q && seq_len_kv) { options.set_padding_mask(true); options.set_seq_len_q(graph.tensor("SEQ_LEN_Q", SEQ_LEN_Q, *seq_len_q)); @@ -247,6 +267,7 @@ DnnGraph build_sdpa_backward_graph( const array& v, bool do_causal, const std::optional& mask_arr, + const std::optional& sinks, const array& o, const array& d_o, const array& stats, @@ -271,6 +292,9 @@ DnnGraph build_sdpa_backward_graph( if (mask_arr) { options.set_bias(graph.tensor("BIAS", BIAS, *mask_arr)); } + if (sinks) { + options.set_sink_token(graph.tensor_4d("SINKS", SINKS, *sinks, 1)); + } auto [d_q_, d_k_, d_v_] = graph.sdpa_backward(q_, k_, v_, o_, d_o_, stats_, options); @@ -333,6 +357,7 @@ void sdpa_cudnn( std::optional& stats, bool do_causal, const std::optional& mask_arr, + const std::optional& sinks, bool output_logsumexp, Stream s) { auto& encoder = cu::get_command_encoder(s); @@ -365,6 +390,9 @@ void sdpa_cudnn( if (mask_arr) { encoder.set_input_array(*mask_arr); } + if (sinks) { + encoder.set_input_array(*sinks); + } if (seq_len_q && seq_len_kv) { encoder.set_input_array(*seq_len_q); encoder.set_input_array(*seq_len_kv); @@ -376,7 +404,7 @@ void sdpa_cudnn( // Search cache. auto cache_key = build_sdpa_cache_key( - encoder, q, k, v, do_causal, mask_arr, decoding, output_logsumexp); + encoder, q, k, v, do_causal, mask_arr, sinks, decoding, output_logsumexp); auto it = sdpa_cache().find(cache_key); if (it == sdpa_cache().end()) { auto graph = build_sdpa_graph( @@ -386,6 +414,7 @@ void sdpa_cudnn( v, do_causal, mask_arr, + sinks, seq_len_q, seq_len_kv, output_logsumexp, @@ -404,6 +433,9 @@ void sdpa_cudnn( if (mask_arr) { variant_pack[BIAS] = gpu_ptr(*mask_arr); } + if (sinks) { + variant_pack[SINKS] = gpu_ptr(*sinks); + } if (seq_len_q && seq_len_kv) { variant_pack[SEQ_LEN_Q] = gpu_ptr(*seq_len_q); variant_pack[SEQ_LEN_KV] = gpu_ptr(*seq_len_kv); @@ -424,6 +456,7 @@ void sdpa_backward_cudnn( const array& stats, bool do_causal, const std::optional& mask_arr, + const std::optional& sinks, const array& d_o, array& d_q, array& d_k, @@ -448,13 +481,29 @@ void sdpa_backward_cudnn( if (mask_arr) { encoder.set_input_array(*mask_arr); } + if (sinks) { + encoder.set_input_array(*sinks); + } // Search cache. - auto cache_key = build_sdpa_cache_key(encoder, q, k, v, do_causal, mask_arr); + auto cache_key = + build_sdpa_cache_key(encoder, q, k, v, do_causal, mask_arr, sinks); auto it = sdpa_backward_cache().find(cache_key); if (it == sdpa_backward_cache().end()) { auto graph = build_sdpa_backward_graph( - handle, q, k, v, do_causal, mask_arr, o, d_o, stats, d_q, d_k, d_v); + handle, + q, + k, + v, + do_causal, + mask_arr, + sinks, + o, + d_o, + stats, + d_q, + d_k, + d_v); it = sdpa_backward_cache().emplace(cache_key, std::move(graph)).first; } auto& graph = it->second; @@ -473,6 +522,9 @@ void sdpa_backward_cudnn( if (mask_arr) { variant_pack[BIAS] = gpu_ptr(*mask_arr); } + if (sinks) { + variant_pack[SINKS] = gpu_ptr(*sinks); + } CHECK_CUDNN_FE_ERROR(graph.encode_graph(encoder, std::move(variant_pack))); } @@ -536,12 +588,19 @@ void ScaledDotProductAttention::eval_gpu( if (has_arr_mask) { mask_arr = prepare_sdpa_input(inputs[3], s); } + std::optional sinks; + if (has_sinks_) { + sinks = inputs.back(); + } std::optional stats; if (output_logsumexp_) { stats = outputs[1]; } if (supports_sdpa_cudnn(q, k, v, has_arr_mask, do_causal_, s)) { + if (sinks) { + sinks = prepare_sdpa_sinks(*sinks, s); + } sdpa_cudnn( q, k, @@ -551,14 +610,11 @@ void ScaledDotProductAttention::eval_gpu( stats, do_causal_, mask_arr, + sinks, output_logsumexp_, s); } else { - if (has_sinks_) { - sdpa_vector(q, k, v, scale_, out, do_causal_, inputs.back(), s); - } else { - sdpa_vector(q, k, v, scale_, out, do_causal_, std::nullopt, s); - } + sdpa_vector(q, k, v, scale_, out, do_causal_, sinks, s); } } @@ -593,6 +649,10 @@ void ScaledDotProductAttentionVJP::eval_gpu( if (has_arr_mask) { mask_arr = prepare_sdpa_input(inputs[3], s); } + std::optional sinks; + if (has_sinks_) { + sinks = prepare_sdpa_sinks(inputs.back(), s); + } assert(outputs.size() == 3); auto& d_q = outputs[0]; @@ -600,7 +660,20 @@ void ScaledDotProductAttentionVJP::eval_gpu( auto& d_v = outputs[2]; sdpa_backward_cudnn( - q, k, v, scale_, o, stats, do_causal_, mask_arr, d_o, d_q, d_k, d_v, s); + q, + k, + v, + scale_, + o, + stats, + do_causal_, + mask_arr, + sinks, + d_o, + d_q, + d_k, + d_v, + s); } } // namespace fast diff --git a/python/tests/test_fast_sdpa.py b/python/tests/test_fast_sdpa.py index 1a3bd167..7606373c 100644 --- a/python/tests/test_fast_sdpa.py +++ b/python/tests/test_fast_sdpa.py @@ -544,19 +544,24 @@ class TestFastSDPA(mlx_tests.MLXTestCase): with self.assertRaises(ValueError): mx.fast.scaled_dot_product_attention(q, k, v, scale=scale, sinks=sinks) - for T_kv in [128, 4096]: - for T_q in [1, 128]: - for N_kv in [2, 8]: - q = mx.random.normal(shape=(B, N_q, T_q, D)) - k = mx.random.normal(shape=(B, N_kv, T_kv, D)) - v = mx.random.normal(shape=(B, N_kv, T_kv, D)) - sinks = 10 * mx.random.normal(shape=(N_q,)) + for T_q, T_kv, N_kv, dtype in product( + (1, 128), + (128, 4096), + (2, 8), + (mx.float16, mx.float32), + ): + with self.subTest(T_q=T_q, T_kv=T_kv, N_kv=N_kv, dtype=dtype): + q = mx.random.normal(shape=(B, N_q, T_q, D), dtype=dtype) + k = mx.random.normal(shape=(B, N_kv, T_kv, D), dtype=dtype) + v = mx.random.normal(shape=(B, N_kv, T_kv, D), dtype=dtype) + sinks = 10 * mx.random.normal(shape=(N_q,), dtype=dtype) - expected = mlx_ref_attn(q, k, v, scale, sinks=sinks) - out = mx.fast.scaled_dot_product_attention( - q, k, v, scale=scale, sinks=sinks - ) - self.assertTrue(mx.allclose(out, expected, atol=1e-5)) + expected = mlx_ref_attn(q, k, v, scale, sinks=sinks) + out = mx.fast.scaled_dot_product_attention( + q, k, v, scale=scale, sinks=sinks + ) + atol = 1e-5 if dtype == mx.float32 else 1e-2 + self.assertTrue(mx.allclose(out, expected, atol=atol)) def test_sdpa_grad(self): # High tolerance due to cuDNN SDPA kernel requiring tf32.