From 38e2679c64073f1e67716ea7429b0b4b1714af48 Mon Sep 17 00:00:00 2001 From: mlx-dev Date: Fri, 25 Sep 2026 01:59:48 -0700 Subject: [PATCH 1/7] Add SDPA VJP - split from gdn-sdpa-vjp branch --- benchmarks/python/sdpa_bench.py | 119 ++++- mlx/backend/metal/kernels/CMakeLists.txt | 1 + .../scaled_dot_product_attention_vjp.metal | 156 ++++++ .../steel/attn/kernels/steel_attention.h | 24 +- .../steel/attn/kernels/steel_attention_nax.h | 35 ++ .../attn/kernels/steel_attention_nax.metal | 2 + .../metal/scaled_dot_product_attention.cpp | 495 +++++++++++++++++- python/tests/test_fast_sdpa.py | 60 +-- 8 files changed, 834 insertions(+), 58 deletions(-) create mode 100644 mlx/backend/metal/kernels/scaled_dot_product_attention_vjp.metal diff --git a/benchmarks/python/sdpa_bench.py b/benchmarks/python/sdpa_bench.py index 7dfc7e0d1d..d420b9b473 100644 --- a/benchmarks/python/sdpa_bench.py +++ b/benchmarks/python/sdpa_bench.py @@ -155,13 +155,100 @@ def bench_shape( return time_mlx_fused, time_mlx_unfused +def set_vjp_fallback(enabled): + os.environ["MLX_SDPA_VJP_FALLBACK"] = "1" if enabled else "0" + + +def mlx_fused_attn_grads(q, k, v, scale, cotan, mask=None, transpose=False): + def f(q_, k_, v_): + return do_attention(mlx_fused_attn, q_, k_, v_, scale, mask, transpose) + + _, grads = mx.vjp(f, [q, k, v], [cotan]) + return grads + + +def do_attention_vjp_bench(q, k, v, scale, cotan, mask=None, transpose=False): + def f(q_, k_, v_): + return do_attention(mlx_fused_attn, q_, k_, v_, scale, mask, transpose) + + dq = q + + for i in range(N_iter_func): + _, (dq, dk, dv) = mx.vjp(f, [dq, k, v], [cotan]) + + mx.eval([dq, dk, dv]) + return dq + + +def peak_mem_vjp(q, k, v, scale, cotan, mask=None, transpose=False): + mx.clear_cache() + mx.reset_peak_memory() + grads = mlx_fused_attn_grads(q, k, v, scale, cotan, mask, transpose) + mx.eval(grads) + peak = mx.get_peak_memory() / float(1024.0**3) + del grads + mx.clear_cache() + return peak + + +def max_rel_diff(a_list, b_list): + rel = 0.0 + for a, b in zip(a_list, b_list): + denom = mx.maximum(mx.max(mx.abs(a)), mx.array(1e-6, mx.float32)) + rel = max(rel, (mx.max(mx.abs(a - b)) / denom).item()) + return rel + + +def bench_shape_vjp( + B, qsl, ksl, head_dim, n_q_heads, n_kv_heads, dtype, transpose=True, mask_in=None +): + q_mx, k_mx, v_mx, scale, mask = prepare_inputs( + B, qsl, ksl, head_dim, n_q_heads, n_kv_heads, mask_in, transpose, dtype + ) + cotan = mx.array(np.random.normal(0.0, 1.0, q_mx.shape).astype(getattr(np, dtype))) + mx.eval(cotan) + + set_vjp_fallback(True) + time_fallback = bench( + do_attention_vjp_bench, q_mx, k_mx, v_mx, scale, cotan, mask, transpose + ) + mem_fallback = peak_mem_vjp(q_mx, k_mx, v_mx, scale, cotan, mask, transpose) + g_fallback = mlx_fused_attn_grads(q_mx, k_mx, v_mx, scale, cotan, mask, transpose) + mx.eval(g_fallback) + + set_vjp_fallback(False) + time_vjp = bench( + do_attention_vjp_bench, q_mx, k_mx, v_mx, scale, cotan, mask, transpose + ) + mem_vjp = peak_mem_vjp(q_mx, k_mx, v_mx, scale, cotan, mask, transpose) + g_vjp = mlx_fused_attn_grads(q_mx, k_mx, v_mx, scale, cotan, mask, transpose) + mx.eval(g_vjp) + + rel = max_rel_diff(g_fallback, g_vjp) + # nax truncates so the accuracy is not full fp32 + atol = 5e-3 + + if rel > atol: + print( + f"Failed at (B: {B}, qsl: {qsl}, ksl: {ksl}, head_dim: {head_dim}, n_qh: {n_q_heads}, n_kvh: {n_kv_heads}, mask: {mask_in}) [tpose = {transpose}] with max rel = {rel:3.2e}" + ) + + return time_vjp, time_fallback, mem_vjp, mem_fallback, rel + + def get_gflop_count(B, M, N, K): return float(2.0 * N_iter_bench * N_iter_func * B * M * N * K) / float(1024.0**3) if __name__ == "__main__": parser = argparse.ArgumentParser(description="Run gemm benchmarks") - + parser.add_argument( + "-bw", + "--backward", + action="store_true", + help="benchmark the vjp against the fallback", + ) + args = parser.parse_args() dtypes = ("float16", "float32")[:1] transposes = (False,) @@ -228,6 +315,36 @@ def get_gflop_count(B, M, N, K): shapes = shapes_64 + shapes_72 + shapes_80 + shapes_96 + shapes_128 + shapes_256 + if args.backward: + masks = [None, "causal"] + + print( + " B, qsl, ksl, hdim, n_qh, n_kvh, t, dtype, mask, t_fall, t_vjp, speedup, m_fall, m_vjp, mem_x, rel" + ) + + for dtype in ["float32"]: + for transpose in transposes: + for B, qsl, ksl, head_dim, n_q_heads, n_kv_heads in shapes: + for mask_in in masks: + t_vjp, t_fall, m_vjp, m_fall, rel = bench_shape_vjp( + B, + qsl, + ksl, + head_dim, + n_q_heads, + n_kv_heads, + dtype, + transpose, + mask_in, + ) + speedup = t_fall / t_vjp + mem_x = m_fall / max(m_vjp, 1e-9) + t_str = 1 if transpose else 0 + print( + f"{B:3d}, {qsl:5d}, {ksl:5d}, {head_dim:4d}, {n_q_heads:4d}, {n_kv_heads:5d}, {t_str:1d}, {dtype}, {str(mask_in):>8}, {t_fall: 2.3f}, {t_vjp: 2.3f}, {speedup:6.2f}x, {m_fall:6.2f}, {m_vjp:6.2f}, {mem_x:5.2f}x, {rel:3.1e}" + ) + exit(0) + masks = [None, "bool", "causal"] print( diff --git a/mlx/backend/metal/kernels/CMakeLists.txt b/mlx/backend/metal/kernels/CMakeLists.txt index 00f90ac862..be4a9f77b3 100644 --- a/mlx/backend/metal/kernels/CMakeLists.txt +++ b/mlx/backend/metal/kernels/CMakeLists.txt @@ -55,6 +55,7 @@ build_kernel(random) build_kernel(rms_norm) build_kernel(rope) build_kernel(scaled_dot_product_attention sdpa_vector.h) +build_kernel(scaled_dot_product_attention_vjp) build_kernel(gated_delta_update gated_delta_update.h) if(MLX_METAL_VERSION GREATER_EQUAL 320) build_kernel(fence) diff --git a/mlx/backend/metal/kernels/scaled_dot_product_attention_vjp.metal b/mlx/backend/metal/kernels/scaled_dot_product_attention_vjp.metal new file mode 100644 index 0000000000..aa5d08a058 --- /dev/null +++ b/mlx/backend/metal/kernels/scaled_dot_product_attention_vjp.metal @@ -0,0 +1,156 @@ +// Copyright © 2026 Apple Inc. + +#include "mlx/backend/metal/kernels/utils.h" + +using namespace metal; + +template +[[kernel]] void sdpa_vjp_odo( + const device InT* o [[buffer(0)]], // [B, H, qL, Dv] + const device InT* cot_o [[buffer(1)]], // [B, H, qL, Dv] + device float* odo [[buffer(2)]], // [B, H, qL] + constant int& qL [[buffer(3)]], + uint3 tpg [[thread_position_in_grid]], + uint simd_lane_id [[thread_index_in_simdgroup]]) { + const uint bh = tpg.z; + const int i = int(tpg.y); + + if (i >= qL) { + return; + } + + const size_t row = (size_t)bh * qL + i; + auto o_i = o + row * Dv; + auto co_i = cot_o + row * Dv; + + float acc = 0.0f; + for (int d = int(simd_lane_id); d < Dv; d += 32) { + acc += float(o_i[d]) * float(co_i[d]); + } + acc = simd_sum(acc); + + if (simd_lane_id == 0) { + odo[row] = acc; + } +} + +struct SDPAVJPTileParams { + int bq; + int bk; + int qL; + int i0; + int j0; + float scale; + int diag_off; + int causal; +}; + +template +[[kernel]] void sdpa_vjp_ds( + const device T* S [[buffer(0)]], // [BH, bq, bk] + const device T* dP [[buffer(1)]], // [BH, bq, bk] + const device float* lse [[buffer(2)]], // [B, H, qL] + const device float* odo [[buffer(3)]], // [B, H, qL] + device T* dS [[buffer(4)]], // [BH, bq, bk] + device T* dSt [[buffer(5)]], // [BH, bk, bq] + device T* Pt [[buffer(6)]], // [BH, bk, bq] + const constant SDPAVJPTileParams& p [[buffer(7)]], + uint3 gid [[thread_position_in_grid]]) { + int col = int(gid.x); + int row = int(gid.y); + int bh = int(gid.z); + if (row >= p.bq || col >= p.bk) { + return; + } + + size_t sbase = (size_t(bh) * size_t(p.bq) + size_t(row)) * size_t(p.bk); + size_t tbase = (size_t(bh) * size_t(p.bk) + size_t(col)) * size_t(p.bq); + size_t qi = size_t(bh) * size_t(p.qL) + size_t(p.i0 + row); + + float l = lse[qi]; + float dlt = odo[qi]; + // Fully masked rows carry lse == -inf; exp(-inf - -inf) would be NaN. + bool dead = (l < 0.0f) && metal::isinf(l); + int lim = p.i0 + row + p.diag_off; + + float pv; + if (dead || (p.causal != 0 && (p.j0 + col) > lim)) { + pv = 0.0f; + } else { + pv = metal::fast::exp(static_cast(S[sbase + col]) * p.scale - l); + } + float dsv = pv * (static_cast(dP[sbase + col]) - dlt) * p.scale; + + dS[sbase + col] = static_cast(dsv); + dSt[tbase + row] = static_cast(dsv); + Pt[tbase + row] = static_cast(pv); +} + +template +[[kernel]] void sdpa_vjp_reduce( + const device T* src [[buffer(0)]], + device float* acc [[buffer(1)]], + device T* out [[buffer(2)]], + const constant int& rows [[buffer(3)]], + const constant int& dim [[buffer(4)]], + const constant int& group [[buffer(5)]], + const constant int& acc_rows [[buffer(6)]], + const constant int& row_off [[buffer(7)]], + uint3 gid [[thread_position_in_grid]]) { + int c = int(gid.x); + int row = int(gid.y); + int bh = int(gid.z); + if (c >= dim || row >= rows) { + return; + } + + size_t sbase = size_t(bh) * size_t(group) * size_t(rows) * size_t(dim); + float sum = 0.0f; + for (int g = 0; g < group; ++g) { + size_t si = sbase + + (size_t(g) * size_t(rows) + size_t(row)) * size_t(dim) + size_t(c); + sum += static_cast(src[si]); + } + + size_t ai = + (size_t(bh) * size_t(acc_rows) + size_t(row_off + row)) * size_t(dim) + + size_t(c); + float total = Accum ? acc[ai] + sum : sum; + acc[ai] = total; + out[ai] = static_cast(total); +} + +// Instantiations +#define instantiate_odo(in_type, dv) \ + instantiate_kernel( \ + "sdpa_vjp_odo_" #in_type "_" #dv, sdpa_vjp_odo, in_type, dv) + +#define instantiate_odo_shapes(in_type) \ + instantiate_odo(in_type, 64); \ + instantiate_odo(in_type, 72); \ + instantiate_odo(in_type, 80); \ + instantiate_odo(in_type, 96); \ + instantiate_odo(in_type, 128); \ + instantiate_odo(in_type, 192); \ + instantiate_odo(in_type, 256); + +instantiate_odo_shapes(bfloat16_t); +instantiate_odo_shapes(float16_t); +instantiate_odo_shapes(float); + +#define instantiate_sdpa_vjp_ds(tname, type) \ + instantiate_kernel("sdpa_vjp_ds_" #tname, sdpa_vjp_ds, type) + +instantiate_sdpa_vjp_ds(float, float); +instantiate_sdpa_vjp_ds(float16_t, float16_t); +instantiate_sdpa_vjp_ds(bfloat16_t, bfloat16_t); + +#define instantiate_sdpa_vjp_reduce(tname, type) \ + instantiate_kernel( \ + "sdpa_vjp_reduce_add_" #tname, sdpa_vjp_reduce, type, true) \ + instantiate_kernel( \ + "sdpa_vjp_reduce_set_" #tname, sdpa_vjp_reduce, type, false) + +instantiate_sdpa_vjp_reduce(float, float); +instantiate_sdpa_vjp_reduce(float16_t, float16_t); +instantiate_sdpa_vjp_reduce(bfloat16_t, bfloat16_t); diff --git a/mlx/backend/metal/kernels/steel/attn/kernels/steel_attention.h b/mlx/backend/metal/kernels/steel/attn/kernels/steel_attention.h index 29fa7ba396..2e644fb6f1 100644 --- a/mlx/backend/metal/kernels/steel/attn/kernels/steel_attention.h +++ b/mlx/backend/metal/kernels/steel/attn/kernels/steel_attention.h @@ -14,6 +14,7 @@ constant bool align_K [[function_constant(201)]]; constant bool has_mask [[function_constant(300)]]; constant bool do_causal [[function_constant(301)]]; constant bool has_sinks [[function_constant(302)]]; +constant bool save_lse [[function_constant(303)]]; struct MaxOp { template @@ -76,6 +77,7 @@ template < const constant AttnMaskParams* mask_params [[buffer(5), function_constant(has_mask)]], const device MaskType* mask [[buffer(6), function_constant(has_mask)]], const device T* sinks [[buffer(7), function_constant(has_sinks)]], + device float* lse [[buffer(8), function_constant(save_lse)]], uint simd_lane_id [[thread_index_in_simdgroup]], uint simd_group_id [[simdgroup_index_in_threadgroup]], uint3 tid [[threadgroup_position_in_grid]], @@ -102,6 +104,12 @@ template < tidl.y * params->O_strides[1] + // Head tidl.x * BQ * params->O_strides[2]; // Sequence + if (save_lse) { + lse += tidl.z * params->H * params->qL + // Batch + tidl.y * params->qL + // Head + tidl.x * BQ; // Sequence + } + if (has_mask) { mask += tidl.z * mask_params->M_strides[0] + // Batch tidl.y * mask_params->M_strides[1]; // Head @@ -460,6 +468,20 @@ template < Otile.template row_bin_op(sum_score); threadgroup_barrier(mem_flags::mem_none); + // Store the logsumexp for the backward pass + if (save_lse && sn == 0) { + using stile_t = decltype(Stile); + const bool is_last_q = int(tid.x) == (params->NQ_aligned); + + STEEL_PRAGMA_UNROLL + for (short i = 0; i < kRowsPT; ++i) { + const short r = tm + sm + i * stile_t::kFragRows; + if (align_Q || !is_last_q || r < params->qL_rem) { + lse[r] = M_LN2_F * (max_score[i] + metal::log2(sum_score[i])); + } + } + } + // Store results O += (tm + sm) * params->O_strides[2] + sn; @@ -473,4 +495,4 @@ template < } else { Otile.template store(O, params->O_strides[2]); } -} +} \ No newline at end of file diff --git a/mlx/backend/metal/kernels/steel/attn/kernels/steel_attention_nax.h b/mlx/backend/metal/kernels/steel/attn/kernels/steel_attention_nax.h index 4a5a9716fd..6166368db3 100644 --- a/mlx/backend/metal/kernels/steel/attn/kernels/steel_attention_nax.h +++ b/mlx/backend/metal/kernels/steel/attn/kernels/steel_attention_nax.h @@ -17,6 +17,7 @@ constant bool align_K [[function_constant(201)]]; constant bool has_mask [[function_constant(300)]]; constant bool do_causal [[function_constant(301)]]; constant bool has_sinks [[function_constant(302)]]; +constant bool save_lse [[function_constant(303)]]; template struct TransformScale { @@ -89,6 +90,7 @@ template < const constant AttnMaskParams* mask_params [[buffer(5), function_constant(has_mask)]], const device MaskType* mask [[buffer(6), function_constant(has_mask)]], const device T* sinks [[buffer(7), function_constant(has_sinks)]], + device float* lse [[buffer(8), function_constant(save_lse)]], uint simd_lane_id [[thread_index_in_simdgroup]], uint simd_group_id [[simdgroup_index_in_threadgroup]], uint3 tid [[threadgroup_position_in_grid]], @@ -116,6 +118,13 @@ template < tidl.y * params->O_strides[1] + // Head tidl.x * BQ * params->O_strides[2]; // Sequence + if (save_lse) { + // [B, H, qL], contiguous. + lse += tidl.z * params->H * params->qL + // Batch + tidl.y * params->qL + // Head + tidl.x * BQ; // Sequence + } + if (has_mask) { mask += tidl.z * mask_params->M_strides[0] + // Batch tidl.y * mask_params->M_strides[1]; // Head @@ -472,6 +481,16 @@ template < Otile.template row_bin_op(rcp); + if (save_lse && sn == 0) { + STEEL_PRAGMA_UNROLL + for (short i = 0; i < kRowsPT; ++i) { + const short r = tm + i * otile_t::kFragRowsJump + sm; + if (align_Q || !is_last_q || r < params->qL_rem) { + lse[r] = max_score[i] * M_LN2_F + metal::log(sum_score[i]); + } + } + } + // Store results O += tm * int(params->O_strides[2]); @@ -518,6 +537,7 @@ template < const constant AttnMaskParams* mask_params [[buffer(5), function_constant(has_mask)]], const device MaskType* mask [[buffer(6), function_constant(has_mask)]], const device T* sinks [[buffer(7), function_constant(has_sinks)]], + device float* LSE [[buffer(8), function_constant(save_lse)]], uint simd_lane_id [[thread_index_in_simdgroup]], uint simd_group_id [[simdgroup_index_in_threadgroup]], uint3 tid [[threadgroup_position_in_grid]], @@ -895,6 +915,21 @@ template < rcp[i] = 1.f / sum_score[i]; } + if (save_lse && d_half == 0 && sn == 0) { + const int lse_row_base = int(tid.x) * BQ + tm; + const int lse_head_off = (int(tid.z) * params->H + int(tid.y)) * params->qL; + + STEEL_PRAGMA_UNROLL + for (short i = 0; i < kRowsPT; ++i) { + const int row = lse_row_base + i * stile_t::kFragRowsJump + sm; + if (row < params->qL) { + LSE[lse_head_off + row] = (sum_score[i] == 0) + ? -INFINITY + : float(M_LN2_F * (max_score[i] + metal::log2(sum_score[i]))); + } + } + } + Otile.template row_bin_op(rcp); if (!align_Q && is_last_q) { diff --git a/mlx/backend/metal/kernels/steel/attn/kernels/steel_attention_nax.metal b/mlx/backend/metal/kernels/steel/attn/kernels/steel_attention_nax.metal index 66d55539ab..cc439d8285 100644 --- a/mlx/backend/metal/kernels/steel/attn/kernels/steel_attention_nax.metal +++ b/mlx/backend/metal/kernels/steel/attn/kernels/steel_attention_nax.metal @@ -19,12 +19,14 @@ #define instantiate_attn_shapes_helper(iname, itype, mname, mtype) \ instantiate_attn_dsplit(iname, itype, 64, 32, 256, 4, 2, mname, mtype) \ + instantiate_attn(iname, itype, 64, 32, 256, 4, 1, mname, mtype) \ instantiate_attn(iname, itype, 64, 32, 128, 4, 1, mname, mtype) \ instantiate_attn(iname, itype, 64, 32, 96, 4, 1, mname, mtype) \ instantiate_attn(iname, itype, 64, 32, 64, 4, 1, mname, mtype) \ instantiate_attn(iname, itype, 64, 64, 128, 4, 1, mname, mtype) \ instantiate_attn(iname, itype, 64, 64, 64, 4, 1, mname, mtype) + #define instantiate_attn_mask_helper(iname, itype) \ instantiate_attn_shapes_helper(iname, itype, iname, itype) \ instantiate_attn_shapes_helper(iname, itype, bool_, bool) diff --git a/mlx/backend/metal/scaled_dot_product_attention.cpp b/mlx/backend/metal/scaled_dot_product_attention.cpp index 6ece8c43fb..a1559951fd 100644 --- a/mlx/backend/metal/scaled_dot_product_attention.cpp +++ b/mlx/backend/metal/scaled_dot_product_attention.cpp @@ -7,6 +7,7 @@ #include "mlx/backend/metal/kernels.h" #include "mlx/backend/metal/kernels/defines.h" #include "mlx/backend/metal/kernels/steel/attn/params.h" +#include "mlx/backend/metal/matmul.h" #include "mlx/backend/metal/utils.h" #include "mlx/fast_primitives.h" #include "mlx/utils.h" @@ -25,14 +26,21 @@ void sdpa_full_self_attention_nax( array& o, bool do_causal_, const std::optional& mask, - const std::optional& sinks) { + const std::optional& sinks, + array* lse) { using namespace mlx::steel; int bd = q.shape(-1); int bq = 64; int bk = 32; - bool split_d = bd == 256; + // The d-split kernel gives wn warps the same query rows and a slice of the + // head dim each, so every one of them runs the softmax bookkeeping for those + // rows. The sink seeds sum_score with exp(z - z) == 1, which then enters the + // denominator once per warp instead of once per row, and the logsumexp store + // lands before the cross-warp combine so each warp writes a partial. Keep + // sinks and lse on the single-warp path until both move after the combine. + bool split_d = (bd == 256) && !sinks.has_value() && (lse == nullptr); int wm = 4; int wn = split_d ? 2 : 1; @@ -79,13 +87,15 @@ void sdpa_full_self_attention_nax( const bool has_mask = mask.has_value(); const bool do_causal = do_causal_; const bool has_sinks = sinks.has_value(); + const bool save_lse = (lse != nullptr); metal::MTLFCList func_consts = { {&align_Q, MTL::DataType::DataTypeBool, 200}, {&align_K, MTL::DataType::DataTypeBool, 201}, {&has_mask, MTL::DataType::DataTypeBool, 300}, {&do_causal, MTL::DataType::DataTypeBool, 301}, - {&has_sinks, MTL::DataType::DataTypeBool, 302}}; + {&has_sinks, MTL::DataType::DataTypeBool, 302}, + {&save_lse, MTL::DataType::DataTypeBool, 303}}; std::string base_name; concatenate( @@ -118,7 +128,9 @@ void sdpa_full_self_attention_nax( "_do_causal_", (do_causal ? 't' : 'n'), "_has_sinks_", - (has_sinks ? 't' : 'n')); + (has_sinks ? 't' : 'n'), + "_save_lse_", + (save_lse ? 't' : 'n')); auto& compute_encoder = metal::get_command_encoder(s); @@ -188,6 +200,9 @@ void sdpa_full_self_attention_nax( if (has_sinks) { compute_encoder.set_input_array(*sinks, 7); } + if (save_lse) { + compute_encoder.set_output_array(*lse, 8); + } MTL::Size grid_dims = MTL::Size(NQ, H, B); MTL::Size group_dims = MTL::Size(32, wm, wn); @@ -206,15 +221,12 @@ void sdpa_full_self_attention_metal( array& o, bool do_causal_, const std::optional& mask, - const std::optional& sinks) { - int B = q.shape(0); - int H = q.shape(1); + const std::optional& sinks, + array* lse /* = nullptr */) { int D = q.shape(3); - int gqa_factor = q.shape(1) / k.shape(1); int qL = q.shape(2); int kL = k.shape(2); - if (metal::is_nax_available() && (D == 64 || D == 96 || D == 128 || D == 256) && (env::enable_tf32() || q.dtype() != float32)) { @@ -228,15 +240,18 @@ void sdpa_full_self_attention_metal( /* array& o = */ o, /* bool do_causal_ = */ do_causal_, /* const std::optional& mask = */ mask, - /* const std::optional& sinks = */ sinks); + /* const std::optional& sinks = */ sinks, + /* array* lse = */ lse); } // Pad head dims 72 and 80 to 96 to reach the NAX kernel. The added lanes // are zero and the caller's scale is retained. Enable by default only for // long, unmasked half-precision attention, where the attention work can // amortize the padding and output copy. The override is read per call. + // + // Skipped when the logsumexp is needed bool pad_default = qL >= 512 && kL >= 512 && !do_causal_ && !mask && !sinks; - if ((D == 72 || D == 80) && metal::is_nax_available() && + if ((D == 72 || D == 80) && lse == nullptr && metal::is_nax_available() && (q.dtype() == float16 || q.dtype() == bfloat16) && env::get_var("MLX_SDPA_PAD_HEAD_DIM", pad_default ? 1 : 0) == 1) { constexpr int pad_to = 96; @@ -283,7 +298,8 @@ void sdpa_full_self_attention_metal( /* array& o = */ op, /* bool do_causal_ = */ do_causal_, /* const std::optional& mask = */ mask, - /* const std::optional& sinks = */ sinks); + /* const std::optional& sinks = */ sinks, + /* array* lse = */ nullptr); // Slice the padded lanes back out into o's caller-chosen strides. copy_gpu_inplace( @@ -303,6 +319,10 @@ void sdpa_full_self_attention_metal( using namespace mlx::steel; + int B = q.shape(0); + int H = q.shape(1); + int gqa_factor = q.shape(1) / k.shape(1); + int wm = 4; int wn = 1; @@ -315,13 +335,15 @@ void sdpa_full_self_attention_metal( const bool has_mask = mask.has_value(); const bool do_causal = do_causal_; const bool has_sinks = sinks.has_value(); + const bool save_lse = (lse != nullptr); metal::MTLFCList func_consts = { {&align_Q, MTL::DataType::DataTypeBool, 200}, {&align_K, MTL::DataType::DataTypeBool, 201}, {&has_mask, MTL::DataType::DataTypeBool, 300}, {&do_causal, MTL::DataType::DataTypeBool, 301}, - {&has_sinks, MTL::DataType::DataTypeBool, 302}}; + {&has_sinks, MTL::DataType::DataTypeBool, 302}, + {&save_lse, MTL::DataType::DataTypeBool, 303}}; std::string base_name; concatenate( @@ -354,7 +376,9 @@ void sdpa_full_self_attention_metal( "_do_causal_", (do_causal ? 't' : 'n'), "_has_sinks_", - (has_sinks ? 't' : 'n')); + (has_sinks ? 't' : 'n'), + "_save_lse_", + (save_lse ? 't' : 'n')); auto& compute_encoder = metal::get_command_encoder(s); @@ -423,6 +447,9 @@ void sdpa_full_self_attention_metal( if (has_sinks) { compute_encoder.set_input_array(*sinks, 7); } + if (save_lse) { + compute_encoder.set_output_array(*lse, 8); + } MTL::Size grid_dims = MTL::Size(NQ, H, B); MTL::Size group_dims = MTL::Size(32, wm, wn); @@ -716,11 +743,12 @@ std::tuple has_fused_kernel( if (s.device != Device::gpu) { return {false, "the fused kernels require a GPU (Metal) stream."}; } - if (output_logsumexp) { + + if (output_logsumexp && has_arr_mask) { return { false, - "the fused forward does not produce the logsumexp required for " - "the fused VJP; use default routing when training."}; + "the backward pass does not support an array mask; only causal masking " + "is implemented."}; } const int value_head_dim = v.shape(-1); @@ -814,6 +842,346 @@ std::tuple has_fused_kernel( return {true, ""}; } +/////////////////////////////////////////////////////////////////////////////// +// Backward pass +/////////////////////////////////////////////////////////////////////////////// + +struct SDPAVJPTileParams { + int bq; + int bk; + int qL; + int i0; + int j0; + float scale; + int diag_off; + int causal; +}; + +array ensure_row_contiguous(const array& x, metal::Device& d, const Stream& s) { + if (!x.flags().row_contiguous) { + array x_copy = contiguous_copy_gpu(x, s); + metal::get_command_encoder(s).add_temporary(x_copy); + return x_copy; + } else { + return x; + } +} + +array vjp_alloc(Shape shape, Dtype dt, const Stream& s) { + array a(std::move(shape), dt, nullptr, {}); + a.set_data(allocator::malloc(a.nbytes())); + metal::get_command_encoder(s).add_temporary(a); + return a; +} + +// Takes a view of the buffer +array vjp_view(const array& base, Shape shape) { + array v(shape, base.dtype(), nullptr, {}); + Strides st(shape.size()); + int64_t acc = 1; + for (int i = static_cast(shape.size()) - 1; i >= 0; --i) { + st[i] = acc; + acc *= shape[i]; + } + array::Flags f{1, 1, shape.size() <= 1}; + v.copy_shared_buffer(base, st, f, static_cast(acc), 0); + return v; +} + +// A [B, H, len, D] window of a row-contiguous [B, H, T, D] input, starting at +// sequence position t0. +array vjp_row_slice(const array& x, int t0, int len) { + const auto& shp = x.shape(); + int64_t Dx = shp[3]; + int64_t Tx = shp[2]; + Shape ns = {shp[0], shp[1], len, static_cast(Dx)}; + Strides st = {shp[1] * Tx * Dx, Tx * Dx, Dx, 1}; + array v(ns, x.dtype(), nullptr, {}); + array::Flags f{1, len == shp[2], false}; + size_t base = static_cast(x.offset()) / x.itemsize(); + v.copy_shared_buffer( + x, + st, + f, + x.data_size(), + base + static_cast(t0) * static_cast(Dx)); + return v; +} + +// Block sizes for the score tile. The best results are usually for 1024x1024 so +// I set those as default. This probably depends on the hardware so it may need +// to be tuned in the future. +std::pair sdpa_vjp_blocks(int B, int H, int qL, int kL, Dtype ctype) { + constexpr int kDefaultBlock = 1024; + + int blk = env::get_var("MLX_SDPA_VJP_BLOCK", kDefaultBlock); + + int bq = std::min(blk, qL); + int bk = std::min(blk, kL); + if (int e = env::get_var("MLX_SDPA_VJP_BQ", 0); e > 0) { + bq = std::min(e, qL); + } + if (int e = env::get_var("MLX_SDPA_VJP_BK", 0); e > 0) { + bk = std::min(e, kL); + } + return {bq, bk}; +} + +void sdpa_vjp_blocked( + const Stream& s, + metal::Device& d, + const array& q, + const array& k, + const array& v, + const array& odo, + const array& lse, + const array& cot_o, + array& dq, + array& dk, + array& dv, + float scale, + bool causal) { + auto& compute_encoder = metal::get_command_encoder(s); + + const int B = q.shape(0); + const int H = q.shape(1); + const int qL = q.shape(2); + const int D = q.shape(3); + const int Hk = k.shape(1); + const int kL = k.shape(2); + const int Dv = v.shape(3); + const int G = H / Hk; + const int BH = B * H; + const int BHk = B * Hk; + // Under a causal mask the diagonal sits at j == i + (kL - qL), which is + // nonzero whenever the queries are a suffix of the keys. + const int diag_off = kL - qL; + + auto blocks = sdpa_vjp_blocks(B, H, qL, kL, q.dtype()); + const int BQ = blocks.first; + const int BK = blocks.second; + + Dtype ctype = q.dtype(); + std::string tname = get_type_string(ctype); + + auto ds_kernel = d.get_kernel("sdpa_vjp_ds_" + tname); + auto red_add = d.get_kernel("sdpa_vjp_reduce_add_" + tname); + auto red_set = d.get_kernel("sdpa_vjp_reduce_set_" + tname); + + // Temporary tiles + array s_buf = vjp_alloc({BH, BQ, BK}, ctype, s); + array dp_buf = vjp_alloc({BH, BQ, BK}, ctype, s); + array dst_buf = vjp_alloc({BH, BK, BQ}, ctype, s); + array pt_buf = vjp_alloc({BH, BK, BQ}, ctype, s); + + // Output of the gemms + array dq_tile = vjp_alloc({BH, BQ, D}, ctype, s); + array dk_tile = vjp_alloc({BH, BK, D}, ctype, s); + array dv_tile = vjp_alloc({BH, BK, Dv}, ctype, s); + + // Accumulators for the result + array dq_acc = vjp_alloc({B, H, qL, D}, float32, s); + array dk_acc = vjp_alloc({B, Hk, kL, D}, float32, s); + array dv_acc = vjp_alloc({B, Hk, kL, Dv}, float32, s); + + std::vector copies; + + auto run_reduce = [&](const array& src, + array& acc, + array& out, + int rows, + int dim, + int group, + int acc_rows, + int row_off, + int nbh, + bool accum) { + compute_encoder.set_compute_pipeline_state(accum ? red_add : red_set); + compute_encoder.set_input_array(src, 0); + compute_encoder.set_input_array(acc, 1); + compute_encoder.set_output_array(acc, 1); + compute_encoder.set_output_array(out, 2); + compute_encoder.set_bytes(rows, 3); + compute_encoder.set_bytes(dim, 4); + compute_encoder.set_bytes(group, 5); + compute_encoder.set_bytes(acc_rows, 6); + compute_encoder.set_bytes(row_off, 7); + compute_encoder.dispatch_threads( + MTL::Size(dim, rows, nbh), MTL::Size(std::min(dim, 32), 8, 1)); + }; + + const int n_i = (qL + BQ - 1) / BQ; + + const Strides q_bs = { + H * int64_t(qL) * D, G * int64_t(qL) * D, int64_t(qL) * D}; + const Strides o_bs = { + H * int64_t(qL) * Dv, G * int64_t(qL) * Dv, int64_t(qL) * Dv}; + // A zero innermost batch stride broadcasts one kv head over its query group. + const Strides k_bs = {Hk * int64_t(kL) * D, int64_t(kL) * D, 0}; + const Strides v_bs = {Hk * int64_t(kL) * Dv, int64_t(kL) * Dv, 0}; + const Shape bshape = {B, Hk, G}; + + for (int j0 = 0; j0 < kL; j0 += BK) { + const int bk_len = std::min(BK, kL - j0); + + array k_sl = vjp_row_slice(k, j0, bk_len); + array v_sl = vjp_row_slice(v, j0, bk_len); + + bool kv_accum = false; + + for (int ii = 0; ii < n_i; ++ii) { + const int i0 = ii * BQ; + const int bq_len = std::min(BQ, qL - i0); + // Skip tiles that lie entirely above the causal diagonal. + if (causal && j0 > i0 + bq_len - 1 + diag_off) { + continue; + } + + array q_sl = vjp_row_slice(q, i0, bq_len); + array o_sl = vjp_row_slice(cot_o, i0, bq_len); + + array s_v = vjp_view(s_buf, {BH, bq_len, bk_len}); + // S = Q @ K.T + steel_matmul( + s, + d, + q_sl, + k_sl, + s_v, + bq_len, + bk_len, + D, + BH, + D, + D, + false, + true, + copies, + bshape, + q_bs, + k_bs); + + array dp_v = vjp_view(dp_buf, {BH, bq_len, bk_len}); + // dP = dO @ V.T + steel_matmul( + s, + d, + o_sl, + v_sl, + dp_v, + bq_len, + bk_len, + Dv, + BH, + Dv, + Dv, + false, + true, + copies, + bshape, + o_bs, + v_bs); + + array ds_v = vjp_view(s_buf, {BH, bq_len, bk_len}); + array dst_v = vjp_view(dst_buf, {BH, bk_len, bq_len}); + array pt_v = vjp_view(pt_buf, {BH, bk_len, bq_len}); + + SDPAVJPTileParams params{ + bq_len, bk_len, qL, i0, j0, scale, diag_off, causal ? 1 : 0}; + + // P = exp(scale * S - lse), dS = P * (dP - delta) * scale + compute_encoder.set_compute_pipeline_state(ds_kernel); + compute_encoder.set_input_array(s_v, 0); + compute_encoder.set_input_array(dp_v, 1); + compute_encoder.set_input_array(lse, 2); + compute_encoder.set_input_array(odo, 3); + compute_encoder.set_output_array(s_v, 4); + compute_encoder.set_output_array(dst_v, 5); + compute_encoder.set_output_array(pt_v, 6); + compute_encoder.set_bytes(params, 7); + compute_encoder.dispatch_threads( + MTL::Size(bk_len, bq_len, BH), MTL::Size(32, 8, 1)); + + int64_t ss2 = static_cast(bq_len) * bk_len; + Strides sc_bs = {H * ss2, G * ss2, ss2}; + + array dq_v = vjp_view(dq_tile, {BH, bq_len, D}); + // dQ = dS @ K + steel_matmul( + s, + d, + ds_v, + k_sl, + dq_v, + bq_len, + D, + bk_len, + BH, + bk_len, + D, + false, + false, + copies, + bshape, + sc_bs, + k_bs); + + array dk_v = vjp_view(dk_tile, {BH, bk_len, D}); + // dK = dS.T @ Q + steel_matmul( + s, + d, + dst_v, + q_sl, + dk_v, + bk_len, + D, + bq_len, + BH, + bq_len, + D, + false, + false, + copies, + bshape, + sc_bs, + q_bs); + + array dv_v = vjp_view(dv_tile, {BH, bk_len, Dv}); + // dV = P.T @ dO + steel_matmul( + s, + d, + pt_v, + o_sl, + dv_v, + bk_len, + Dv, + bq_len, + BH, + bq_len, + Dv, + false, + false, + copies, + bshape, + sc_bs, + o_bs); + + // dQ[i0] += dQ_tile, dK[j0] += sum_G dK_tile, dV[j0] += sum_G dV_tile + run_reduce(dq_v, dq_acc, dq, bq_len, D, 1, qL, i0, BH, j0 != 0); + run_reduce(dk_v, dk_acc, dk, bk_len, D, G, kL, j0, BHk, kv_accum); + run_reduce(dv_v, dv_acc, dv, bk_len, Dv, G, kL, j0, BHk, kv_accum); + + kv_accum = true; + } + } + + for (auto& c : copies) { + compute_encoder.add_temporary(c); + } +} + } // namespace bool ScaledDotProductAttention::use_fallback( @@ -840,15 +1208,15 @@ bool ScaledDotProductAttention::use_fallback( return false; } - if (is_training) { - // It's faster for training on Metal to use the unfused SDPA for both - // forward and backward. - return true; - } if (!has_fused) { return true; } + // The logsumexp only comes out of the fused forward + if (output_logsumexp) { + return false; + } + const int query_sequence_length = q.shape(2); const int query_head_dim = q.shape(-1); const int value_head_dim = v.shape(-1); @@ -886,7 +1254,6 @@ void ScaledDotProductAttention::eval_gpu( auto& o = outputs[0]; std::vector copies; - // Define some copy functions to ensure the layout of the inputs is as // expected. copies.reserve(inputs.size()); @@ -914,6 +1281,12 @@ void ScaledDotProductAttention::eval_gpu( // We are in vector mode ie single query if (q_pre.shape(2) <= 8) { + if (outputs.size() > 1) { + throw std::runtime_error( + "[scaled_dot_product_attention] the vector kernels do not produce a " + "logsumexp; the VJP path requires the full-attention kernel."); + } + auto q_copy_unless = [](const array& arr) { if (arr.flags().row_contiguous) { return true; @@ -1012,21 +1385,91 @@ void ScaledDotProductAttention::eval_gpu( ? std::optional{copy_unless(is_matrix_contiguous, inputs[3])} : std::nullopt; + array* lse = nullptr; + if (outputs.size() > 1) { + outputs[1].set_data(allocator::malloc(outputs[1].nbytes())); + lse = &outputs[1]; + } + sdpa_full_self_attention_metal( - s, d, q, k, v, scale_, o, do_causal_, mask, sinks); + s, d, q, k, v, scale_, o, do_causal_, mask, sinks, lse); } metal::get_command_encoder(s).add_temporaries(std::move(copies)); } +// The vjp uses matmuls directly so it should be fine to call it every time as +// it uses a lot less memory. It can be overriden by setting +// MLX_SDPA_VJP_FALLBACK=1 bool ScaledDotProductAttentionVJP::use_fallback(const array& q, Stream s) { - return true; + if (s.device != Device::gpu) { + return true; + } + if (env::get_var("MLX_SDPA_VJP_FALLBACK", 0) != 0) { + return true; + } + auto dt = q.dtype(); + return !(dt == float32 || dt == float16 || dt == bfloat16); } void ScaledDotProductAttentionVJP::eval_gpu( const std::vector& inputs, std::vector& outputs) { - throw std::runtime_error("NYI"); + auto& s = stream(); + auto& d = metal::device(s.device); + + if (has_sinks_) { + throw std::runtime_error( + "[ScaledDotProductAttentionVJP] NYI: attention sinks"); + } + + auto q = ensure_row_contiguous(inputs[0], d, s); + auto k = ensure_row_contiguous(inputs[1], d, s); + auto v = ensure_row_contiguous(inputs[2], d, s); + const int n_in = static_cast(inputs.size()); + auto o = ensure_row_contiguous(inputs[n_in - 3], d, s); + auto lse = ensure_row_contiguous(inputs[n_in - 2], d, s); + auto cot_o = ensure_row_contiguous(inputs[n_in - 1], d, s); + + const int B = q.shape(0); + const int H = q.shape(1); + const int qL = q.shape(2); + const int Dv = v.shape(3); + + auto& dq = outputs[0]; + auto& dk = outputs[1]; + auto& dv = outputs[2]; + + auto& compute_encoder = metal::get_command_encoder(s); + + dq.set_data(allocator::malloc(dq.nbytes())); + dk.set_data(allocator::malloc(dk.nbytes())); + dv.set_data(allocator::malloc(dv.nbytes())); + + array odo({B, H, qL}, float32, nullptr, {}); + odo.set_data(allocator::malloc(odo.nbytes())); + compute_encoder.add_temporary(odo); + + std::string odo_name; + concatenate( + odo_name, + "sdpa_vjp_odo_", + get_type_string(q.dtype()), + "_", + std::to_string(Dv)); + + auto odo_kernel = d.get_kernel(odo_name); + + compute_encoder.set_compute_pipeline_state(odo_kernel); + compute_encoder.set_input_array(o, 0); + compute_encoder.set_input_array(cot_o, 1); + compute_encoder.set_output_array(odo, 2); + compute_encoder.set_bytes(qL, 3); + compute_encoder.dispatch_threads( + MTL::Size(32, qL, B * H), MTL::Size(32, 1, 1)); + + sdpa_vjp_blocked( + s, d, q, k, v, odo, lse, cot_o, dq, dk, dv, scale_, do_causal_); } } // namespace mlx::core::fast diff --git a/python/tests/test_fast_sdpa.py b/python/tests/test_fast_sdpa.py index fefe7325b3..fd13814b84 100644 --- a/python/tests/test_fast_sdpa.py +++ b/python/tests/test_fast_sdpa.py @@ -136,7 +136,7 @@ def test_sdpa_head_dim_72(self): ) if dtype == mx.float32: - atol = 1e-5 + atol = 7e-5 elif dtype == mx.bfloat16: atol = 5e-3 else: @@ -172,7 +172,7 @@ def test_sdpa_head_dim_80(self): ) if dtype == mx.float32: - atol = 1e-5 + atol = 8e-5 elif dtype == mx.bfloat16: atol = 5e-3 else: @@ -217,7 +217,7 @@ def test_sdpa_head_dim_96(self): ) if dtype == mx.float32: - atol = 1e-5 + atol = 7e-5 elif dtype == mx.bfloat16: atol = 5e-3 else: @@ -287,7 +287,7 @@ def test_sdpa_full_head_dim_256(self): if dtype == mx.float32: # The fused shapes run through tf32 tensor ops when # MLX_ENABLE_TF32 is on (the default). - tol = 1e-3 if qL >= 2048 else 1e-4 + tol = 1e-3 if qL >= 2048 else 5e-4 else: tol = 5e-3 self.assertTrue(mx.allclose(ref, out, atol=tol, rtol=tol)) @@ -322,7 +322,7 @@ def test_sdpa_vector_kv_transposed_head_seq(self): scale=scale, mask=m, ) - self.assertTrue(mx.allclose(ref, out, atol=1e-4, rtol=1e-4)) + self.assertTrue(mx.allclose(ref, out, atol=5e-4, rtol=5e-4)) def test_sdpa_vector(self): D = 64 @@ -365,7 +365,7 @@ def test_sdpa_vector(self): scale=scale, mask=m, ) - self.assertTrue(mx.allclose(ref, out, atol=1e-4, rtol=1e-4)) + self.assertTrue(mx.allclose(ref, out, atol=5e-4, rtol=5e-4)) L = 4096 scale = 1.0 @@ -394,7 +394,7 @@ def test_sdpa_vector(self): scale=scale, mask=m, ) - self.assertTrue(mx.allclose(ref, out, atol=1e-4, rtol=1e-4)) + self.assertTrue(mx.allclose(ref, out, atol=5e-4, rtol=5e-4)) def test_sdpa_vector_gqa_long(self): scale = 1.0 @@ -408,7 +408,7 @@ def test_sdpa_vector_gqa_long(self): vr = mx.repeat(v, Nq // Nkv, axis=1) ref = mlx_primitives_sdpa(q, kr, vr, scale) out = mx.fast.scaled_dot_product_attention(q, k, v, scale=scale) - self.assertTrue(mx.allclose(ref, out, atol=1e-4, rtol=1e-4)) + self.assertTrue(mx.allclose(ref, out, atol=5e-4, rtol=5e-4)) @unittest.skipIf(not mx.metal.is_available(), "Metal kernel path only") def test_sdpa_vector_head_dim_512(self): @@ -431,7 +431,7 @@ def test_sdpa_vector_head_dim_512(self): out = mx.fast.scaled_dot_product_attention( q, k, v, scale=scale, force_fused=True ) - atol = 1e-5 if dtype == mx.float32 else 2e-2 + atol = 7e-5 if dtype == mx.float32 else 2e-2 self.assertTrue(mx.allclose(ref, out, atol=atol)) # Test 2-pass kernel. @@ -446,7 +446,7 @@ def test_sdpa_vector_head_dim_512(self): out = mx.fast.scaled_dot_product_attention( q, k, v, scale=scale, force_fused=True ) - atol = 1e-5 if dtype == mx.float32 else 2e-2 + atol = 7e-5 if dtype == mx.float32 else 2e-2 self.assertTrue(mx.allclose(ref, out, atol=atol)) # Test other heads. @@ -463,7 +463,7 @@ def test_sdpa_vector_head_dim_512(self): scale=scale, force_fused=True, ) - self.assertTrue(mx.allclose(ref, out, atol=1e-4, rtol=1e-4)) + self.assertTrue(mx.allclose(ref, out, atol=5e-4, rtol=5e-4)) # Test batched. B = 2 @@ -483,7 +483,7 @@ def test_sdpa_vector_head_dim_512(self): sinks=s_in, force_fused=True, ) - self.assertTrue(mx.allclose(ref, out, atol=1e-4, rtol=1e-4)) + self.assertTrue(mx.allclose(ref, out, atol=5e-4, rtol=5e-4)) def test_sdpa_fully_masked(self): Lkv = 8 @@ -507,7 +507,7 @@ def test_sdpa_inf_score(self): k[..., 0, :] = -float("inf") ref = mlx_primitives_sdpa(q, k, v, scale=1, mask=None) out = mx.fast.scaled_dot_product_attention(q, k, v, mask=None, scale=1) - self.assertTrue(mx.allclose(ref, out, atol=1e-4, rtol=1e-4)) + self.assertTrue(mx.allclose(ref, out, atol=5e-4, rtol=5e-4)) def test_sdpa_few_query(self): D = 64 @@ -539,7 +539,7 @@ def test_sdpa_few_query(self): scale=scale, mask=m, ) - self.assertTrue(mx.allclose(ref, out, atol=1e-4, rtol=1e-4)) + self.assertTrue(mx.allclose(ref, out, atol=5e-4, rtol=5e-4)) L = 4096 scale = 1.0 @@ -565,7 +565,7 @@ def test_sdpa_few_query(self): scale=scale, mask=m, ) - self.assertTrue(mx.allclose(ref, out, atol=1e-4, rtol=1e-4)) + self.assertTrue(mx.allclose(ref, out, atol=5e-4, rtol=5e-4)) @unittest.skip("Different head and value dims is not enabled") def test_sdpa_vector_value_dims(self): @@ -582,7 +582,7 @@ def test_sdpa_vector_value_dims(self): v = 5e-1 * mx.random.normal(shape=(1, Nkv, L, V)) ref = mlx_primitives_sdpa(q, k, v, scale) out = mx.fast.scaled_dot_product_attention(q, k, v, scale=scale) - self.assertTrue(mx.allclose(ref, out, atol=1e-4, rtol=1e-4)) + self.assertTrue(mx.allclose(ref, out, atol=5e-4, rtol=5e-4)) def test_sdpa_vector_batched(self): D = 64 @@ -592,29 +592,29 @@ def test_sdpa_vector_batched(self): out = mx.fast.scaled_dot_product_attention(q, k, v, mask=None, scale=1.0) ref = mlx_ref_attn(q, k, v) - self.assertTrue(mx.allclose(ref, out, atol=1e-4, rtol=1e-4)) + self.assertTrue(mx.allclose(ref, out, atol=5e-4, rtol=5e-4)) q = mx.random.normal(shape=(2, 4, 3, D)) out = mx.fast.scaled_dot_product_attention(q, k, v, mask=None, scale=1.0) ref = mlx_ref_attn(q, k, v) - self.assertTrue(mx.allclose(ref, out, atol=1e-4, rtol=1e-4)) + self.assertTrue(mx.allclose(ref, out, atol=5e-4, rtol=5e-4)) q = mx.random.normal(shape=(2, 3, 4, D)).swapaxes(1, 2) out = mx.fast.scaled_dot_product_attention(q, k, v, mask=None, scale=1.0) ref = mlx_ref_attn(q, k, v) - self.assertTrue(mx.allclose(ref, out, atol=1e-4, rtol=1e-4)) + self.assertTrue(mx.allclose(ref, out, atol=5e-4, rtol=5e-4)) k = mx.random.normal(shape=(2, 3, 1, D)).swapaxes(1, 2) out = mx.fast.scaled_dot_product_attention(q, k, v, mask=None, scale=1.0) ref = mlx_ref_attn(q, k, v) - self.assertTrue(mx.allclose(ref, out, atol=1e-4, rtol=1e-4)) + self.assertTrue(mx.allclose(ref, out, atol=5e-4, rtol=5e-4)) q = mx.random.normal(shape=(2, 4, 3, D)) k = mx.random.normal(shape=(2, 3, 2, D)).swapaxes(1, 2) v = mx.random.normal(shape=(2, 2, 3, D)) out = mx.fast.scaled_dot_product_attention(q, k, v, mask=None, scale=1.0) ref = mlx_ref_attn(q, k, v) - self.assertTrue(mx.allclose(ref, out, atol=1e-4, rtol=1e-4)) + self.assertTrue(mx.allclose(ref, out, atol=5e-4, rtol=5e-4)) q = mx.random.normal(shape=(2, 4, 3, D)) k = mx.random.normal(shape=(2, 1, 3, D)) @@ -622,7 +622,7 @@ def test_sdpa_vector_batched(self): mask = 10 * mx.random.normal(shape=(1, 2, 3, 3)).swapaxes(0, 1) out = mx.fast.scaled_dot_product_attention(q, k, v, mask=mask, scale=1.0) ref = mlx_ref_attn(q, k, v, mask=mask) - self.assertTrue(mx.allclose(ref, out, atol=1e-4, rtol=1e-4)) + self.assertTrue(mx.allclose(ref, out, atol=5e-4, rtol=5e-4)) @unittest.skipIf(not mx.is_available(mx.gpu), "GPU kernel path only") def test_sdpa_blocks_env_override(self): @@ -637,7 +637,7 @@ def test_sdpa_blocks_env_override(self): for blocks in (16, 33, 48, 100): with mlx_tests.scoped_env(MLX_SDPA_BLOCKS=str(blocks)): out = mx.fast.scaled_dot_product_attention(q, k, v, scale=D**-0.5) - self.assertTrue(mx.allclose(ref, out, atol=1e-4, rtol=1e-4)) + self.assertTrue(mx.allclose(ref, out, atol=5e-4, rtol=5e-4)) @unittest.skipIf(not mx.is_available(mx.gpu), "too slow on CPU") def test_sdpa(self): @@ -715,7 +715,7 @@ def test_sdpa(self): out_ref = out_ref[:, :, offset:, :] out_fst = out_fst[:, :, offset:, :] - atol = 2e-5 if dtype == mx.float32 else 3e-4 + atol = 6e-5 if dtype == mx.float32 else 3e-4 self.assertListEqual(list(out_ref.shape), list(out_fst.shape)) @@ -775,7 +775,7 @@ def test_sdpa_broadcast_mask(self): v = 5e-1 * mx.random.normal(shape=(1, Nkv, L, D)) ref = mlx_primitives_sdpa(q, k, v, scale, mask=mask) out = mx.fast.scaled_dot_product_attention(q, k, v, scale=scale, mask=mask) - self.assertTrue(mx.allclose(ref, out, atol=1e-4, rtol=1e-4)) + self.assertTrue(mx.allclose(ref, out, atol=5e-4, rtol=5e-4)) def test_sdpa_noncontiguous_inputs(self): mask = mx.ones(shape=(4, 1, 7, 7), dtype=mx.bool_) @@ -786,7 +786,7 @@ def test_sdpa_noncontiguous_inputs(self): v = mx.random.normal(shape=(4, 7, 8, 64)).swapaxes(1, 2) out = mx.fast.scaled_dot_product_attention(q, k, v, scale=1.0, mask=mask) ref = mlx_ref_attn(q, k, v, scale=1.0, mask=mask) - self.assertTrue(mx.allclose(ref, out, atol=1e-4, rtol=1e-4)) + self.assertTrue(mx.allclose(ref, out, atol=5e-4, rtol=5e-4)) def test_sdpa_promote_mask(self): mask = mx.array(2.0, mx.bfloat16) @@ -802,7 +802,7 @@ def test_sdpa_promote_mask(self): v = 5e-1 * mx.random.normal(shape=(1, Nkv, L, D)) ref = mlx_primitives_sdpa(q, k, v, scale, mask=mask) out = mx.fast.scaled_dot_product_attention(q, k, v, scale=scale, mask=mask) - self.assertTrue(mx.allclose(ref, out, atol=1e-4, rtol=1e-4)) + self.assertTrue(mx.allclose(ref, out, atol=5e-4, rtol=5e-4)) def test_sdpa_nan_bug(self): N = 128 @@ -822,7 +822,7 @@ def test_sdpa_nan_bug(self): out = mx.fast.scaled_dot_product_attention(q, k, v, mask=mask, scale=1.0) expected = mlx_ref_attn(q, k, v, mask=mask, scale=1.0) self.assertFalse(mx.isnan(out).any().item()) - self.assertLessEqual(mx.abs(out - expected).max().item(), 1e-4) + self.assertLessEqual(mx.abs(out - expected).max().item(), 5e-4) # And an additive one mask = mx.log(mask) @@ -830,7 +830,7 @@ def test_sdpa_nan_bug(self): out = mx.fast.scaled_dot_product_attention(q, k, v, mask=mask, scale=1.0) expected = mlx_ref_attn(q, k, v, mask=mask, scale=1.0) self.assertFalse(mx.isnan(out).any().item()) - self.assertLessEqual(mx.abs(out - expected).max().item(), 1e-4) + self.assertLessEqual(mx.abs(out - expected).max().item(), 5e-4) def test_sdpa_attention_sinks(self): B = 2 @@ -879,7 +879,7 @@ def test_sdpa_attention_sinks(self): out = mx.fast.scaled_dot_product_attention( q, k, v, scale=scale, sinks=sinks ) - atol = 1e-5 if dtype == mx.float32 else 1e-2 + atol = 6e-4 if dtype == mx.float32 else 1e-2 self.assertTrue(mx.allclose(out, expected, atol=atol)) def test_sdpa_grad(self): From b2964e88fc5755a7a7a5bd2289eed874993e6557 Mon Sep 17 00:00:00 2001 From: mlx-dev Date: Fri, 25 Sep 2026 05:19:29 -0700 Subject: [PATCH 2/7] Revert test_fast_sdpa.py --- python/tests/test_fast_sdpa.py | 60 +++++++++++++++++----------------- 1 file changed, 30 insertions(+), 30 deletions(-) diff --git a/python/tests/test_fast_sdpa.py b/python/tests/test_fast_sdpa.py index 04bb618b15..18c64ed77f 100644 --- a/python/tests/test_fast_sdpa.py +++ b/python/tests/test_fast_sdpa.py @@ -136,7 +136,7 @@ def test_sdpa_head_dim_72(self): ) if dtype == mx.float32: - atol = 7e-5 + atol = 1e-5 elif dtype == mx.bfloat16: atol = 5e-3 else: @@ -172,7 +172,7 @@ def test_sdpa_head_dim_80(self): ) if dtype == mx.float32: - atol = 8e-5 + atol = 1e-5 elif dtype == mx.bfloat16: atol = 5e-3 else: @@ -217,7 +217,7 @@ def test_sdpa_head_dim_96(self): ) if dtype == mx.float32: - atol = 7e-5 + atol = 1e-5 elif dtype == mx.bfloat16: atol = 5e-3 else: @@ -287,7 +287,7 @@ def test_sdpa_full_head_dim_256(self): if dtype == mx.float32: # The fused shapes run through tf32 tensor ops when # MLX_ENABLE_TF32 is on (the default). - tol = 1e-3 if qL >= 2048 else 5e-4 + tol = 1e-3 if qL >= 2048 else 1e-4 else: tol = 5e-3 self.assertTrue(mx.allclose(ref, out, atol=tol, rtol=tol)) @@ -322,7 +322,7 @@ def test_sdpa_vector_kv_transposed_head_seq(self): scale=scale, mask=m, ) - self.assertTrue(mx.allclose(ref, out, atol=5e-4, rtol=5e-4)) + self.assertTrue(mx.allclose(ref, out, atol=1e-4, rtol=1e-4)) def test_sdpa_vector(self): D = 64 @@ -365,7 +365,7 @@ def test_sdpa_vector(self): scale=scale, mask=m, ) - self.assertTrue(mx.allclose(ref, out, atol=5e-4, rtol=5e-4)) + self.assertTrue(mx.allclose(ref, out, atol=1e-4, rtol=1e-4)) L = 4096 scale = 1.0 @@ -394,7 +394,7 @@ def test_sdpa_vector(self): scale=scale, mask=m, ) - self.assertTrue(mx.allclose(ref, out, atol=5e-4, rtol=5e-4)) + self.assertTrue(mx.allclose(ref, out, atol=1e-4, rtol=1e-4)) def test_sdpa_vector_gqa_long(self): scale = 1.0 @@ -408,7 +408,7 @@ def test_sdpa_vector_gqa_long(self): vr = mx.repeat(v, Nq // Nkv, axis=1) ref = mlx_primitives_sdpa(q, kr, vr, scale) out = mx.fast.scaled_dot_product_attention(q, k, v, scale=scale) - self.assertTrue(mx.allclose(ref, out, atol=5e-4, rtol=5e-4)) + self.assertTrue(mx.allclose(ref, out, atol=1e-4, rtol=1e-4)) @unittest.skipIf(not mx.metal.is_available(), "Metal kernel path only") def test_sdpa_vector_head_dim_512(self): @@ -431,7 +431,7 @@ def test_sdpa_vector_head_dim_512(self): out = mx.fast.scaled_dot_product_attention( q, k, v, scale=scale, force_fused=True ) - atol = 7e-5 if dtype == mx.float32 else 2e-2 + atol = 1e-5 if dtype == mx.float32 else 2e-2 self.assertTrue(mx.allclose(ref, out, atol=atol)) # Test 2-pass kernel. @@ -446,7 +446,7 @@ def test_sdpa_vector_head_dim_512(self): out = mx.fast.scaled_dot_product_attention( q, k, v, scale=scale, force_fused=True ) - atol = 7e-5 if dtype == mx.float32 else 2e-2 + atol = 1e-5 if dtype == mx.float32 else 2e-2 self.assertTrue(mx.allclose(ref, out, atol=atol)) # Test other heads. @@ -463,7 +463,7 @@ def test_sdpa_vector_head_dim_512(self): scale=scale, force_fused=True, ) - self.assertTrue(mx.allclose(ref, out, atol=5e-4, rtol=5e-4)) + self.assertTrue(mx.allclose(ref, out, atol=1e-4, rtol=1e-4)) # Test batched. B = 2 @@ -483,7 +483,7 @@ def test_sdpa_vector_head_dim_512(self): sinks=s_in, force_fused=True, ) - self.assertTrue(mx.allclose(ref, out, atol=5e-4, rtol=5e-4)) + self.assertTrue(mx.allclose(ref, out, atol=1e-4, rtol=1e-4)) def test_sdpa_fully_masked(self): Lkv = 8 @@ -507,7 +507,7 @@ def test_sdpa_inf_score(self): k[..., 0, :] = -float("inf") ref = mlx_primitives_sdpa(q, k, v, scale=1, mask=None) out = mx.fast.scaled_dot_product_attention(q, k, v, mask=None, scale=1) - self.assertTrue(mx.allclose(ref, out, atol=5e-4, rtol=5e-4)) + self.assertTrue(mx.allclose(ref, out, atol=1e-4, rtol=1e-4)) def test_sdpa_few_query(self): D = 64 @@ -539,7 +539,7 @@ def test_sdpa_few_query(self): scale=scale, mask=m, ) - self.assertTrue(mx.allclose(ref, out, atol=5e-4, rtol=5e-4)) + self.assertTrue(mx.allclose(ref, out, atol=1e-4, rtol=1e-4)) L = 4096 scale = 1.0 @@ -565,7 +565,7 @@ def test_sdpa_few_query(self): scale=scale, mask=m, ) - self.assertTrue(mx.allclose(ref, out, atol=5e-4, rtol=5e-4)) + self.assertTrue(mx.allclose(ref, out, atol=1e-4, rtol=1e-4)) @unittest.skip("Different head and value dims is not enabled") def test_sdpa_vector_value_dims(self): @@ -582,7 +582,7 @@ def test_sdpa_vector_value_dims(self): v = 5e-1 * mx.random.normal(shape=(1, Nkv, L, V)) ref = mlx_primitives_sdpa(q, k, v, scale) out = mx.fast.scaled_dot_product_attention(q, k, v, scale=scale) - self.assertTrue(mx.allclose(ref, out, atol=5e-4, rtol=5e-4)) + self.assertTrue(mx.allclose(ref, out, atol=1e-4, rtol=1e-4)) def test_sdpa_vector_batched(self): D = 64 @@ -592,29 +592,29 @@ def test_sdpa_vector_batched(self): out = mx.fast.scaled_dot_product_attention(q, k, v, mask=None, scale=1.0) ref = mlx_ref_attn(q, k, v) - self.assertTrue(mx.allclose(ref, out, atol=5e-4, rtol=5e-4)) + self.assertTrue(mx.allclose(ref, out, atol=1e-4, rtol=1e-4)) q = mx.random.normal(shape=(2, 4, 3, D)) out = mx.fast.scaled_dot_product_attention(q, k, v, mask=None, scale=1.0) ref = mlx_ref_attn(q, k, v) - self.assertTrue(mx.allclose(ref, out, atol=5e-4, rtol=5e-4)) + self.assertTrue(mx.allclose(ref, out, atol=1e-4, rtol=1e-4)) q = mx.random.normal(shape=(2, 3, 4, D)).swapaxes(1, 2) out = mx.fast.scaled_dot_product_attention(q, k, v, mask=None, scale=1.0) ref = mlx_ref_attn(q, k, v) - self.assertTrue(mx.allclose(ref, out, atol=5e-4, rtol=5e-4)) + self.assertTrue(mx.allclose(ref, out, atol=1e-4, rtol=1e-4)) k = mx.random.normal(shape=(2, 3, 1, D)).swapaxes(1, 2) out = mx.fast.scaled_dot_product_attention(q, k, v, mask=None, scale=1.0) ref = mlx_ref_attn(q, k, v) - self.assertTrue(mx.allclose(ref, out, atol=5e-4, rtol=5e-4)) + self.assertTrue(mx.allclose(ref, out, atol=1e-4, rtol=1e-4)) q = mx.random.normal(shape=(2, 4, 3, D)) k = mx.random.normal(shape=(2, 3, 2, D)).swapaxes(1, 2) v = mx.random.normal(shape=(2, 2, 3, D)) out = mx.fast.scaled_dot_product_attention(q, k, v, mask=None, scale=1.0) ref = mlx_ref_attn(q, k, v) - self.assertTrue(mx.allclose(ref, out, atol=5e-4, rtol=5e-4)) + self.assertTrue(mx.allclose(ref, out, atol=1e-4, rtol=1e-4)) q = mx.random.normal(shape=(2, 4, 3, D)) k = mx.random.normal(shape=(2, 1, 3, D)) @@ -622,7 +622,7 @@ def test_sdpa_vector_batched(self): mask = 10 * mx.random.normal(shape=(1, 2, 3, 3)).swapaxes(0, 1) out = mx.fast.scaled_dot_product_attention(q, k, v, mask=mask, scale=1.0) ref = mlx_ref_attn(q, k, v, mask=mask) - self.assertTrue(mx.allclose(ref, out, atol=5e-4, rtol=5e-4)) + self.assertTrue(mx.allclose(ref, out, atol=1e-4, rtol=1e-4)) @unittest.skipIf(not mx.is_available(mx.gpu), "GPU kernel path only") def test_sdpa_blocks_env_override(self): @@ -637,7 +637,7 @@ def test_sdpa_blocks_env_override(self): for blocks in (16, 33, 48, 100): with mlx_tests.scoped_env(MLX_SDPA_BLOCKS=str(blocks)): out = mx.fast.scaled_dot_product_attention(q, k, v, scale=D**-0.5) - self.assertTrue(mx.allclose(ref, out, atol=5e-4, rtol=5e-4)) + self.assertTrue(mx.allclose(ref, out, atol=1e-4, rtol=1e-4)) @unittest.skipIf(not mx.is_available(mx.gpu), "too slow on CPU") def test_sdpa(self): @@ -715,7 +715,7 @@ def test_sdpa(self): out_ref = out_ref[:, :, offset:, :] out_fst = out_fst[:, :, offset:, :] - atol = 6e-5 if dtype == mx.float32 else 3e-4 + atol = 2e-5 if dtype == mx.float32 else 3e-4 self.assertListEqual(list(out_ref.shape), list(out_fst.shape)) @@ -775,7 +775,7 @@ def test_sdpa_broadcast_mask(self): v = 5e-1 * mx.random.normal(shape=(1, Nkv, L, D)) ref = mlx_primitives_sdpa(q, k, v, scale, mask=mask) out = mx.fast.scaled_dot_product_attention(q, k, v, scale=scale, mask=mask) - self.assertTrue(mx.allclose(ref, out, atol=5e-4, rtol=5e-4)) + self.assertTrue(mx.allclose(ref, out, atol=1e-4, rtol=1e-4)) def test_sdpa_noncontiguous_inputs(self): mask = mx.ones(shape=(4, 1, 7, 7), dtype=mx.bool_) @@ -786,7 +786,7 @@ def test_sdpa_noncontiguous_inputs(self): v = mx.random.normal(shape=(4, 7, 8, 64)).swapaxes(1, 2) out = mx.fast.scaled_dot_product_attention(q, k, v, scale=1.0, mask=mask) ref = mlx_ref_attn(q, k, v, scale=1.0, mask=mask) - self.assertTrue(mx.allclose(ref, out, atol=5e-4, rtol=5e-4)) + self.assertTrue(mx.allclose(ref, out, atol=1e-4, rtol=1e-4)) def test_sdpa_promote_mask(self): mask = mx.array(2.0, mx.bfloat16) @@ -802,7 +802,7 @@ def test_sdpa_promote_mask(self): v = 5e-1 * mx.random.normal(shape=(1, Nkv, L, D)) ref = mlx_primitives_sdpa(q, k, v, scale, mask=mask) out = mx.fast.scaled_dot_product_attention(q, k, v, scale=scale, mask=mask) - self.assertTrue(mx.allclose(ref, out, atol=5e-4, rtol=5e-4)) + self.assertTrue(mx.allclose(ref, out, atol=1e-4, rtol=1e-4)) def test_sdpa_nan_bug(self): N = 128 @@ -822,7 +822,7 @@ def test_sdpa_nan_bug(self): out = mx.fast.scaled_dot_product_attention(q, k, v, mask=mask, scale=1.0) expected = mlx_ref_attn(q, k, v, mask=mask, scale=1.0) self.assertFalse(mx.isnan(out).any().item()) - self.assertLessEqual(mx.abs(out - expected).max().item(), 5e-4) + self.assertLessEqual(mx.abs(out - expected).max().item(), 1e-4) # And an additive one mask = mx.log(mask) @@ -830,7 +830,7 @@ def test_sdpa_nan_bug(self): out = mx.fast.scaled_dot_product_attention(q, k, v, mask=mask, scale=1.0) expected = mlx_ref_attn(q, k, v, mask=mask, scale=1.0) self.assertFalse(mx.isnan(out).any().item()) - self.assertLessEqual(mx.abs(out - expected).max().item(), 5e-4) + self.assertLessEqual(mx.abs(out - expected).max().item(), 1e-4) def test_sdpa_attention_sinks(self): B = 2 @@ -879,7 +879,7 @@ def test_sdpa_attention_sinks(self): out = mx.fast.scaled_dot_product_attention( q, k, v, scale=scale, sinks=sinks ) - atol = 6e-4 if dtype == mx.float32 else 1e-2 + atol = 1e-5 if dtype == mx.float32 else 1e-2 self.assertTrue(mx.allclose(out, expected, atol=atol)) def test_sdpa_grad(self): From c5770e6566a4f979677254e3058878920494a087 Mon Sep 17 00:00:00 2001 From: mlx-dev Date: Mon, 28 Sep 2026 01:19:45 -0700 Subject: [PATCH 3/7] fix dim 512 --- .../metal/kernels/steel/attn/kernels/steel_attention_nax.h | 2 +- mlx/backend/metal/scaled_dot_product_attention.cpp | 3 +-- 2 files changed, 2 insertions(+), 3 deletions(-) diff --git a/mlx/backend/metal/kernels/steel/attn/kernels/steel_attention_nax.h b/mlx/backend/metal/kernels/steel/attn/kernels/steel_attention_nax.h index 1ef1c3b4cb..13cf24f78a 100644 --- a/mlx/backend/metal/kernels/steel/attn/kernels/steel_attention_nax.h +++ b/mlx/backend/metal/kernels/steel/attn/kernels/steel_attention_nax.h @@ -923,7 +923,7 @@ template < rcp[i] = 1.f / sum_score[i]; } - if (save_lse && d_half == 0 && sn == 0) { + if (save_lse && d_group == 0 && sn == 0) { const int lse_row_base = int(tid.x) * BQ + tm; const int lse_head_off = (int(tid.z) * params->H + int(tid.y)) * params->qL; diff --git a/mlx/backend/metal/scaled_dot_product_attention.cpp b/mlx/backend/metal/scaled_dot_product_attention.cpp index 590f69f297..0b58333dab 100644 --- a/mlx/backend/metal/scaled_dot_product_attention.cpp +++ b/mlx/backend/metal/scaled_dot_product_attention.cpp @@ -35,8 +35,7 @@ void sdpa_full_self_attention_nax( int bq = bd == 512 ? 32 : 64; int bk = 32; - - bool split_d = (bd == 256 || bd == 512) && !sinks.has_value() && (lse == nullptr); + bool split_d = (bd == 256 || bd == 512); int wm = bd == 512 ? 2 : 4; int wn = split_d ? bd / 128 : 1; int B = q.shape(0); From a56111f7c8814eca17221956b1e0bea2e1276af5 Mon Sep 17 00:00:00 2001 From: mlx-dev Date: Mon, 28 Sep 2026 01:30:06 -0700 Subject: [PATCH 4/7] add OdO for D=512 --- .../metal/kernels/scaled_dot_product_attention_vjp.metal | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/mlx/backend/metal/kernels/scaled_dot_product_attention_vjp.metal b/mlx/backend/metal/kernels/scaled_dot_product_attention_vjp.metal index aa5d08a058..3a804d94f7 100644 --- a/mlx/backend/metal/kernels/scaled_dot_product_attention_vjp.metal +++ b/mlx/backend/metal/kernels/scaled_dot_product_attention_vjp.metal @@ -132,7 +132,8 @@ template instantiate_odo(in_type, 96); \ instantiate_odo(in_type, 128); \ instantiate_odo(in_type, 192); \ - instantiate_odo(in_type, 256); + instantiate_odo(in_type, 256); \ + instantiate_odo(in_type, 512); \ instantiate_odo_shapes(bfloat16_t); instantiate_odo_shapes(float16_t); From be9ef2fd8e118ede141236ed5a8a76f7a84d72f4 Mon Sep 17 00:00:00 2001 From: mlx-dev Date: Mon, 28 Sep 2026 06:30:21 -0700 Subject: [PATCH 5/7] remove a not needed buffer --- .../kernels/scaled_dot_product_attention_vjp.metal | 6 ++---- mlx/backend/metal/scaled_dot_product_attention.cpp | 10 ++++------ 2 files changed, 6 insertions(+), 10 deletions(-) diff --git a/mlx/backend/metal/kernels/scaled_dot_product_attention_vjp.metal b/mlx/backend/metal/kernels/scaled_dot_product_attention_vjp.metal index 3a804d94f7..da196b7d53 100644 --- a/mlx/backend/metal/kernels/scaled_dot_product_attention_vjp.metal +++ b/mlx/backend/metal/kernels/scaled_dot_product_attention_vjp.metal @@ -52,9 +52,8 @@ template const device float* lse [[buffer(2)]], // [B, H, qL] const device float* odo [[buffer(3)]], // [B, H, qL] device T* dS [[buffer(4)]], // [BH, bq, bk] - device T* dSt [[buffer(5)]], // [BH, bk, bq] - device T* Pt [[buffer(6)]], // [BH, bk, bq] - const constant SDPAVJPTileParams& p [[buffer(7)]], + device T* Pt [[buffer(5)]], // [BH, bk, bq] + const constant SDPAVJPTileParams& p [[buffer(6)]], uint3 gid [[thread_position_in_grid]]) { int col = int(gid.x); int row = int(gid.y); @@ -82,7 +81,6 @@ template float dsv = pv * (static_cast(dP[sbase + col]) - dlt) * p.scale; dS[sbase + col] = static_cast(dsv); - dSt[tbase + row] = static_cast(dsv); Pt[tbase + row] = static_cast(pv); } diff --git a/mlx/backend/metal/scaled_dot_product_attention.cpp b/mlx/backend/metal/scaled_dot_product_attention.cpp index 0b58333dab..1ca9cbb948 100644 --- a/mlx/backend/metal/scaled_dot_product_attention.cpp +++ b/mlx/backend/metal/scaled_dot_product_attention.cpp @@ -1094,7 +1094,6 @@ void sdpa_vjp_blocked( v_bs); array ds_v = vjp_view(s_buf, {BH, bq_len, bk_len}); - array dst_v = vjp_view(dst_buf, {BH, bk_len, bq_len}); array pt_v = vjp_view(pt_buf, {BH, bk_len, bq_len}); SDPAVJPTileParams params{ @@ -1107,9 +1106,8 @@ void sdpa_vjp_blocked( compute_encoder.set_input_array(lse, 2); compute_encoder.set_input_array(odo, 3); compute_encoder.set_output_array(s_v, 4); - compute_encoder.set_output_array(dst_v, 5); - compute_encoder.set_output_array(pt_v, 6); - compute_encoder.set_bytes(params, 7); + compute_encoder.set_output_array(pt_v, 5); + compute_encoder.set_bytes(params, 6); compute_encoder.dispatch_threads( MTL::Size(bk_len, bq_len, BH), MTL::Size(32, 8, 1)); @@ -1142,7 +1140,7 @@ void sdpa_vjp_blocked( steel_matmul( s, d, - dst_v, + ds_v, q_sl, dk_v, bk_len, @@ -1151,7 +1149,7 @@ void sdpa_vjp_blocked( BH, bq_len, D, - false, + true, false, copies, bshape, From 79bb531b230221b66e79bc7f59a9395dcfeb1793 Mon Sep 17 00:00:00 2001 From: mlx-dev Date: Mon, 28 Sep 2026 06:31:56 -0700 Subject: [PATCH 6/7] implement Cheng suggestion --- mlx/backend/metal/scaled_dot_product_attention.cpp | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/mlx/backend/metal/scaled_dot_product_attention.cpp b/mlx/backend/metal/scaled_dot_product_attention.cpp index 1ca9cbb948..22ad94a7b2 100644 --- a/mlx/backend/metal/scaled_dot_product_attention.cpp +++ b/mlx/backend/metal/scaled_dot_product_attention.cpp @@ -968,9 +968,7 @@ void sdpa_vjp_blocked( // nonzero whenever the queries are a suffix of the keys. const int diag_off = kL - qL; - auto blocks = sdpa_vjp_blocks(B, H, qL, kL, q.dtype()); - const int BQ = blocks.first; - const int BK = blocks.second; + auto [BQ, BK] = sdpa_vjp_blocks(B, H, qL, kL, q.dtype()); Dtype ctype = q.dtype(); std::string tname = get_type_string(ctype); From 3a2abe7de655ca56e97772e50c1c48b9c9b6a243 Mon Sep 17 00:00:00 2001 From: mlx-dev Date: Mon, 28 Sep 2026 10:40:05 -0700 Subject: [PATCH 7/7] Remove one buffer --- .../scaled_dot_product_attention_vjp.metal | 5 ++--- .../metal/scaled_dot_product_attention.cpp | 15 +++++++-------- 2 files changed, 9 insertions(+), 11 deletions(-) diff --git a/mlx/backend/metal/kernels/scaled_dot_product_attention_vjp.metal b/mlx/backend/metal/kernels/scaled_dot_product_attention_vjp.metal index da196b7d53..ddd9f4bcda 100644 --- a/mlx/backend/metal/kernels/scaled_dot_product_attention_vjp.metal +++ b/mlx/backend/metal/kernels/scaled_dot_product_attention_vjp.metal @@ -52,7 +52,7 @@ template const device float* lse [[buffer(2)]], // [B, H, qL] const device float* odo [[buffer(3)]], // [B, H, qL] device T* dS [[buffer(4)]], // [BH, bq, bk] - device T* Pt [[buffer(5)]], // [BH, bk, bq] + device T* P [[buffer(5)]], // [BH, bk, bq] const constant SDPAVJPTileParams& p [[buffer(6)]], uint3 gid [[thread_position_in_grid]]) { int col = int(gid.x); @@ -63,7 +63,6 @@ template } size_t sbase = (size_t(bh) * size_t(p.bq) + size_t(row)) * size_t(p.bk); - size_t tbase = (size_t(bh) * size_t(p.bk) + size_t(col)) * size_t(p.bq); size_t qi = size_t(bh) * size_t(p.qL) + size_t(p.i0 + row); float l = lse[qi]; @@ -81,7 +80,7 @@ template float dsv = pv * (static_cast(dP[sbase + col]) - dlt) * p.scale; dS[sbase + col] = static_cast(dsv); - Pt[tbase + row] = static_cast(pv); + P[sbase + col] = static_cast(pv); } template diff --git a/mlx/backend/metal/scaled_dot_product_attention.cpp b/mlx/backend/metal/scaled_dot_product_attention.cpp index 22ad94a7b2..2d81c35d8d 100644 --- a/mlx/backend/metal/scaled_dot_product_attention.cpp +++ b/mlx/backend/metal/scaled_dot_product_attention.cpp @@ -980,8 +980,7 @@ void sdpa_vjp_blocked( // Temporary tiles array s_buf = vjp_alloc({BH, BQ, BK}, ctype, s); array dp_buf = vjp_alloc({BH, BQ, BK}, ctype, s); - array dst_buf = vjp_alloc({BH, BK, BQ}, ctype, s); - array pt_buf = vjp_alloc({BH, BK, BQ}, ctype, s); + array p_buf = vjp_alloc({BH, BQ, BK}, ctype, s); // Output of the gemms array dq_tile = vjp_alloc({BH, BQ, D}, ctype, s); @@ -1092,7 +1091,7 @@ void sdpa_vjp_blocked( v_bs); array ds_v = vjp_view(s_buf, {BH, bq_len, bk_len}); - array pt_v = vjp_view(pt_buf, {BH, bk_len, bq_len}); + array p_v = vjp_view(p_buf, {BH, bq_len, bk_len}); SDPAVJPTileParams params{ bq_len, bk_len, qL, i0, j0, scale, diag_off, causal ? 1 : 0}; @@ -1104,7 +1103,7 @@ void sdpa_vjp_blocked( compute_encoder.set_input_array(lse, 2); compute_encoder.set_input_array(odo, 3); compute_encoder.set_output_array(s_v, 4); - compute_encoder.set_output_array(pt_v, 5); + compute_encoder.set_output_array(p_v, 5); compute_encoder.set_bytes(params, 6); compute_encoder.dispatch_threads( MTL::Size(bk_len, bq_len, BH), MTL::Size(32, 8, 1)); @@ -1145,7 +1144,7 @@ void sdpa_vjp_blocked( D, bq_len, BH, - bq_len, + bk_len, D, true, false, @@ -1159,16 +1158,16 @@ void sdpa_vjp_blocked( steel_matmul( s, d, - pt_v, + p_v, o_sl, dv_v, bk_len, Dv, bq_len, BH, - bq_len, + bk_len, Dv, - false, + true, false, copies, bshape,