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 d5513dd030..4da27a5843 100644 --- a/mlx/backend/metal/kernels/CMakeLists.txt +++ b/mlx/backend/metal/kernels/CMakeLists.txt @@ -56,6 +56,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..ddd9f4bcda --- /dev/null +++ b/mlx/backend/metal/kernels/scaled_dot_product_attention_vjp.metal @@ -0,0 +1,154 @@ +// 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* P [[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); + 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 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); + P[sbase + col] = 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(in_type, 512); \ + +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 d5cb00e427..d8f5efdc11 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 @@ -77,6 +78,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]], @@ -103,6 +105,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 @@ -532,6 +540,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] + d_half * BVh + sn; @@ -545,4 +567,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 7a7099b8f4..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 @@ -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 { @@ -90,6 +91,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]], @@ -117,6 +119,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 @@ -475,6 +484,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]); @@ -516,6 +535,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]], @@ -903,6 +923,21 @@ template < rcp[i] = 1.f / sum_score[i]; } + 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; + + 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 2296187fea..eba254155d 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 @@ -22,11 +22,13 @@ instantiate_attn_dsplit(iname, itype, 64, 32, 256, 4, 2, mname, mtype) \ instantiate_attn(iname, itype, 64, 64, 128, 128, 4, 1, mname, mtype) \ instantiate_attn(iname, itype, 64, 64, 64, 64, 4, 1, mname, mtype) \ + instantiate_attn(iname, itype, 64, 32, 256, 256, 4, 1, mname, mtype) \ instantiate_attn(iname, itype, 64, 32, 128, 128, 4, 1, mname, mtype) \ instantiate_attn(iname, itype, 64, 32, 96, 96, 4, 1, mname, mtype) \ instantiate_attn(iname, itype, 64, 32, 96, 64, 4, 1, mname, mtype) \ instantiate_attn(iname, itype, 64, 32, 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 3f5ed83368..2d81c35d8d 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,7 +26,8 @@ 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); @@ -33,7 +35,7 @@ void sdpa_full_self_attention_nax( int bq = bd == 512 ? 32 : 64; int bk = 32; - bool split_d = bd == 256 || bd == 512; + 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); @@ -79,13 +81,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( @@ -120,7 +124,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); @@ -191,6 +197,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); @@ -209,15 +218,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 || D == 512) && (env::enable_tf32() || q.dtype() != float32)) { @@ -231,15 +237,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; @@ -286,7 +295,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( @@ -306,6 +316,9 @@ 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); char devc = d.get_architecture().back(); int bd = q.shape(-1); int bv = v.shape(-1); @@ -322,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( @@ -363,7 +378,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); @@ -433,6 +450,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); @@ -726,11 +746,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); @@ -832,6 +853,341 @@ 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 [BQ, BK] = sdpa_vjp_blocks(B, H, qL, kL, q.dtype()); + + 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 p_buf = vjp_alloc({BH, BQ, BK}, 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 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}; + + // 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(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)); + + 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, + ds_v, + q_sl, + dk_v, + bk_len, + D, + bq_len, + BH, + bk_len, + D, + true, + 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, + p_v, + o_sl, + dv_v, + bk_len, + Dv, + bq_len, + BH, + bk_len, + Dv, + true, + 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( @@ -858,15 +1214,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); @@ -922,7 +1278,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()); @@ -950,6 +1305,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; @@ -1048,21 +1409,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