From eb1b13ede85ec54899e54e6f9c73506ec18e0dc1 Mon Sep 17 00:00:00 2001 From: mlx-dev Date: Fri, 25 Sep 2026 05:15:51 -0700 Subject: [PATCH 01/11] Add GDN VJP - split from gdn-spda-vjp branch --- mlx/backend/metal/CMakeLists.txt | 1 + mlx/backend/metal/gated_delta_update.cpp | 479 ++++- mlx/backend/metal/jit_kernels.cpp | 38 + mlx/backend/metal/kernels.h | 24 + mlx/backend/metal/kernels/CMakeLists.txt | 3 + .../metal/kernels/gated_delta_nax_ops.h | 622 ++++++ .../metal/kernels/gated_delta_update.h | 27 +- .../metal/kernels/gated_delta_update.metal | 48 +- .../metal/kernels/gated_delta_update_nax.h | 389 +--- .../kernels/gated_delta_update_nax.metal | 31 +- .../kernels/gated_delta_update_nax_vjp.h | 1736 +++++++++++++++++ .../kernels/gated_delta_update_nax_vjp.metal | 44 + .../metal/kernels/gated_delta_update_vjp.h | 184 ++ .../kernels/gated_delta_update_vjp.metal | 35 + mlx/backend/metal/nojit_kernels.cpp | 16 + mlx/fast.cpp | 53 +- mlx/fast_primitives.h | 39 + 17 files changed, 3315 insertions(+), 454 deletions(-) create mode 100644 mlx/backend/metal/kernels/gated_delta_nax_ops.h create mode 100644 mlx/backend/metal/kernels/gated_delta_update_nax_vjp.h create mode 100644 mlx/backend/metal/kernels/gated_delta_update_nax_vjp.metal create mode 100644 mlx/backend/metal/kernels/gated_delta_update_vjp.h create mode 100644 mlx/backend/metal/kernels/gated_delta_update_vjp.metal diff --git a/mlx/backend/metal/CMakeLists.txt b/mlx/backend/metal/CMakeLists.txt index c86cfdeb07..28c4072944 100644 --- a/mlx/backend/metal/CMakeLists.txt +++ b/mlx/backend/metal/CMakeLists.txt @@ -103,6 +103,7 @@ if(MLX_METAL_JIT) make_jit_source(steel/attn/kernels/steel_attention_nax) make_jit_source(gated_delta_update_nax) + make_jit_source(gated_delta_update_nax_vjp) else() message( diff --git a/mlx/backend/metal/gated_delta_update.cpp b/mlx/backend/metal/gated_delta_update.cpp index 45f41677c0..ef9eca5fab 100644 --- a/mlx/backend/metal/gated_delta_update.cpp +++ b/mlx/backend/metal/gated_delta_update.cpp @@ -12,38 +12,41 @@ namespace mlx::core::fast { -bool GatedDeltaUpdate::use_fallback( - const int Hk, - const int Dk, - const int Hv, - const int Dv, - const bool has_mask, - Stream s) { - if (s.device == Device::cpu) { - return true; - } +namespace { - if (has_mask) { - return true; - } +inline int gated_delta_chunk_size(int T) { + int C = metal::is_nax_available() ? 16 : 8; + C = T > 8 ? C : 1; + return env::get_var("GATED_DELTA_CHUNK", C); +} - if (Dk != 128 || Dv != 128) { - return true; - } +inline int gated_delta_chunk_size_vjp(int T) { + int C = metal::is_nax_available() ? 16 : 1; + return env::get_var("GATED_DELTA_CHUNK_VJP", C); +} - const bool supported_heads = (Hk == 24 && Hv == 24) || - (Hk == 32 && Hv == 32) || (Hk == 16 && Hv == 32) || - (Hk == 16 && Hv == 48) || (Hk == 16 && Hv == 16) || - (Hk == 16 && Hv == 64); - if (!supported_heads) { - return true; +inline int gated_delta_ckpt(int chunk) { + int c = env::get_var("GATED_DELTA_CKPT", 16); + if (c != 1 && c != 4 && c != 8 && c != 16) { + c = 16; } + return (chunk == 16) ? c : 1; +} - return false; +inline int gated_delta_n_ckpt(int n_chunks, int ckpt) { + return (n_chunks + ckpt - 1) / ckpt; } -inline array -ensure_row_contiguous(const array& x, metal::Device& d, const Stream& s) { +bool supported_gated_delta_shape(int Hk, int Dk, int Hv, int Dv) { + if (Dk != 128 || Dv != 128) { + return false; + } + return (Hk == 24 && Hv == 24) || (Hk == 32 && Hv == 32) || + (Hk == 16 && Hv == 32) || (Hk == 16 && Hv == 48) || + (Hk == 16 && Hv == 16) || (Hk == 16 && Hv == 64); +} + +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); @@ -53,6 +56,31 @@ ensure_row_contiguous(const array& x, metal::Device& d, const Stream& s) { } } +array scratch_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; +} + +} // namespace + +bool GatedDeltaUpdate::use_fallback( + const int Hk, + const int Dk, + const int Hv, + const int Dv, + const bool has_mask, + Stream s) { + if (s.device == Device::cpu) { + return true; + } + if (has_mask) { + return true; + } + return !supported_gated_delta_shape(Hk, Dk, Hv, Dv); +} + void GatedDeltaUpdate::eval_gpu( const std::vector& inputs, std::vector& outputs) { @@ -76,92 +104,379 @@ void GatedDeltaUpdate::eval_gpu( int Hv = v.shape(2); int Dv = v.shape(3); - int C = 1; - int threshold = env::get_var("GATED_DELTA_THRESH", 16); - if (T >= threshold) - C = metal::is_nax_available() ? 16 : 8; - - C = env::get_var("GATED_DELTA_CHUNK", C); - - if (!metal::is_nax_available()) - C = std::min(C, 8); // override in case nax is not available. - - std::string suffix; - concatenate( - suffix, - get_type_string(q.dtype()), - "_", - std::to_string(Dk), - "_", - std::to_string(Dv), - "_", - std::to_string(Hk), - "_", - std::to_string(Hv)); + int C = gated_delta_chunk_size(T); + const int ckpt = gated_delta_ckpt(C); + + std::string suffix = get_type_string(q.dtype()) + "_" + std::to_string(Dk) + + "_" + std::to_string(Dv) + "_" + std::to_string(Hk) + "_" + + std::to_string(Hv); auto& compute_encoder = metal::get_command_encoder(s); out.set_data(allocator::malloc(out.nbytes())); hf.set_data(allocator::malloc(hf.nbytes())); + fill_gpu(array(0, out.dtype()), out, s); + + // The forward is inference-only here: the vjp runs its own save pass. The + // kernel declares buffers 9, 10 and 11 unconditionally, so bind a placeholder + // and turn the stores off with the function constant. + bool save_state = false; + metal::MTLFCList func_consts = { + {&save_state, MTL::DataType::DataTypeBool, 200}, + }; + + array dummy = scratch_alloc({1}, float32, s); + switch (C) { case 16: { - std::string kernel_name = "gated_delta_fused_nax_"; - std::string base_name; - concatenate(base_name, kernel_name, suffix, "_16"); - std::string hash_name = base_name; - metal::MTLFCList func_consts = {}; + std::string base_name = "gated_delta_fused_nax_" + suffix + "_" + + std::to_string(C) + "_" + std::to_string(ckpt); auto delta_kernel = - get_gated_delta_nax_kernel(d, base_name, hash_name, func_consts); + get_gated_delta_nax_kernel(d, base_name, base_name, func_consts); + compute_encoder.set_compute_pipeline_state(delta_kernel); + compute_encoder.set_input_array(q, 0); + compute_encoder.set_input_array(k, 1); + compute_encoder.set_input_array(v, 2); + compute_encoder.set_input_array(h0, 3); // initial state in + compute_encoder.set_input_array(g, 4); + compute_encoder.set_input_array(beta, 5); + compute_encoder.set_output_array(out, 6); + compute_encoder.set_output_array(hf, 7); // final state out + compute_encoder.set_bytes(T, 8); + compute_encoder.set_output_array(dummy, 9); // state_cache + compute_encoder.set_output_array(dummy, 10); // chunk_mats + compute_encoder.set_output_array(dummy, 11); // chunk_delta + + auto grid = MTL::Size(32, Dv / 16, B * Hv); + auto threads = MTL::Size(32, 4, 1); + compute_encoder.dispatch_threads(grid, threads); break; } case 8: { - std::string kernel_name = "gated_delta_fused_chunk_"; - std::string base_name; - concatenate(base_name, kernel_name, suffix, "_8"); - std::string hash_name = base_name; - metal::MTLFCList func_consts = {}; + std::string base_name = + "gated_delta_fused_chunk_" + suffix + "_" + std::to_string(C); auto delta_kernel = - get_gated_delta_kernel(d, base_name, hash_name, func_consts); + get_gated_delta_kernel(d, base_name, base_name, func_consts); + compute_encoder.set_compute_pipeline_state(delta_kernel); + compute_encoder.set_input_array(q, 0); + compute_encoder.set_input_array(k, 1); + compute_encoder.set_input_array(v, 2); + compute_encoder.set_input_array(h0, 3); // initial state in + compute_encoder.set_input_array(g, 4); + compute_encoder.set_input_array(beta, 5); + compute_encoder.set_output_array(out, 6); + compute_encoder.set_output_array(hf, 7); // final state out + compute_encoder.set_bytes(T, 8); + + auto grid = MTL::Size(32, Dv / 8, B * Hv); + auto threads = MTL::Size(32, 4, 1); + compute_encoder.dispatch_threads(grid, threads); break; } - case 0: - C = 1; // avoiding a div by zero on the grid. - case 1: { - std::string kernel_name = "seq_gated_delta_"; - std::string base_name; - concatenate(base_name, kernel_name, suffix); - std::string hash_name = base_name; - metal::MTLFCList func_consts = {}; + case 1: + case 0: { + // Ckpt is a template parameter now, so even the inference path has to + // name one; it just never stores. + std::string base_name = + "seq_gated_delta_" + suffix + "_" + std::to_string(ckpt); auto delta_kernel = - get_gated_delta_kernel(d, base_name, hash_name, func_consts); + get_gated_delta_kernel(d, base_name, base_name, func_consts); compute_encoder.set_compute_pipeline_state(delta_kernel); + + // Order must match the kernel signature: state_in is slot 3, g is 4. + compute_encoder.set_input_array(q, 0); + compute_encoder.set_input_array(k, 1); + compute_encoder.set_input_array(v, 2); + compute_encoder.set_input_array(h0, 3); + compute_encoder.set_input_array(g, 4); + compute_encoder.set_input_array(beta, 5); + compute_encoder.set_output_array(out, 6); + compute_encoder.set_output_array(hf, 7); + compute_encoder.set_bytes(T, 8); + // Slot 9 is gated off by the function constant, so it stays unbound. + + auto grid = MTL::Size(32, Dv, B * Hv); + auto threads = MTL::Size(32, 4, 1); + compute_encoder.dispatch_threads(grid, threads); break; } default: { throw std::runtime_error( - "NYI: Only sequential (C=1) and chunk size 8,16 are supported"); + "NYI: Only sequential and chunk size 8,16 are supported"); } } - compute_encoder.set_input_array(q, 0); - compute_encoder.set_input_array(k, 1); - compute_encoder.set_input_array(v, 2); - compute_encoder.set_input_array(h0, 3); // initial state in - compute_encoder.set_input_array(g, 4); - compute_encoder.set_input_array(beta, 5); - compute_encoder.set_output_array(out, 6); - compute_encoder.set_output_array(hf, 7); // final state out - compute_encoder.set_bytes(T, 8); - - auto grid = MTL::Size(32, Dv / C, B * Hv); - auto threads = MTL::Size(32, 4, 1); - compute_encoder.dispatch_threads(grid, threads); +} + +///////////////////////// +// BACKWARD PASS STUFF // +///////////////////////// +bool GatedDeltaUpdateVJP::use_fallback( + const int Hk, + const int Dk, + const int Hv, + const int Dv, + Stream s) { + if (s.device == Device::cpu) { + return true; + } + if (env::get_var("GATED_DELTA_VJP_FALLBACK", 0) != 0) { + return true; + } + return !supported_gated_delta_shape(Hk, Dk, Hv, Dv); +} + +void GatedDeltaUpdateVJP::eval_gpu( + const std::vector& inputs, + std::vector& outputs) { + // Inputs are q, k, v, g, b, h0, cot_o, cot_h + // Outputs are dq, dk, dv, dg, db, dh0 + auto& s = stream(); + auto& d = metal::device(s.device); + + 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); + auto g = ensure_row_contiguous(inputs[3], d, s); + auto beta = ensure_row_contiguous(inputs[4], d, s); + auto h0 = ensure_row_contiguous(inputs[5], d, s); + auto cot_o = ensure_row_contiguous(inputs[6], d, s); + auto cot_h = ensure_row_contiguous(inputs[7], d, s); + + int B = q.shape(0); + int T = q.shape(1); + int Hk = q.shape(2); + int Dk = q.shape(3); + int Hv = v.shape(2); + int Dv = v.shape(3); + + // 16 = chunked NAX, anything else = sequential. + int C = gated_delta_chunk_size_vjp(T); + const bool chunked = (C == 16); + + const int n_chunks = (T + 15) / 16; + const int ckpt = gated_delta_ckpt(C); + const int n_states = chunked ? gated_delta_n_ckpt(n_chunks, ckpt) : T; + + auto& dq = outputs[0]; + auto& dk = outputs[1]; + auto& dv = outputs[2]; + auto& dg = outputs[3]; + auto& db = outputs[4]; + auto& dh = outputs[5]; + + 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())); + dg.set_data(allocator::malloc(dg.nbytes())); + db.set_data(allocator::malloc(db.nbytes())); + dh.set_data(allocator::malloc(dh.nbytes())); + + // The kernels accumulate their gradients as fp32, so we may need a cast. + const bool stage_fp32 = (q.dtype() != float32); + + auto dq_acc = stage_fp32 ? scratch_alloc(dq.shape(), float32, s) : dq; + auto dk_acc = stage_fp32 ? scratch_alloc(dk.shape(), float32, s) : dk; + auto dv_acc = stage_fp32 ? scratch_alloc(dv.shape(), float32, s) : dv; + auto dg_acc = stage_fp32 ? scratch_alloc(dg.shape(), float32, s) : dg; + auto db_acc = stage_fp32 ? scratch_alloc(db.shape(), float32, s) : db; + + fill_gpu(array(0, float32), dq_acc, s); + fill_gpu(array(0, float32), dk_acc, s); + fill_gpu(array(0, float32), dv_acc, s); + fill_gpu(array(0, float32), dg_acc, s); + fill_gpu(array(0, float32), db_acc, s); + fill_gpu(array(0, dh.dtype()), dh, s); + + array state_cache = scratch_alloc({B, Hv, n_states, Dv, Dk}, float32, s); + array seg_states = scratch_alloc( + chunked ? Shape({B, Hv, ckpt, Dv, Dk}) : Shape({1}), float32, s); + array chunk_mats = scratch_alloc( + chunked ? Shape({B, Hv, n_chunks, 3, 16, 16}) : Shape({1}), float32, s); + array chunk_delta = scratch_alloc( + chunked ? Shape({B, Hv, n_chunks, 16, Dv}) : Shape({1}), float32, s); + + if (chunked) { + fill_gpu(array(0, float32), state_cache, s); + fill_gpu(array(0, float32), chunk_mats, s); + fill_gpu(array(0, float32), chunk_delta, s); + } + + array y_scratch = scratch_alloc({B, T, Hv, Dv}, q.dtype(), s); + array hf_scratch = scratch_alloc({B, Hv, Dv, Dk}, float32, s); + + std::string suffix = get_type_string(q.dtype()) + "_" + std::to_string(Dk) + + "_" + std::to_string(Dv) + "_" + std::to_string(Hk) + "_" + + std::to_string(Hv); + + bool save_state = true; + metal::MTLFCList save_consts = { + {&save_state, MTL::DataType::DataTypeBool, 200}, + }; + metal::MTLFCList no_consts = {}; + + switch (C) { + case 16: { + const std::string ckpt_suffix = + "_" + std::to_string(C) + "_" + std::to_string(ckpt); + + // Forward save pass. + { + std::string base_name = "gated_delta_fused_nax_" + suffix + ckpt_suffix; + + auto delta_kernel = get_gated_delta_nax_kernel( + d, base_name, base_name + "_save", save_consts); + + compute_encoder.set_compute_pipeline_state(delta_kernel); + compute_encoder.set_input_array(q, 0); + compute_encoder.set_input_array(k, 1); + compute_encoder.set_input_array(v, 2); + compute_encoder.set_input_array(h0, 3); + compute_encoder.set_input_array(g, 4); + compute_encoder.set_input_array(beta, 5); + compute_encoder.set_output_array(y_scratch, 6); + compute_encoder.set_output_array(hf_scratch, 7); + compute_encoder.set_bytes(T, 8); + compute_encoder.set_output_array(state_cache, 9); + compute_encoder.set_output_array(chunk_mats, 10); + compute_encoder.set_output_array(chunk_delta, 11); + + auto grid = MTL::Size(32, Dv / 16, B * Hv); + auto threads = MTL::Size(32, 4, 1); + compute_encoder.dispatch_threads(grid, threads); + } + + { + std::string base_name = + "gated_delta_vjp_fused_nax_" + suffix + ckpt_suffix; + + auto delta_kernel = + get_gated_delta_vjp_nax_kernel(d, base_name, base_name, no_consts); + + compute_encoder.set_compute_pipeline_state(delta_kernel); + compute_encoder.set_input_array(q, 0); + compute_encoder.set_input_array(k, 1); + compute_encoder.set_input_array(v, 2); + compute_encoder.set_input_array(g, 3); + compute_encoder.set_input_array(beta, 4); + compute_encoder.set_input_array(cot_o, 5); + compute_encoder.set_input_array(cot_h, 6); + compute_encoder.set_input_array(state_cache, 7); + compute_encoder.set_bytes(T, 8); + compute_encoder.set_output_array(dq_acc, 9); + compute_encoder.set_output_array(dk_acc, 10); + compute_encoder.set_output_array(dv_acc, 11); + compute_encoder.set_output_array(dg_acc, 12); + compute_encoder.set_output_array(db_acc, 13); + compute_encoder.set_output_array(dh, 14); + compute_encoder.set_input_array(chunk_mats, 15); + compute_encoder.set_input_array(chunk_delta, 16); + compute_encoder.set_output_array(seg_states, 17); + + auto grid = MTL::Size(32, Dv / 16, B * Hv); + auto threads = MTL::Size(32, Dv / 16, 1); + compute_encoder.dispatch_threads(grid, threads); + } + + { + int n_total = B * Hv; + + std::string base_name = "gated_delta_dgamma_to_dg_" + + get_type_string(q.dtype()) + "_" + std::to_string(C); + + auto dgamma_kernel = + get_gated_delta_vjp_nax_kernel(d, base_name, base_name, no_consts); + + compute_encoder.set_compute_pipeline_state(dgamma_kernel); + compute_encoder.set_input_array(g, 0); + compute_encoder.set_output_array(dg_acc, 1); + compute_encoder.set_bytes(T, 2); + compute_encoder.set_bytes(Hv, 3); + compute_encoder.set_bytes(n_total, 4); + + auto grid = MTL::Size(n_total, n_chunks, 1); + auto threads = MTL::Size(std::min(n_total, 32), 1, 1); + compute_encoder.dispatch_threads(grid, threads); + } + break; + } + case 1: + case 0: { + const std::string ckpt_suffix = "_" + std::to_string(ckpt); + + { + std::string base_name = "seq_gated_delta_" + suffix + ckpt_suffix; + + auto delta_kernel = get_gated_delta_kernel( + d, base_name, base_name + "_save", save_consts); + + compute_encoder.set_compute_pipeline_state(delta_kernel); + compute_encoder.set_input_array(q, 0); + compute_encoder.set_input_array(k, 1); + compute_encoder.set_input_array(v, 2); + compute_encoder.set_input_array(h0, 3); + compute_encoder.set_input_array(g, 4); + compute_encoder.set_input_array(beta, 5); + compute_encoder.set_output_array(y_scratch, 6); + compute_encoder.set_output_array(hf_scratch, 7); + compute_encoder.set_bytes(T, 8); + compute_encoder.set_output_array(state_cache, 9); + + auto grid = MTL::Size(32, Dv, B * Hv); + auto threads = MTL::Size(32, 4, 1); + compute_encoder.dispatch_threads(grid, threads); + } + + { + std::string base_name = "seq_gated_delta_vjp_" + suffix + ckpt_suffix; + + auto delta_kernel = + get_gated_delta_vjp_kernel(d, base_name, base_name, no_consts); + + compute_encoder.set_compute_pipeline_state(delta_kernel); + compute_encoder.set_input_array(q, 0); + compute_encoder.set_input_array(k, 1); + compute_encoder.set_input_array(v, 2); + compute_encoder.set_input_array(g, 3); + compute_encoder.set_input_array(beta, 4); + compute_encoder.set_input_array(cot_o, 5); + compute_encoder.set_input_array(cot_h, 6); + compute_encoder.set_input_array(state_cache, 7); + compute_encoder.set_bytes(T, 8); + compute_encoder.set_output_array(dq_acc, 9); + compute_encoder.set_output_array(dk_acc, 10); + compute_encoder.set_output_array(dv_acc, 11); + compute_encoder.set_output_array(dg_acc, 12); + compute_encoder.set_output_array(db_acc, 13); + compute_encoder.set_output_array(dh, 14); + + auto grid = MTL::Size(32, Dv, B * Hv); + auto threads = MTL::Size(32, 4, 1); + compute_encoder.dispatch_threads(grid, threads); + } + break; + } + default: { + throw std::runtime_error( + "NYI: Only sequential and chunk size 16 are supported for vjp"); + } + } + + if (stage_fp32) { + copy_gpu(dq_acc, dq, CopyType::General, s); + copy_gpu(dk_acc, dk, CopyType::General, s); + copy_gpu(dv_acc, dv, CopyType::General, s); + copy_gpu(dg_acc, dg, CopyType::General, s); + copy_gpu(db_acc, db, CopyType::General, s); + } } } // namespace mlx::core::fast diff --git a/mlx/backend/metal/jit_kernels.cpp b/mlx/backend/metal/jit_kernels.cpp index 1c052e6314..516c980a59 100644 --- a/mlx/backend/metal/jit_kernels.cpp +++ b/mlx/backend/metal/jit_kernels.cpp @@ -1389,6 +1389,22 @@ MTL::ComputePipelineState* get_steel_attention_nax_kernel( return d.get_kernel(kernel_name, lib, hash_name, func_consts); } +MTL::ComputePipelineState* get_sdpa_vjp_kernel( + metal::Device& d, + const std::string& kernel_name, + const std::string& hash_name, + const metal::MTLFCList& func_consts) { + return d.get_kernel(kernel_name, hash_name, func_consts); +} + +MTL::ComputePipelineState* get_sdpa_vjp_nax_kernel( + metal::Device& d, + const std::string& kernel_name, + const std::string& hash_name, + const metal::MTLFCList& func_consts) { + return d.get_kernel(kernel_name, hash_name, func_consts); +} + MTL::ComputePipelineState* get_gated_delta_kernel( metal::Device& d, const std::string& kernel_name, @@ -1397,6 +1413,14 @@ MTL::ComputePipelineState* get_gated_delta_kernel( return d.get_kernel(kernel_name, hash_name, func_consts); } +MTL::ComputePipelineState* get_gated_delta_vjp_kernel( + metal::Device& d, + const std::string& kernel_name, + const std::string& hash_name, + const metal::MTLFCList& func_consts) { + return d.get_kernel(kernel_name, hash_name, func_consts); +} + MTL::ComputePipelineState* get_gated_delta_nax_kernel( metal::Device& d, const std::string& kernel_name, @@ -1411,4 +1435,18 @@ MTL::ComputePipelineState* get_gated_delta_nax_kernel( return d.get_kernel(kernel_name, lib, hash_name, func_consts); } +MTL::ComputePipelineState* get_gated_delta_vjp_nax_kernel( + metal::Device& d, + const std::string& kernel_name, + const std::string& hash_name, + const metal::MTLFCList& func_consts) { + const auto& lib_name = kernel_name; + auto lib = d.get_library(lib_name, [&]() { + std::string kernel_source; + concatenate(kernel_source, metal::utils(), metal::gated_delta_update_nax()); + return kernel_source; + }); + return d.get_kernel(kernel_name, lib, hash_name, func_consts); +} + } // namespace mlx::core diff --git a/mlx/backend/metal/kernels.h b/mlx/backend/metal/kernels.h index 28ac3ba940..e4d791e4e1 100644 --- a/mlx/backend/metal/kernels.h +++ b/mlx/backend/metal/kernels.h @@ -443,6 +443,30 @@ MTL::ComputePipelineState* get_gated_delta_nax_kernel( const std::string& hash_name, const metal::MTLFCList& func_consts); +MTL::ComputePipelineState* get_gated_delta_kernel( + metal::Device& d, + const std::string& kernel_name, + const std::string& hash_name, + const metal::MTLFCList& func_consts); + +MTL::ComputePipelineState* get_gated_delta_nax_kernel( + metal::Device& d, + const std::string& kernel_name, + const std::string& hash_name, + const metal::MTLFCList& func_consts); + +MTL::ComputePipelineState* get_gated_delta_vjp_kernel( + metal::Device& d, + const std::string& kernel_name, + const std::string& hash_name, + const metal::MTLFCList& func_consts); + +MTL::ComputePipelineState* get_gated_delta_vjp_nax_kernel( + metal::Device& d, + const std::string& kernel_name, + const std::string& hash_name, + const metal::MTLFCList& func_consts); + // Create a GPU kernel template definition for JIT compilation template std::string get_template_definition( diff --git a/mlx/backend/metal/kernels/CMakeLists.txt b/mlx/backend/metal/kernels/CMakeLists.txt index 00f90ac862..bfff5bb9bd 100644 --- a/mlx/backend/metal/kernels/CMakeLists.txt +++ b/mlx/backend/metal/kernels/CMakeLists.txt @@ -56,6 +56,7 @@ build_kernel(rms_norm) build_kernel(rope) build_kernel(scaled_dot_product_attention sdpa_vector.h) build_kernel(gated_delta_update gated_delta_update.h) +build_kernel(gated_delta_update_vjp gated_delta_update_vjp.h) if(MLX_METAL_VERSION GREATER_EQUAL 320) build_kernel(fence) endif() @@ -172,6 +173,8 @@ if(NOT MLX_METAL_JIT) build_kernel(quantized_nax quantized_nax.h ${STEEL_NAX_HEADERS}) build_kernel(gated_delta_update_nax gated_delta_update_nax.h ${STEEL_NAX_HEADERS}) + build_kernel(gated_delta_update_nax_vjp gated_delta_update_nax_vjp.h + ${STEEL_NAX_HEADERS}) build_kernel(fp_quantized_nax fp4.h fp8.h fp_quantized_nax.h ${STEEL_NAX_HEADERS}) diff --git a/mlx/backend/metal/kernels/gated_delta_nax_ops.h b/mlx/backend/metal/kernels/gated_delta_nax_ops.h new file mode 100644 index 0000000000..60e46865fe --- /dev/null +++ b/mlx/backend/metal/kernels/gated_delta_nax_ops.h @@ -0,0 +1,622 @@ +#pragma once + +#include + +#include +#include + +#include "mlx/backend/metal/kernels/steel/gemm/nax.h" + +using namespace metal; +using namespace mpp; +using namespace mpp::tensor_ops; + +typedef mlx::steel::NAXTile _M16x16; +typedef mlx::steel::NAXTile _M16x32; + +// NAX MACROS I can probably do a nice template instead of doing this +// fm = base_fm + (idx >> 2) * 8; // idx>>2 = idx/4 -> 0 for idx 0-3, 1 for +// idx 4-7 fn = base_fn + (idx % 4); // 4 consecutive columns +#define AT_NAX(TILE, IDX) TILE.elems()[IDX] + +// out = a - b +template +METAL_FUNC T operator-(const thread T& a, const thread T& b) { + T out; + STEEL_PRAGMA_UNROLL + for (short f = 0; f < T::kNumFrags; f++) { + out.frag_at(0, f) = a.frag_at(0, f) - b.frag_at(0, f); + } + return out; +} + +template +METAL_FUNC T operator-(const thread T& a) { + T out; + STEEL_PRAGMA_UNROLL + for (short f = 0; f < T::kNumFrags; f++) { + out.frag_at(0, f) = -a.frag_at(0, f); + } + return out; +} + +// out = a + b +template +METAL_FUNC T operator+(const thread T& a, const thread T& b) { + T out; + STEEL_PRAGMA_UNROLL + for (short f = 0; f < T::kNumFrags; f++) { + out.frag_at(0, f) = a.frag_at(0, f) + b.frag_at(0, f); + } + return out; +} + +template +METAL_FUNC T operator+(const thread T& a, float addend) { + T out; + STEEL_PRAGMA_UNROLL + for (short f = 0; f < T::kNumFrags; f++) { + out.frag_at(0, f) = a.frag_at(0, f) + addend; + } + return out; +} + +#define TRIL_NAX(TILE0, TILE1) \ + { \ + STEEL_PRAGMA_UNROLL \ + for (short _i = 0; _i < decltype(TILE0)::kElemsPerFrag; _i++) { \ + const short2 _c = mlx::steel::BaseNAXFrag::get_coord(_i); /* {fn, fm} */ \ + AT_NAX(TILE0, _i) = (_c.x >= _c.y) ? 0.f : AT_NAX(TILE1, _i); \ + } \ + } + +// out = a * b, elementwise. Not a matmul. +template +METAL_FUNC T operator*(const thread T& a, const thread T& b) { + T out; + STEEL_PRAGMA_UNROLL + for (short f = 0; f < T::kNumFrags; f++) { + out.frag_at(0, f) = a.frag_at(0, f) * b.frag_at(0, f); + } + return out; +} + +template +METAL_FUNC T operator*(const thread T& a, float s) { + T out; + STEEL_PRAGMA_UNROLL + for (short f = 0; f < T::kNumFrags; f++) { + out.frag_at(0, f) = a.frag_at(0, f) * s; + } + return out; +} + +template +METAL_FUNC T fmadd(const thread T& a, const thread T& b, const thread T& c) { + T out; + STEEL_PRAGMA_UNROLL + for (short f = 0; f < T::kNumFrags; f++) { + out.frag_at(0, f) = a.frag_at(0, f) * b.frag_at(0, f) + c.frag_at(0, f); + } + return out; +} + +template +METAL_FUNC _M16x16 reduce(const thread T& a) { + _M16x16 out; + out.frag_at(0, 0) = a.frag_at(0, 0); + STEEL_PRAGMA_UNROLL + for (short f = 1; f < T::kNumFrags; f++) { + out.frag_at(0, 0) += a.frag_at(0, f); + } + return out; +} + +template +METAL_FUNC T scale_rows(const thread T& tile, const thread float* s) { + typename T::frag_type sv; + STEEL_PRAGMA_UNROLL + for (short i = 0; i < T::kElemsPerFrag; i++) { + sv[i] = s[i >> 2]; + } + + T out; + STEEL_PRAGMA_UNROLL + for (short f = 0; f < T::kNumFrags; f++) { + out.frag_at(0, f) = tile.frag_at(0, f) * sv; + } + return out; +} + +template +METAL_FUNC void row_sum(thread float* dst, const thread T& tile) { + STEEL_PRAGMA_UNROLL + for (short f = 0; f < T::kNumFrags; f++) { + STEEL_PRAGMA_UNROLL + for (short i = 0; i < T::kElemsPerFrag; i++) { + dst[i >> 2] += tile.frag_at(0, f)[i]; + } + } +} + +// Scales row i by DEC2[group(i)] == exp(gamma_{C-1} - gamma_i). +#define SCALE2_P(TILE0, DEC2) \ + { \ + STEEL_PRAGMA_UNROLL \ + for (short _i = 0; _i < decltype(TILE0)::kElemsPerTile; _i++) { \ + const short _w = _i % mlx::steel::BaseNAXFrag::kElemsPerFrag; \ + AT_NAX(TILE0, _i) *= (DEC2)[_w >> 2]; \ + } \ + } + +#define SUB_NAX(TILE0, TILE1, TILE2) \ + { \ + STEEL_PRAGMA_UNROLL \ + for (short _i = 0; _i < decltype(TILE0)::kElemsPerTile; _i++) { \ + AT_NAX(TILE0, _i) = AT_NAX(TILE1, _i) - AT_NAX(TILE2, _i); \ + } \ + } + +#define ADD_NAX(TILE0, TILE1, TILE2) \ + { \ + STEEL_PRAGMA_UNROLL \ + for (short _i = 0; _i < decltype(TILE0)::kElemsPerTile; _i++) { \ + AT_NAX(TILE0, _i) = AT_NAX(TILE1, _i) + AT_NAX(TILE2, _i); \ + } \ + } + +#define FMA_NAX(TILE0, S, TILE1, TILE2) \ + { \ + STEEL_PRAGMA_UNROLL \ + for (short _i = 0; _i < mlx::steel::BaseNAXFrag::kElemsPerFrag; _i++) { \ + (TILE0)[_i] = (S) * (TILE1)[_i] + (TILE2)[_i]; \ + } \ + } + +#define SCALE_NAX(TILE0, S) \ + { \ + STEEL_PRAGMA_UNROLL \ + for (short _i = 0; _i < decltype(TILE0)::kElemsPerTile; _i++) { \ + AT_NAX(TILE0, _i) *= (S); \ + } \ + } + +#define SCALE_ROW_NAX(TILE0, S) \ + { \ + STEEL_PRAGMA_UNROLL \ + for (short _i = 0; _i < decltype(TILE0)::kElemsPerTile; _i++) { \ + const short _w = _i % mlx::steel::BaseNAXFrag::kElemsPerFrag; \ + AT_NAX(TILE0, _i) *= \ + metal::fast::exp((S)[mlx::steel::BaseNAXFrag::get_coord(_w).y]); \ + } \ + } + +#define TRIL_NAX(TILE0, TILE1) \ + { \ + STEEL_PRAGMA_UNROLL \ + for (short _i = 0; _i < decltype(TILE0)::kElemsPerFrag; _i++) { \ + const short2 _c = mlx::steel::BaseNAXFrag::get_coord(_i); /* {fn, fm} */ \ + AT_NAX(TILE0, _i) = (_c.x >= _c.y) ? 0.f : AT_NAX(TILE1, _i); \ + } \ + } + +#define SCALE_BETA_NAX(TILE0, BETA2) \ + { \ + STEEL_PRAGMA_UNROLL \ + for (short _i = 0; _i < decltype(TILE0)::kElemsPerTile; _i++) { \ + const short _w = _i % mlx::steel::BaseNAXFrag::kElemsPerFrag; \ + AT_NAX(TILE0, _i) *= (BETA2)[_w >> 2]; \ + } \ + } + +#define SCALE2_NAX(TILE0, GAMMA) \ + { \ + STEEL_PRAGMA_UNROLL \ + for (short _i = 0; _i < decltype(TILE0)::kElemsPerTile; _i++) { \ + const short _w = _i % mlx::steel::BaseNAXFrag::kElemsPerFrag; \ + const short _fm = mlx::steel::BaseNAXFrag::get_coord(_w).y; \ + AT_NAX(TILE0, _i) *= metal::fast::exp((GAMMA)[(C) - 1] - (GAMMA)[_fm]); \ + } \ + } + +#define SCALE_TRI_NAX(TILE0, GAMMA) \ + { \ + STEEL_PRAGMA_UNROLL \ + for (short _i = 0; _i < decltype(TILE0)::kElemsPerFrag; _i++) { \ + const short2 _c = mlx::steel::BaseNAXFrag::get_coord(_i); /* {fn, fm} */ \ + AT_NAX(TILE0, _i) *= (_c.x > _c.y) \ + ? 0.f \ + : metal::fast::exp((GAMMA)[_c.y] - (GAMMA)[_c.x]); \ + } \ + } + +#define SCALE_TRIEQ_NAX1(TILE0, BETA) \ + { \ + STEEL_PRAGMA_UNROLL \ + for (short _i = 0; _i < decltype(TILE0)::kElemsPerFrag; _i++) { \ + const short2 _c = mlx::steel::BaseNAXFrag::get_coord(_i); /* {fn, fm} */ \ + AT_NAX(TILE0, _i) *= (_c.x >= _c.y) ? 0.f : (BETA)[_i >> 2]; \ + } \ + } + +namespace mlx { +namespace steel { +template < + typename CType, + typename AType, + typename BType, + bool transpose_a = false, + bool transpose_b = false, + mpp::tensor_ops::matmul2d_descriptor::mode Mode = + mpp::tensor_ops::matmul2d_descriptor::mode::multiply_accumulate> +METAL_FUNC static constexpr void mma( + thread BaseNAXFrag::dtype_frag_t& C, + const thread BaseNAXFrag::dtype_frag_t& A0, + const thread BaseNAXFrag::dtype_frag_t& A1, + metal::bool_constant, + const thread BaseNAXFrag::dtype_frag_t& B0, + const thread BaseNAXFrag::dtype_frag_t& B1, + metal::bool_constant) { + // M=16, N=16, K=32: A and B each two K-fragments, single 16x16 C. + constexpr auto desc = mpp::tensor_ops::matmul2d_descriptor( + 16, 16, 32, transpose_a, transpose_b, true, Mode); + + mpp::tensor_ops::matmul2d gemm_op; + + auto ct_a = + gemm_op.template get_left_input_cooperative_tensor(); + auto ct_b = + gemm_op + .template get_right_input_cooperative_tensor(); + auto ct_c = gemm_op.template get_destination_cooperative_tensor< + decltype(ct_a), + decltype(ct_b), + CType>(); + + STEEL_PRAGMA_UNROLL + for (short i = 0; i < BaseNAXFrag::kElemsPerFrag; i++) { + ct_a[i] = A0[i]; + ct_a[BaseNAXFrag::kElemsPerFrag + i] = A1[i]; + ct_b[i] = B0[i]; + ct_b[BaseNAXFrag::kElemsPerFrag + i] = B1[i]; + ct_c[i] = C[i]; + } + + gemm_op.run(ct_a, ct_b, ct_c); + + STEEL_PRAGMA_UNROLL + for (short i = 0; i < BaseNAXFrag::kElemsPerFrag; i++) { + C[i] = ct_c[i]; + } +} + +template < + typename CType, + typename AType, + typename BType, + bool transpose_a, + bool transpose_b, + mpp::tensor_ops::matmul2d_descriptor::mode Mode = + mpp::tensor_ops::matmul2d_descriptor::mode::multiply_accumulate> +METAL_FUNC static constexpr void mma( + thread BaseNAXFrag::dtype_frag_t& C, + const thread BaseNAXFrag::dtype_frag_t& A, + metal::bool_constant, + const thread BaseNAXFrag::dtype_frag_t& B, + metal::bool_constant) { + constexpr auto desc = mpp::tensor_ops::matmul2d_descriptor( + 16, 32, 16, transpose_a, transpose_b, true, Mode); + + mpp::tensor_ops::matmul2d gemm_op; + + auto ct_a = + gemm_op.template get_left_input_cooperative_tensor(); + auto ct_b = + gemm_op + .template get_right_input_cooperative_tensor(); + auto ct_c = gemm_op.template get_destination_cooperative_tensor< + decltype(ct_a), + decltype(ct_b), + CType>(); + + STEEL_PRAGMA_UNROLL + for (short i = 0; i < BaseNAXFrag::kElemsPerFrag; i++) { + ct_a[i] = A[i]; + ct_b[i] = B[i]; + ct_b[BaseNAXFrag::kElemsPerFrag + i] = 0.0; + ct_c[i] = C[i]; + ct_c[BaseNAXFrag::kElemsPerFrag + i] = 0.0; + } + + gemm_op.run(ct_a, ct_b, ct_c); + + STEEL_PRAGMA_UNROLL + for (short i = 0; i < BaseNAXFrag::kElemsPerFrag; i++) { + C[i] = ct_c[i]; + } +} + +template < + typename CType, + typename AType, + typename BType, + bool transpose_a = false, + bool transpose_b = false, + mpp::tensor_ops::matmul2d_descriptor::mode Mode = + mpp::tensor_ops::matmul2d_descriptor::mode::multiply_accumulate> +METAL_FUNC static constexpr void mman( + thread BaseNAXFrag::dtype_frag_t& Cn0, + thread BaseNAXFrag::dtype_frag_t& Cn1, + const thread BaseNAXFrag::dtype_frag_t& A, + metal::bool_constant, + const thread BaseNAXFrag::dtype_frag_t& Bn0, + const thread BaseNAXFrag::dtype_frag_t& Bn1, + metal::bool_constant) { + // M=16, N=32, K=16: single A (K=16), B and C two N-fragments each. + constexpr auto desc = mpp::tensor_ops::matmul2d_descriptor( + 16, 32, 16, transpose_a, transpose_b, true, Mode); + + // Create matmul op + mpp::tensor_ops::matmul2d gemm_op; + + // Create matmul operands in registers + auto ct_a = + gemm_op.template get_left_input_cooperative_tensor(); + auto ct_b = + gemm_op + .template get_right_input_cooperative_tensor(); + + // Create matmul output in register + auto ct_c = gemm_op.template get_destination_cooperative_tensor< + decltype(ct_a), + decltype(ct_b), + CType>(); + + // Load A in to left operand registers + STEEL_PRAGMA_UNROLL + for (short i = 0; i < BaseNAXFrag::kElemsPerFrag; i++) { + ct_a[i] = A[i]; + ct_b[i] = Bn0[i]; + ct_b[BaseNAXFrag::kElemsPerFrag + i] = Bn1[i]; + ct_c[i] = Cn0[i]; + ct_c[BaseNAXFrag::kElemsPerFrag + i] = Cn1[i]; + } + + // Do matmul + gemm_op.run(ct_a, ct_b, ct_c); + + // Copy out results + STEEL_PRAGMA_UNROLL + for (short i = 0; i < BaseNAXFrag::kElemsPerFrag; i++) { + Cn0[i] = ct_c[i]; + Cn1[i] = ct_c[BaseNAXFrag::kElemsPerFrag + i]; + } +} + +} // namespace steel +} // namespace mlx + +// Operand and accumulator types are taken from the tiles rather than pinned to +// float. Pinning them meant a bf16 tile was widened to 32 bits in registers and +// multiplied on the fp32 pipe, which doubled the operand register cost and +// threw away the native bf16 throughput. elem_type is NAXTile's element type, +// so declaring the tile as NAXTile is now enough to select the right +// instantiation. +#define MM16x16x16(C, CO, A, TA, AO, B, TB, BO) \ + mlx::steel::mma< \ + typename decltype(C)::elem_type, \ + typename decltype(A)::elem_type, \ + typename decltype(B)::elem_type, \ + TA, \ + TB, \ + mpp::tensor_ops::matmul2d_descriptor::mode::multiply>( \ + C.frag_at(0, (CO)), \ + A.frag_at(0, (AO)), \ + metal::bool_constant{}, \ + B.frag_at(0, (BO)), \ + metal::bool_constant{}); + +#define MMA16x16x16(C, CO, A, TA, AO, B, TB, BO) \ + mlx::steel::mma< \ + typename decltype(C)::elem_type, \ + typename decltype(A)::elem_type, \ + typename decltype(B)::elem_type, \ + TA, \ + TB, \ + mpp::tensor_ops::matmul2d_descriptor::mode::multiply_accumulate>( \ + C.frag_at(0, (CO)), \ + A.frag_at(0, (AO)), \ + metal::bool_constant{}, \ + B.frag_at(0, (BO)), \ + metal::bool_constant{}); + +#define MMA16x16x32(C, CO, A, TA, AO, B, TB, BO) \ + mlx::steel::mma< \ + typename decltype(C)::elem_type, \ + typename decltype(A)::elem_type, \ + typename decltype(B)::elem_type, \ + TA, \ + TB, \ + mpp::tensor_ops::matmul2d_descriptor::mode::multiply_accumulate>( \ + C.frag_at(0, (CO)), \ + A.frag_at(0, (AO)), \ + A.frag_at(0, (AO) + 1), \ + metal::bool_constant{}, \ + B.frag_at(0, (BO)), \ + B.frag_at(0, (BO) + 1), \ + metal::bool_constant{}); + +#define MM16x32x16(C, CO, A, TA, AO, B, TB, BO) \ + mlx::steel::mman< \ + typename decltype(C)::elem_type, \ + typename decltype(A)::elem_type, \ + typename decltype(B)::elem_type, \ + TA, \ + TB, \ + mpp::tensor_ops::matmul2d_descriptor::mode::multiply>( \ + C.frag_at(0, (CO)), \ + C.frag_at(0, (CO) + 1), \ + A.frag_at(0, (AO)), \ + metal::bool_constant{}, \ + B.frag_at(0, (BO)), \ + B.frag_at(0, (BO) + 1), \ + metal::bool_constant{}); + +#define MMA16x32x16(C, CO, A, TA, AO, B, TB, BO) \ + mlx::steel::mman< \ + typename decltype(C)::elem_type, \ + typename decltype(A)::elem_type, \ + typename decltype(B)::elem_type, \ + TA, \ + TB, \ + mpp::tensor_ops::matmul2d_descriptor::mode::multiply_accumulate>( \ + C.frag_at(0, (CO)), \ + C.frag_at(0, (CO) + 1), \ + A.frag_at(0, (AO)), \ + metal::bool_constant{}, \ + B.frag_at(0, (BO)), \ + B.frag_at(0, (BO) + 1), \ + metal::bool_constant{}); + +// Row sum of an elementwise product, reduced over the fragment's column axis. +// DST[0] is the fm row group, DST[1] the fm + kElemRowsJump group. +#define ROWSUM_NAX(DST, TILE0, TILE1) \ + { \ + STEEL_PRAGMA_UNROLL \ + for (short _i = 0; _i < decltype(TILE0)::kElemsPerTile; _i++) { \ + const short _w = _i % mlx::steel::BaseNAXFrag::kElemsPerFrag; \ + (DST)[_w >> 2] += AT_NAX(TILE0, _i) * AT_NAX(TILE1, _i); \ + } \ + } + +// already generic +#define SCALE_BETA_NAX_O(DST, TILE, BETA2) \ + { \ + STEEL_PRAGMA_UNROLL \ + for (short _i = 0; _i < decltype(TILE)::kElemsPerTile; _i++) { \ + const short _w = _i % mlx::steel::BaseNAXFrag::kElemsPerFrag; \ + AT_NAX(DST, _i) = AT_NAX(TILE, _i) * (BETA2)[_w >> 2]; \ + } \ + } + +#define ROWSUM1_NAX(DST, TILE0) \ + { \ + STEEL_PRAGMA_UNROLL \ + for (short _i = 0; _i < decltype(TILE0)::kElemsPerTile; _i++) { \ + const short _w = _i % mlx::steel::BaseNAXFrag::kElemsPerFrag; \ + (DST)[_w >> 2] += AT_NAX(TILE0, _i); \ + } \ + } + +#define MUL_NAX(TILE0, TILE1, TILE2) \ + { \ + STEEL_PRAGMA_UNROLL \ + for (short _i = 0; _i < decltype(TILE0)::kElemsPerTile; _i++) { \ + AT_NAX(TILE0, _i) = AT_NAX(TILE1, _i) * AT_NAX(TILE2, _i); \ + } \ + } + +#define MULA_NAX(TILE0, TILE1, TILE2) \ + { \ + STEEL_PRAGMA_UNROLL \ + for (short _i = 0; _i < decltype(TILE0)::kElemsPerTile; _i++) { \ + AT_NAX(TILE0, _i) += AT_NAX(TILE1, _i) * AT_NAX(TILE2, _i); \ + } \ + } + +#define MULS_NAX(TILE0, TILE1, TILE2) \ + { \ + STEEL_PRAGMA_UNROLL \ + for (short _i = 0; _i < decltype(TILE0)::kElemsPerTile; _i++) { \ + AT_NAX(TILE0, _i) -= AT_NAX(TILE1, _i) * AT_NAX(TILE2, _i); \ + } \ + } + +#define NSCALE_ROW_NAX(TILE0, S) \ + { \ + STEEL_PRAGMA_UNROLL \ + for (short _i = 0; _i < decltype(TILE0)::kElemsPerTile; _i++) { \ + const short _w = _i % mlx::steel::BaseNAXFrag::kElemsPerFrag; \ + AT_NAX(TILE0, _i) = -AT_NAX(TILE0, _i) * \ + metal::fast::exp((S)[mlx::steel::BaseNAXFrag::get_coord(_w).y]); \ + } \ + } + +#define SCALE_ROW_P(TILE0, ROW2) \ + { \ + STEEL_PRAGMA_UNROLL \ + for (short _i = 0; _i < decltype(TILE0)::kElemsPerTile; _i++) { \ + const short _w = _i % mlx::steel::BaseNAXFrag::kElemsPerFrag; \ + AT_NAX(TILE0, _i) *= (ROW2)[_w >> 2]; \ + } \ + } + +// As SCALE_ROW_P but also negates, in one pass. +#define NSCALE_ROW_P(TILE0, ROW2) \ + { \ + STEEL_PRAGMA_UNROLL \ + for (short _i = 0; _i < decltype(TILE0)::kElemsPerTile; _i++) { \ + const short _w = _i % mlx::steel::BaseNAXFrag::kElemsPerFrag; \ + AT_NAX(TILE0, _i) = -AT_NAX(TILE0, _i) * (ROW2)[_w >> 2]; \ + } \ + } + +// Scales row i by DEC2[group(i)] == exp(gamma_{C-1} - gamma_i). +#define SCALE2_P(TILE0, DEC2) \ + { \ + STEEL_PRAGMA_UNROLL \ + for (short _i = 0; _i < decltype(TILE0)::kElemsPerTile; _i++) { \ + const short _w = _i % mlx::steel::BaseNAXFrag::kElemsPerFrag; \ + AT_NAX(TILE0, _i) *= (DEC2)[_w >> 2]; \ + } \ + } + +// sdpa stuff + +// Operand-typed variants. The mma template was always generic over +// AType/BType; only the macros pinned them to float, which silently forced +// every bf16/fp16 matmul onto the fp32 pipe. C stays float: accumulating a +// long reduction in 16 bits loses too much. +#define MMA16x16x32_OP(OT, C, CO, A, TA, AO, B, TB, BO) \ + mlx::steel::mma< \ + float, \ + OT, \ + OT, \ + TA, \ + TB, \ + mpp::tensor_ops::matmul2d_descriptor::mode::multiply_accumulate>( \ + C.frag_at(0, (CO)), \ + A.frag_at(0, (AO)), \ + A.frag_at(0, (AO) + 1), \ + metal::bool_constant{}, \ + B.frag_at(0, (BO)), \ + B.frag_at(0, (BO) + 1), \ + metal::bool_constant{}); + +#define MMA16x32x16_OP(OT, C, CO, A, TA, AO, B, TB, BO) \ + mlx::steel::mman< \ + float, \ + OT, \ + OT, \ + TA, \ + TB, \ + mpp::tensor_ops::matmul2d_descriptor::mode::multiply_accumulate>( \ + C.frag_at(0, (CO)), \ + C.frag_at(0, (CO) + 1), \ + A.frag_at(0, (AO)), \ + metal::bool_constant{}, \ + B.frag_at(0, (BO)), \ + B.frag_at(0, (BO) + 1), \ + metal::bool_constant{}); + +// Narrow a float tile into the operand dtype. The conversion has to be +// explicit: elems() exposes the raw element type and MSL will not implicitly +// convert float to bfloat. +#define CAST_NAX(DST, SRC) \ + { \ + using _dst_elem_t = typename decltype(DST)::elem_type; \ + STEEL_PRAGMA_UNROLL \ + for (short _i = 0; _i < decltype(SRC)::kElemsPerTile; _i++) { \ + AT_NAX(DST, _i) = static_cast<_dst_elem_t>(AT_NAX(SRC, _i)); \ + } \ + } diff --git a/mlx/backend/metal/kernels/gated_delta_update.h b/mlx/backend/metal/kernels/gated_delta_update.h index 9553356c7a..78171e4b48 100644 --- a/mlx/backend/metal/kernels/gated_delta_update.h +++ b/mlx/backend/metal/kernels/gated_delta_update.h @@ -3,6 +3,8 @@ #include #include "mlx/backend/metal/kernels/utils.h" +constant bool save_state [[function_constant(200)]]; + #define AT(TILE, IDX) TILE.thread_elements()[IDX] #define SUB(TILE0, TILE1, TILE2) \ { \ @@ -263,25 +265,24 @@ template } /* - auto grid = MTL::Size(32, Dv, B * Hv); + auto grid = MTL::Size(32, Dv, B * Hv); auto threads = MTL::Size(32, 4, 1); */ -template +template [[kernel]] void gated_delta_seq( const device InT* q [[buffer(0)]], const device InT* k [[buffer(1)]], const device InT* v [[buffer(2)]], const device float* state_in [[buffer(3)]], - const device InT* g [[buffer(4)]], // [B, T, Hv] or [B, T, Hv, Dk] + const device InT* g [[buffer(4)]], // [B, T, Hv] const device InT* beta [[buffer(5)]], // [B, T, Hv] - // [B, Hv, Dv, Dk] device InT* y [[buffer(6)]], // [B, T, Hv, Dv] device float* state_out [[buffer(7)]], // [B, Hv, Dv, Dk] constant int& T [[buffer(8)]], + device float* state_cache [[buffer(9), function_constant(save_state)]], uint3 thread_position_in_grid [[thread_position_in_grid]], uint3 thread_position_in_threadgroup [[thread_position_in_threadgroup]], uint thread_index_in_simdgroup [[thread_index_in_simdgroup]]) { - // kernel implementation auto n = thread_position_in_grid.z; auto b_idx = n / Hv; auto hv_idx = n % Hv; @@ -313,7 +314,23 @@ template auto g_ = g + b_idx * T * Hv; auto beta_ = beta + b_idx * T * Hv; + // Checkpoint every Ckpt-th step rather than every step: at T = 4096 with + // B = 4, Hv = 32 a dense cache is 34 GB, which the backward's replay makes + // unnecessary. + const int n_ckpt = (T + Ckpt - 1) / Ckpt; + auto c_state_base = state_cache + n * n_ckpt * Dv * Dk + dv_idx * Dk; + for (int t = 0; t < T; ++t) { + // Stored before the decay, so slot t/Ckpt holds the state at *entry* to + // step t. The backward replays forward from here, so it must be the entry + // state and not the exit state. + if (save_state && (t % Ckpt) == 0) { + auto c_state = c_state_base + (t / Ckpt) * Dv * Dk; + for (int i = 0; i < n_per_t; ++i) { + c_state[n_per_t * dk_idx + i] = state[i]; + } + } + float kv_mem = 0.0f; for (int i = 0; i < n_per_t; ++i) { auto s_idx = n_per_t * dk_idx + i; diff --git a/mlx/backend/metal/kernels/gated_delta_update.metal b/mlx/backend/metal/kernels/gated_delta_update.metal index 21238df430..f0797e80b6 100644 --- a/mlx/backend/metal/kernels/gated_delta_update.metal +++ b/mlx/backend/metal/kernels/gated_delta_update.metal @@ -3,15 +3,17 @@ using namespace metal; -#define instantiate_gated_delta_update_seq(in_type, dk, dv, hk, hv) \ - instantiate_kernel( \ - "seq_gated_delta_" #in_type "_" #dk "_" #dv "_" #hk "_" #hv, \ - gated_delta_seq, \ - in_type, \ - dk, \ - dv, \ - hk, \ - hv) +#define instantiate_gated_delta_update_seq(in_type, dk, dv, hk, hv, ckpt) \ + instantiate_kernel( \ + "seq_gated_delta_" #in_type "_" #dk "_" #dv "_" #hk "_" #hv \ + "_" #ckpt, \ + gated_delta_seq, \ + in_type, \ + dk, \ + dv, \ + hk, \ + hv, \ + ckpt) #define instantiate_gated_delta_update_fused_chunk(in_type, dk, dv, hk, hv, c) \ instantiate_kernel( \ @@ -25,18 +27,24 @@ using namespace metal; hv, \ c) -#define instantiate_gated_delta_dims(in_type, dk, dv, hk, hv) \ - instantiate_gated_delta_update_seq(in_type, dk, dv, hk, hv) \ - instantiate_gated_delta_update_fused_chunk(in_type, dk, dv, hk, hv, 8) +// Ckpt is a template parameter of the sequential kernel, so the inference path +// has to name one even though it never stores. These must cover every value +// gated_delta_ckpt() can return, since the host puts it in the kernel name. +#define instantiate_gated_delta_dims(in_type, dk, dv, hk, hv) \ + instantiate_gated_delta_update_seq(in_type, dk, dv, hk, hv, 1) \ + instantiate_gated_delta_update_seq(in_type, dk, dv, hk, hv, 4) \ + instantiate_gated_delta_update_seq(in_type, dk, dv, hk, hv, 8) \ + instantiate_gated_delta_update_seq(in_type, dk, dv, hk, hv, 16) \ + instantiate_gated_delta_update_fused_chunk(in_type, dk, dv, hk, hv, 8) -#define instantiate_gated_delta(in_type) \ - instantiate_gated_delta_dims(in_type, 128, 128, 24, 24) \ - instantiate_gated_delta_dims(in_type, 128, 128, 32, 32) \ - instantiate_gated_delta_dims(in_type, 128, 128, 16, 32) \ - instantiate_gated_delta_dims(in_type, 128, 128, 16, 48) \ - instantiate_gated_delta_dims(in_type, 128, 128, 16, 16) \ - instantiate_gated_delta_dims(in_type, 128, 128, 16, 64) +#define instantiate_gated_delta(in_type) \ + instantiate_gated_delta_dims(in_type, 128, 128, 24, 24) \ + instantiate_gated_delta_dims(in_type, 128, 128, 32, 32) \ + instantiate_gated_delta_dims(in_type, 128, 128, 16, 32) \ + instantiate_gated_delta_dims(in_type, 128, 128, 16, 48) \ + instantiate_gated_delta_dims(in_type, 128, 128, 16, 16) \ + instantiate_gated_delta_dims(in_type, 128, 128, 16, 64) instantiate_gated_delta(float); instantiate_gated_delta(bfloat16_t); -instantiate_gated_delta(float16_t); \ No newline at end of file +instantiate_gated_delta(float16_t); diff --git a/mlx/backend/metal/kernels/gated_delta_update_nax.h b/mlx/backend/metal/kernels/gated_delta_update_nax.h index f4ae518264..a69e9df359 100644 --- a/mlx/backend/metal/kernels/gated_delta_update_nax.h +++ b/mlx/backend/metal/kernels/gated_delta_update_nax.h @@ -1,336 +1,18 @@ #pragma once -#include - -#include -#include - -#include "mlx/backend/metal/kernels/steel/gemm/nax.h" +#include "mlx/backend/metal/kernels/gated_delta_nax_ops.h" using namespace metal; using namespace mpp; using namespace mpp::tensor_ops; -// NAX MACROS I can probably do a nice template instead of doing this -// fm = base_fm + (idx >> 2) * 8; // idx>>2 = idx/4 -> 0 for idx 0-3, 1 for -// idx 4-7 fn = base_fn + (idx % 4); // 4 consecutive columns -#define AT_NAX(TILE, IDX) TILE.elems()[IDX] - -#define SUB_NAX(TILE0, TILE1, TILE2) \ - { \ - STEEL_PRAGMA_UNROLL \ - for (short _i = 0; _i < decltype(TILE0)::kElemsPerFrag; _i++) { \ - AT_NAX(TILE0, _i) = AT_NAX(TILE1, _i) - AT_NAX(TILE2, _i); \ - } \ - } - -#define ADD_NAX(TILE0, TILE1, TILE2) \ - { \ - STEEL_PRAGMA_UNROLL \ - for (short _i = 0; _i < decltype(TILE0)::kElemsPerFrag; _i++) { \ - AT_NAX(TILE0, _i) = AT_NAX(TILE1, _i) + AT_NAX(TILE2, _i); \ - } \ - } - -#define FMA_NAX(TILE0, S, TILE1, TILE2) \ - { \ - STEEL_PRAGMA_UNROLL \ - for (short _i = 0; _i < mlx::steel::BaseNAXFrag::kElemsPerFrag; _i++) { \ - (TILE0)[_i] = (S) * (TILE1)[_i] + (TILE2)[_i]; \ - } \ - } - -#define SCALE_NAX(TILE0, S) \ - { \ - STEEL_PRAGMA_UNROLL \ - for (short _i = 0; _i < decltype(TILE0)::kElemsPerTile; _i++) { \ - AT_NAX(TILE0, _i) *= (S); \ - } \ - } - -#define SCALE_ROW_NAX(TILE0, S) \ - { \ - STEEL_PRAGMA_UNROLL \ - for (short _i = 0; _i < decltype(TILE0)::kElemsPerTile; _i++) { \ - const short _w = _i % mlx::steel::BaseNAXFrag::kElemsPerFrag; \ - AT_NAX(TILE0, _i) *= \ - metal::fast::exp((S)[mlx::steel::BaseNAXFrag::get_coord(_w).y]); \ - } \ - } - -#define SCALE_BETA_NAX(TILE0, BETA2) \ - { \ - STEEL_PRAGMA_UNROLL \ - for (short _i = 0; _i < decltype(TILE0)::kElemsPerTile; _i++) { \ - const short _w = _i % mlx::steel::BaseNAXFrag::kElemsPerFrag; \ - AT_NAX(TILE0, _i) *= (BETA2)[_w >> 2]; \ - } \ - } - -#define SCALE2_NAX(TILE0, GAMMA) \ - { \ - STEEL_PRAGMA_UNROLL \ - for (short _i = 0; _i < decltype(TILE0)::kElemsPerTile; _i++) { \ - const short _w = _i % mlx::steel::BaseNAXFrag::kElemsPerFrag; \ - const short _fm = mlx::steel::BaseNAXFrag::get_coord(_w).y; \ - AT_NAX(TILE0, _i) *= metal::fast::exp((GAMMA)[(C) - 1] - (GAMMA)[_fm]); \ - } \ - } - -#define SCALE_TRI_NAX(TILE0, GAMMA) \ - { \ - STEEL_PRAGMA_UNROLL \ - for (short _i = 0; _i < decltype(TILE0)::kElemsPerFrag; _i++) { \ - const short2 _c = mlx::steel::BaseNAXFrag::get_coord(_i); /* {fn, fm} */ \ - AT_NAX(TILE0, _i) *= (_c.x > _c.y) \ - ? 0.f \ - : metal::fast::exp((GAMMA)[_c.y] - (GAMMA)[_c.x]); \ - } \ - } - -#define SCALE_TRIEQ_NAX1(TILE0, BETA) \ - { \ - STEEL_PRAGMA_UNROLL \ - for (short _i = 0; _i < decltype(TILE0)::kElemsPerFrag; _i++) { \ - const short2 _c = mlx::steel::BaseNAXFrag::get_coord(_i); /* {fn, fm} */ \ - AT_NAX(TILE0, _i) *= (_c.x >= _c.y) ? 0.f : (BETA)[_i >> 2]; \ - } \ - } - -namespace mlx { -namespace steel { -template < - typename CType, - typename AType, - typename BType, - bool transpose_a = false, - bool transpose_b = false, - mpp::tensor_ops::matmul2d_descriptor::mode Mode = - mpp::tensor_ops::matmul2d_descriptor::mode::multiply_accumulate> -METAL_FUNC static constexpr void mma( - thread BaseNAXFrag::dtype_frag_t& C, - const thread BaseNAXFrag::dtype_frag_t& A0, - const thread BaseNAXFrag::dtype_frag_t& A1, - metal::bool_constant, - const thread BaseNAXFrag::dtype_frag_t& B0, - const thread BaseNAXFrag::dtype_frag_t& B1, - metal::bool_constant) { - // M=16, N=16, K=32: A and B each two K-fragments, single 16x16 C. - constexpr auto desc = mpp::tensor_ops::matmul2d_descriptor( - 16, 16, 32, transpose_a, transpose_b, true, Mode); - - mpp::tensor_ops::matmul2d gemm_op; - - auto ct_a = - gemm_op.template get_left_input_cooperative_tensor(); - auto ct_b = - gemm_op - .template get_right_input_cooperative_tensor(); - auto ct_c = gemm_op.template get_destination_cooperative_tensor< - decltype(ct_a), - decltype(ct_b), - CType>(); - - STEEL_PRAGMA_UNROLL - for (short i = 0; i < BaseNAXFrag::kElemsPerFrag; i++) { - ct_a[i] = A0[i]; - ct_a[BaseNAXFrag::kElemsPerFrag + i] = A1[i]; - ct_b[i] = B0[i]; - ct_b[BaseNAXFrag::kElemsPerFrag + i] = B1[i]; - ct_c[i] = C[i]; - } - - gemm_op.run(ct_a, ct_b, ct_c); - - STEEL_PRAGMA_UNROLL - for (short i = 0; i < BaseNAXFrag::kElemsPerFrag; i++) { - C[i] = ct_c[i]; - } -} - -template < - typename CType, - typename AType, - typename BType, - bool transpose_a, - bool transpose_b, - mpp::tensor_ops::matmul2d_descriptor::mode Mode = - mpp::tensor_ops::matmul2d_descriptor::mode::multiply_accumulate> -METAL_FUNC static constexpr void mma( - thread BaseNAXFrag::dtype_frag_t& C, - const thread BaseNAXFrag::dtype_frag_t& A, - metal::bool_constant, - const thread BaseNAXFrag::dtype_frag_t& B, - metal::bool_constant) { - constexpr auto desc = mpp::tensor_ops::matmul2d_descriptor( - 16, 32, 16, transpose_a, transpose_b, true, Mode); - - mpp::tensor_ops::matmul2d gemm_op; - - auto ct_a = - gemm_op.template get_left_input_cooperative_tensor(); - auto ct_b = - gemm_op - .template get_right_input_cooperative_tensor(); - auto ct_c = gemm_op.template get_destination_cooperative_tensor< - decltype(ct_a), - decltype(ct_b), - CType>(); - - STEEL_PRAGMA_UNROLL - for (short i = 0; i < BaseNAXFrag::kElemsPerFrag; i++) { - ct_a[i] = A[i]; - ct_b[i] = B[i]; - ct_b[BaseNAXFrag::kElemsPerFrag + i] = 0.0; - ct_c[i] = C[i]; - ct_c[BaseNAXFrag::kElemsPerFrag + i] = 0.0; - } +/////////////////////////////////////////////////////////////////////////////// +// Function constants +/////////////////////////////////////////////////////////////////////////////// - gemm_op.run(ct_a, ct_b, ct_c); +constant bool save_state_cache [[function_constant(200)]]; - STEEL_PRAGMA_UNROLL - for (short i = 0; i < BaseNAXFrag::kElemsPerFrag; i++) { - C[i] = ct_c[i]; - } -} - -template < - typename CType, - typename AType, - typename BType, - bool transpose_a = false, - bool transpose_b = false, - mpp::tensor_ops::matmul2d_descriptor::mode Mode = - mpp::tensor_ops::matmul2d_descriptor::mode::multiply_accumulate> -METAL_FUNC static constexpr void mman( - thread BaseNAXFrag::dtype_frag_t& Cn0, - thread BaseNAXFrag::dtype_frag_t& Cn1, - const thread BaseNAXFrag::dtype_frag_t& A, - metal::bool_constant, - const thread BaseNAXFrag::dtype_frag_t& Bn0, - const thread BaseNAXFrag::dtype_frag_t& Bn1, - metal::bool_constant) { - // M=16, N=32, K=16: single A (K=16), B and C two N-fragments each. - constexpr auto desc = mpp::tensor_ops::matmul2d_descriptor( - 16, 32, 16, transpose_a, transpose_b, true, Mode); - - // Create matmul op - mpp::tensor_ops::matmul2d gemm_op; - - // Create matmul operands in registers - auto ct_a = - gemm_op.template get_left_input_cooperative_tensor(); - auto ct_b = - gemm_op - .template get_right_input_cooperative_tensor(); - - // Create matmul output in register - auto ct_c = gemm_op.template get_destination_cooperative_tensor< - decltype(ct_a), - decltype(ct_b), - CType>(); - - // Load A in to left operand registers - STEEL_PRAGMA_UNROLL - for (short i = 0; i < BaseNAXFrag::kElemsPerFrag; i++) { - ct_a[i] = A[i]; - ct_b[i] = Bn0[i]; - ct_b[BaseNAXFrag::kElemsPerFrag + i] = Bn1[i]; - ct_c[i] = Cn0[i]; - ct_c[BaseNAXFrag::kElemsPerFrag + i] = Cn1[i]; - } - - // Do matmul - gemm_op.run(ct_a, ct_b, ct_c); - - // Copy out results - STEEL_PRAGMA_UNROLL - for (short i = 0; i < BaseNAXFrag::kElemsPerFrag; i++) { - Cn0[i] = ct_c[i]; - Cn1[i] = ct_c[BaseNAXFrag::kElemsPerFrag + i]; - } -} - -} // namespace steel -} // namespace mlx - -#define MM16x16x16(C, CO, A, TA, AO, B, TB, BO) \ - mlx::steel::mma< \ - float, \ - float, \ - float, \ - TA, \ - TB, \ - mpp::tensor_ops::matmul2d_descriptor::mode::multiply>( \ - C.frag_at(0, (CO)), \ - A.frag_at(0, (AO)), \ - metal::bool_constant{}, \ - B.frag_at(0, (BO)), \ - metal::bool_constant{}); - -#define MMA16x16x16(C, CO, A, TA, AO, B, TB, BO) \ - mlx::steel::mma< \ - float, \ - float, \ - float, \ - TA, \ - TB, \ - mpp::tensor_ops::matmul2d_descriptor::mode::multiply_accumulate>( \ - C.frag_at(0, (CO)), \ - A.frag_at(0, (AO)), \ - metal::bool_constant{}, \ - B.frag_at(0, (BO)), \ - metal::bool_constant{}); - -#define MMA16x16x32(C, CO, A, TA, AO, B, TB, BO) \ - mlx::steel::mma< \ - float, \ - float, \ - float, \ - TA, \ - TB, \ - mpp::tensor_ops::matmul2d_descriptor::mode::multiply_accumulate>( \ - C.frag_at(0, (CO)), \ - A.frag_at(0, (AO)), \ - A.frag_at(0, (AO) + 1), \ - metal::bool_constant{}, \ - B.frag_at(0, (BO)), \ - B.frag_at(0, (BO) + 1), \ - metal::bool_constant{}); - -#define MM16x32x16(C, CO, A, TA, AO, B, TB, BO) \ - mlx::steel::mman< \ - float, \ - float, \ - float, \ - TA, \ - TB, \ - mpp::tensor_ops::matmul2d_descriptor::mode::multiply>( \ - C.frag_at(0, (CO)), \ - C.frag_at(0, (CO) + 1), \ - A.frag_at(0, (AO)), \ - metal::bool_constant{}, \ - B.frag_at(0, (BO)), \ - B.frag_at(0, (BO) + 1), \ - metal::bool_constant{}); - -#define MMA16x32x16(C, CO, A, TA, AO, B, TB, BO) \ - mlx::steel::mman< \ - float, \ - float, \ - float, \ - TA, \ - TB, \ - mpp::tensor_ops::matmul2d_descriptor::mode::multiply_accumulate>( \ - C.frag_at(0, (CO)), \ - C.frag_at(0, (CO) + 1), \ - A.frag_at(0, (AO)), \ - metal::bool_constant{}, \ - B.frag_at(0, (BO)), \ - B.frag_at(0, (BO) + 1), \ - metal::bool_constant{}); - -template +template [[kernel]] void gated_delta_fused_nax( const device InT* q [[buffer(0)]], const device InT* k [[buffer(1)]], @@ -341,6 +23,9 @@ template device InT* y [[buffer(6)]], device float* state_out [[buffer(7)]], constant int& T [[buffer(8)]], + device float* state_cache [[buffer(9)]], // [B, Hv, n_ckpt, Dv, Dk] + device float* chunk_mats [[buffer(10)]], // [B, Hv, n_chunks, 3, 16, 16] + device float* chunk_delta [[buffer(11)]], // [B, Hv, n_chunks, C, Dv] uint3 thread_position_in_grid [[thread_position_in_grid]], uint3 thread_position_in_threadgroup [[thread_position_in_threadgroup]], uint thread_index_in_simdgroup [[thread_index_in_simdgroup]]) { @@ -373,6 +58,13 @@ template auto i_state = state_in + (n * Dv + dv_idx) * Dk; auto o_state = state_out + (n * Dv + dv_idx) * Dk; + const int n_chunks = (T + C - 1) / C; + const int n_ckpt = (n_chunks + Ckpt - 1) / Ckpt; + auto c_state = state_cache + (n * n_ckpt * Dv + dv_idx) * Dk; + + auto o_chunk_mats = chunk_mats + n * n_chunks * 3 * 256; + auto o_chunk_delta = chunk_delta + n * n_chunks * C * Dv + dv_idx; + threadgroup float gamma_all[C * 4]; threadgroup float* gamma = gamma_all + sg_id * C; @@ -405,6 +97,10 @@ template } mlx::steel::NAXTile TMP_tile; + // Which chunk process_chunk is on, so it knows whether this is a segment + // boundary. + int chunk_idx = 0; + auto process_chunk = [&](const short valid_rows, auto bounded_tag) __attribute__((always_inline)) { constexpr bool B = decltype(bounded_tag)::value; @@ -417,19 +113,23 @@ template } }; + // Checkpoint the state entering this chunk before anything mutates it, on + // segment boundaries only. + if (save_state_cache && (chunk_idx % Ckpt) == 0) { + S_tile.store(c_state, Dk); + c_state += Dv * Dk; + } + float g_val = (thread_index_in_simdgroup < (uint)valid_rows) ? metal::fast::log( metal::max( - static_cast( - g_[thread_index_in_simdgroup * Hv + hv_idx]), - 1e-6f)) + float(g_[thread_index_in_simdgroup * Hv + hv_idx]), 1e-6)) : 0.0f; auto gamma_val = simd_prefix_inclusive_sum(g_val); if (thread_index_in_simdgroup < C) { gamma[thread_index_in_simdgroup] = static_cast(gamma_val); } - simdgroup_barrier(mem_flags::mem_threadgroup); beta_fm[0] = (fm < valid_rows) ? beta_[fm * Hv + hv_idx] : 0.0f; const short fm1 = fm + mlx::steel::BaseNAXFrag::kElemRowsJump; @@ -441,6 +141,12 @@ template MMA16x16x32(KKt_tile, 0, K_tile, false, 0, K_tile, true, 0); } + // KKt_raw is identical for every Dv tile of this chunk; only the first + // tile needs to write it out for the vjp kernel to reuse. + if (save_state_cache && dv_idx == 0) { + KKt_tile.store(o_chunk_mats, 16); + } + KKtK_tile = KKt_tile; SCALE_TRIEQ_NAX1(KKtK_tile, beta_fm); @@ -451,6 +157,12 @@ template SUB_NAX(Tinv_tile, I_tile, TMP_tile); } + // Same as KKt_raw: cache the pre-decay inverse before it is overwritten + // below by the gamma-scaled (TUinv) version. + if (save_state_cache && dv_idx == 0) { + Tinv_tile.store(o_chunk_mats + 256, 16); + } + STEEL_PRAGMA_UNROLL for (short nn = 0; nn < Dk / 16; nn += 2) { load_seq(K_tile, k_ + nn * 16, Dk * Hk); @@ -472,20 +184,29 @@ template SUB_NAX(delta_tile, U_tile, WS_tile) - tmp_tile.clear(); QKt_tile.clear(); + tmp_tile.clear(); for (int kk = 0; kk < Dk; kk += 32) { load_seq(Q_tile, q_ + kk, Hk * Dk); load_seq(K_tile, k_ + kk, Hk * Dk); - MMA16x16x32(QKt_tile, 0, Q_tile, false, 0, K_tile, true, 0); - SCALE_ROW_NAX(Q_tile, gamma); MMA16x16x32(tmp_tile, 0, Q_tile, false, 0, S_tile, true, kk / 16); } - SCALE_TRI_NAX(QKt_tile, gamma) + if (save_state_cache) { + if (dv_idx == 0) { + QKt_tile.store(o_chunk_mats + 512, 16); + } + // Every Dv tile stores its own slice. + delta_tile.store(o_chunk_delta, Dv); + o_chunk_mats += 3 * 256; + o_chunk_delta += C * Dv; + } + // Output path runs unconditionally: the forward serves both inference and + // the training save, so y and state_out are always needed. + SCALE_TRI_NAX(QKt_tile, gamma) out_tile = tmp_tile; MMA16x16x16(out_tile, 0, QKt_tile, false, 0, delta_tile, false, 0); @@ -507,6 +228,8 @@ template SCALE2_NAX(K_tile, gamma); MMA16x32x16(S_tile, kk / 16, delta_tile, true, 0, K_tile, false, 0); } + + chunk_idx++; }; int t = 0; @@ -524,4 +247,4 @@ template } S_tile.store(o_state, Dk); -} \ No newline at end of file +} diff --git a/mlx/backend/metal/kernels/gated_delta_update_nax.metal b/mlx/backend/metal/kernels/gated_delta_update_nax.metal index 35790fa49f..de84d866ce 100644 --- a/mlx/backend/metal/kernels/gated_delta_update_nax.metal +++ b/mlx/backend/metal/kernels/gated_delta_update_nax.metal @@ -3,20 +3,25 @@ using namespace metal; -#define instantiate_gated_delta_update_fused_nax(in_type, dk, dv, hk, hv, c) \ - instantiate_kernel( \ - "gated_delta_fused_nax_" #in_type "_" #dk "_" #dv "_" #hk "_" #hv \ - "_" #c, \ - gated_delta_fused_nax, \ - in_type, \ - dk, \ - dv, \ - hk, \ - hv, \ - c) +#define instantiate_gated_delta_update_fused_nax( \ + in_type, dk, dv, hk, hv, c, ckpt) \ + instantiate_kernel( \ + "gated_delta_fused_nax_" #in_type "_" #dk "_" #dv \ + "_" #hk "_" #hv "_" #c "_" #ckpt, \ + gated_delta_fused_nax, \ + in_type, \ + dk, \ + dv, \ + hk, \ + hv, \ + c, \ + ckpt) -#define instantiate_gated_delta_dims(in_type, dk, dv, hk, hv) \ - instantiate_gated_delta_update_fused_nax(in_type, dk, dv, hk, hv, 16) +#define instantiate_gated_delta_dims(in_type, dk, dv, hk, hv) \ + instantiate_gated_delta_update_fused_nax(in_type, dk, dv, hk, hv, 16, 1) \ + instantiate_gated_delta_update_fused_nax(in_type, dk, dv, hk, hv, 16, 4) \ + instantiate_gated_delta_update_fused_nax(in_type, dk, dv, hk, hv, 16, 8) \ + instantiate_gated_delta_update_fused_nax(in_type, dk, dv, hk, hv, 16, 16) #define instantiate_gated_delta(in_type) \ instantiate_gated_delta_dims(in_type, 128, 128, 24, 24) \ diff --git a/mlx/backend/metal/kernels/gated_delta_update_nax_vjp.h b/mlx/backend/metal/kernels/gated_delta_update_nax_vjp.h new file mode 100644 index 0000000000..f00ea00a1a --- /dev/null +++ b/mlx/backend/metal/kernels/gated_delta_update_nax_vjp.h @@ -0,0 +1,1736 @@ +#pragma once + +#include "mlx/backend/metal/kernels/gated_delta_nax_ops.h" + +#include +#include "mlx/backend/metal/kernels/atomic.h" + +using namespace metal; +using namespace mpp; +using namespace mpp::tensor_ops; + +template +METAL_FUNC void reduce_tile_tg( + thread TileT& acc, + threadgroup float* scratch, + device mlx_atomic* out, + int row_stride, + int col_off, + short valid_rows, + const short sg_id, + const ushort simd_lane_id) { + constexpr short kE = TileT::kElemsPerTile; + + STEEL_PRAGMA_UNROLL + for (short _i = 0; _i < kE; _i++) { + scratch[(sg_id * 32 + simd_lane_id) * kE + _i] = AT_NAX(acc, _i); + } + threadgroup_barrier(mem_flags::mem_threadgroup); + + if (sg_id == 0) { + STEEL_PRAGMA_UNROLL + for (short _s = 1; _s < kNSG; _s++) { + STEEL_PRAGMA_UNROLL + for (short _i = 0; _i < kE; _i++) { + AT_NAX(acc, _i) += scratch[(_s * 32 + simd_lane_id) * kE + _i]; + } + } + } + threadgroup_barrier(mem_flags::mem_threadgroup); + + // Simdgroup 0 alone holds the total. + if (sg_id != 0) { + return; + } + + STEEL_PRAGMA_UNROLL + for (short _i = 0; _i < kE; _i++) { + const short _f = _i / mlx::steel::BaseNAXFrag::kElemsPerFrag; + const short _w = _i % mlx::steel::BaseNAXFrag::kElemsPerFrag; + const short2 _c = mlx::steel::BaseNAXFrag::get_coord(_w); // {fn, fm} + if (_c.y < valid_rows) { + const int idx = _c.y * row_stride + col_off + _c.x + _f * 16; + if (kGQA) { + mlx_atomic_fetch_add_explicit(out, AT_NAX(acc, _i), idx); + } else { + mlx_atomic_store_explicit(out, AT_NAX(acc, _i), idx); + } + } + } +} + +template +[[kernel]] void gated_delta_vjp_fused_nax( + const device InT* q [[buffer(0)]], // [B, T, Hk, Dk] + const device InT* k [[buffer(1)]], // [B, T, Hk, Dk] + const device InT* v [[buffer(2)]], // [B, T, Hv, Dv] + const device InT* g [[buffer(3)]], // [B, T, Hv] + const device InT* beta [[buffer(4)]], // [B, T, Hv] + const device InT* cot_o [[buffer(5)]], // [B, T, Hv, Dv] + const device float* cot_h [[buffer(6)]], // [B, Hv, Dv, Dk] + const device float* state_cache [[buffer(7)]], // [B, Hv, n_ckpt, Dv, Dk] + constant int& T [[buffer(8)]], + device mlx_atomic* dq [[buffer(9)]], + device mlx_atomic* dk [[buffer(10)]], + device float* dv [[buffer(11)]], + device float* dg [[buffer(12)]], + device float* db [[buffer(13)]], + device float* dh [[buffer(14)]], + const device float* chunk_mats [[buffer(15)]], + const device float* chunk_delta [[buffer(16)]], // [B, Hv, n_chunks, C, Dv] + device float* seg_states [[buffer(17)]], + uint3 thread_position_in_grid [[thread_position_in_grid]], + uint3 thread_position_in_threadgroup [[thread_position_in_threadgroup]], + uint thread_index_in_simdgroup [[thread_index_in_simdgroup]]) { + using _M16xDk = mlx::steel::NAXTile; + + auto n = thread_position_in_grid.z; + auto b_idx = n / Hv; + auto hv_idx = n % Hv; + auto hk_idx = hv_idx / (Hv / Hk); + + auto dv_idx = thread_position_in_grid.y * 16; + const short sg_id = thread_position_in_threadgroup.y; + + const ushort simd_lane_id = __metal_get_thread_index_in_simdgroup(ushort()); + const short qid = simd_lane_id >> 2; + const short fm = ((qid & 4) | ((simd_lane_id >> 1) & 3)); + const short fm1 = fm + mlx::steel::BaseNAXFrag::kElemRowsJump; + + const int n_chunks = (T + C - 1) / C; + // Ceiling: the floor form gives 0 when n_chunks < Ckpt, which would leave the + // allocation empty and fault on the first read. + const int n_ckpt = (n_chunks + Ckpt - 1) / Ckpt; + const short tail = short(T - (n_chunks - 1) * C); + + // Bases, indexed by chunk rather than walked. With two nested loops the + // incremental decrements are no longer expressible. + auto q_base = q + b_idx * T * Hk * Dk + hk_idx * Dk; + auto k_base = k + b_idx * T * Hk * Dk + hk_idx * Dk; + auto dq_base = dq + b_idx * T * Hk * Dk + hk_idx * Dk; + auto dk_base = dk + b_idx * T * Hk * Dk + hk_idx * Dk; + + auto v_base = v + b_idx * T * Hv * Dv + hv_idx * Dv; + auto dv_base = dv + b_idx * T * Hv * Dv + hv_idx * Dv; + auto co_base = cot_o + b_idx * T * Hv * Dv + hv_idx * Dv; + + auto g_base = g + b_idx * T * Hv; + auto beta_base = beta + b_idx * T * Hv; + auto dg_base = dg + b_idx * T * Hv; + auto db_base = db + b_idx * T * Hv; + + auto cm_base = chunk_mats + n * n_chunks * 3 * 256; + auto cd_base = chunk_delta + n * n_chunks * C * Dv + dv_idx; + auto ck_base = state_cache + (n * n_ckpt * Dv + dv_idx) * Dk; + // Each simdgroup owns a distinct dv slice, so these regions are disjoint and + // the replay needs no cross-simdgroup synchronisation. + auto seg_base = seg_states + (n * Ckpt * Dv + dv_idx) * Dk; + + // Per-chunk pointers, assigned from the bases before each process_chunk call. + auto q_ = q_base; + auto k_ = k_base; + auto dq_ = dq_base; + auto dk_ = dk_base; + auto v_ = v_base; + auto dv_ = dv_base; + auto co_ = co_base; + auto g_ = g_base; + auto beta_ = beta_base; + auto dg_ = dg_base; + auto db_ = db_base; + auto i_chunk_mats = cm_base; + auto i_chunk_delta = cd_base; + auto i_seg = seg_base; + + auto i_cot_h = cot_h + (n * Dv + dv_idx) * Dk; + auto o_dh = dh + (n * Dv + dv_idx) * Dk; + + // One threadgroup per (b, hv) covers every Dv slice, so the Dv reduction + // never leaves threadgroup memory. + constexpr int kNSG = Dv / 16; + // Under GQA several threadgroups still share an hk, so dq/dk keep an atomic + // for that axis -- but with kNSG writers folded away first. + constexpr bool kGQA = (Hv != Hk); + + threadgroup float gamma_all[C * kNSG]; + threadgroup float* gamma = gamma_all + sg_id * C; + + threadgroup float red_scratch[(kNSG / 1) * 32 * 16]; + threadgroup float db_stage[kNSG * C]; + threadgroup float dg_stage[kNSG * C]; + + float beta_fm[2]; + + // Carried state gradient, dL/dS. Same [dv, dk] orientation as S_tile. + _M16xDk dS_tile; + dS_tile.load(i_cot_h, Dk); + + // Forward state at chunk entry, taken from the replayed segment. + _M16xDk S_tile; + + _M16x32 K_tile, Q_tile; + _M16x16 V_tile; + _M16x16 delta_tile; + _M16x16 QKt_tile, QKt_raw; + _M16x16 KKt_tile; + _M16x16 TWinv_tile, TUinv_tile; + _M16x16 TMP_tile; + + _M16x16 dout_tile; + _M16x16 ddelta_tile; + _M16x16 dTU_tile, dTW_tile; + _M16x16 dTinv_tile; + _M16x16 dA_tile, G_tile; + _M16x16 tri_decay; + + _M16x16 I_tile; + STEEL_PRAGMA_UNROLL + for (short _i = 0; _i < decltype(I_tile)::kElemsPerFrag; _i++) { + const short2 _c = mlx::steel::BaseNAXFrag::get_coord(_i); /* {fn, fm} */ + AT_NAX(I_tile, _i) = (_c.x == _c.y) ? 1.0f : 0.0f; + } + + _M16x16 Ones_tile; + STEEL_PRAGMA_UNROLL + for (short _i = 0; _i < decltype(Ones_tile)::kElemsPerFrag; _i++) { + AT_NAX(Ones_tile, _i) = 1.0f; + } + + _M16x16 tril_mask; + STEEL_PRAGMA_UNROLL + for (short _i = 0; _i < decltype(tril_mask)::kElemsPerFrag; _i++) { + const short2 _c = mlx::steel::BaseNAXFrag::get_coord(_i); + AT_NAX(tril_mask, _i) = (_c.x >= _c.y) ? 0.0f : 1.0f; + } + + auto process_chunk = [&](const short valid_rows, + auto bounded_tag) __attribute__((always_inline)) { + constexpr bool B = decltype(bounded_tag)::value; + + auto load_seq = [&](thread auto& tile, auto src, int ld) { + if constexpr (B) { + tile.load_rows(src, ld, valid_rows); + } else { + tile.load(src, ld); + } + }; + + // From the replayed segment rather than a per-chunk checkpoint. + S_tile.load(i_seg, Dk); + + float g_val = (thread_index_in_simdgroup < (uint)valid_rows) + ? metal::fast::log( + metal::max(g_[thread_index_in_simdgroup * Hv + hv_idx], 1e-6)) + : 0.0f; + + auto gamma_val = simd_prefix_inclusive_sum(g_val); + if (thread_index_in_simdgroup < C) { + gamma[thread_index_in_simdgroup] = static_cast(gamma_val); + } + simdgroup_barrier(mem_flags::mem_threadgroup); + + const float gamma_last = gamma[C - 1]; + const float gamma_last_exp = metal::fast::exp(gamma_last); + + beta_fm[0] = (fm < valid_rows) ? beta_[fm * Hv + hv_idx] : 0.0f; + beta_fm[1] = (fm1 < valid_rows) ? beta_[fm1 * Hv + hv_idx] : 0.0f; + + // Per-lane decay factors. gamma[j] == gamma[valid_rows-1] for + // j >= valid_rows (the prefix sum of zeros), so the tail chunk needs no + // special case here. + const float row_exp[2] = { + metal::fast::exp(gamma[fm]), metal::fast::exp(gamma[fm1])}; + const float dec_exp[2] = { + metal::fast::exp(gamma_last - gamma[fm]), + metal::fast::exp(gamma_last - gamma[fm1])}; + + STEEL_PRAGMA_UNROLL + for (short _i = 0; _i < decltype(tri_decay)::kElemsPerFrag; _i++) { + const short2 _c = mlx::steel::BaseNAXFrag::get_coord(_i); // {fn, fm} + AT_NAX(tri_decay, _i) = + (_c.x > _c.y) ? 0.f : metal::fast::exp(gamma[_c.y] - gamma[_c.x]); + } + + KKt_tile.load(i_chunk_mats, 16); + TWinv_tile.load(i_chunk_mats + 256, 16); + QKt_raw.load(i_chunk_mats + 512, 16); + + QKt_tile = QKt_raw * tri_decay; + TUinv_tile = TWinv_tile * tri_decay; + + delta_tile.load(i_chunk_delta, Dv); + + // dgamma accumulators: + // dgam_row : per-row + // dgam_pair : row-minus-col + // dgam_last : the gamma_{C-1} + _M16x16 dgam_row, dgam_pair; + dgam_row.clear(); + dgam_pair.clear(); + float dgam_last = 0.0f; + float dbeta_acc[2] = {0.0f, 0.0f}; + + load_seq(dout_tile, co_ + dv_idx, Hv * Dv); + + // dODT = (dO @ delta.T * D) + _M16x16 dODT; + MM16x16x16(dODT, 0, dout_tile, false, 0, delta_tile, true, 0); + dODT = dODT * tri_decay; + + // cannot fuse this because ddelta is used by the next loop + MM16x16x16(ddelta_tile, 0, QKt_tile, true, 0, dout_tile, false, 0); + for (int kk = 0; kk < Dk; kk += 32) { + load_seq(K_tile, k_ + kk, Dk * Hk); + K_tile = scale_rows(K_tile, dec_exp); + MMA16x16x32(ddelta_tile, 0, K_tile, false, 0, dS_tile, true, kk / 16); + } + + // ddelta = M^T @ dO + K_dec @ dS^T + // dq = gamma * (dO @ S) + dODT @ K + dTW_tile.clear(); + for (int kk = 0; kk < Dk; kk += 32) { + _M16x32 dW_raw; + _M16x32 dq_acc; + load_seq(K_tile, k_ + kk, Dk * Hk); + + MM16x32x16(dW_raw, 0, ddelta_tile, false, 0, S_tile, false, kk / 16); + dW_raw = scale_rows(-dW_raw, row_exp); + _M16x32 Kb_tile = scale_rows(K_tile, beta_fm); + MMA16x16x32(dTW_tile, 0, dW_raw, false, 0, Kb_tile, true, 0); + + MM16x32x16(dq_acc, 0, dout_tile, false, 0, S_tile, false, kk / 16); + dq_acc = scale_rows(dq_acc, row_exp); + MMA16x32x16(dq_acc, 0, dODT, false, 0, K_tile, false, 0); + + // dq reduces over Dv, which is entirely inside this threadgroup. + reduce_tile_tg( + dq_acc, + red_scratch, + dq_, + Hk * Dk, + kk, + valid_rows, + sg_id, + simd_lane_id); + } + + // dv = B(Tu.T @ ddelta). Indexed by dv_idx, so no reduction. + _M16x16 dVb_tile; + MM16x16x16(dVb_tile, 0, TUinv_tile, true, 0, ddelta_tile, false, 0); + _M16x16 TUddelta = dVb_tile; // pre-beta, reused by dbeta + dVb_tile = scale_rows(dVb_tile, beta_fm); + + dVb_tile.store_rows(dv_ + dv_idx, Hv * Dv, valid_rows); + + // dT_U = ddelta @ (beta * V).T + load_seq(V_tile, v_ + dv_idx, Dv * Hv); + V_tile = scale_rows(V_tile, beta_fm); + MM16x16x16(dTU_tile, 0, ddelta_tile, false, 0, V_tile, true, 0); + dTU_tile = dTU_tile * tri_decay; + + dTinv_tile = dTW_tile + dTU_tile; + + // dA = -T.T @ dTinv @ T.T + MM16x16x16(TMP_tile, 0, TWinv_tile, true, 0, dTinv_tile, false, 0); + MM16x16x16(dA_tile, 0, TMP_tile, false, 0, TWinv_tile, true, 0); + dA_tile = dA_tile * -1.0f; + + // G = tril_(dA) * beta, and GGt = G + G.T + G_tile = scale_rows(dA_tile, beta_fm) * tril_mask; + _M16x16 GGt_tile = G_tile; + MMA16x16x16(GGt_tile, 0, I_tile, true, 0, G_tile, true, 0); + + // dgamma pair sites + _M16x16 P_tile = QKt_raw * dODT; + _M16x16 R_tile = TWinv_tile * dTU_tile; + dgam_pair = P_tile + R_tile; + + _M16x32 KdKb; // rowsum(K * dK_b) source, needed for dbeta + KdKb.clear(); + _M16x32 dgam_row32; + dgam_row32.clear(); + for (int kk = 0; kk < Dk; kk += 32) { + _M16x32 dk_acc, dW_raw, t1; + + load_seq(Q_tile, q_ + kk, Hk * Dk); + load_seq(K_tile, k_ + kk, Dk * Hk); + + { + _M16x32 QgdOS; + MM16x32x16(QgdOS, 0, dout_tile, false, 0, S_tile, false, kk / 16); + dgam_row32 = dgam_row32 + scale_rows(QgdOS, row_exp) * Q_tile; + } + + MM16x32x16(dk_acc, 0, dODT, true, 0, Q_tile, false, 0); + MMA16x32x16(dk_acc, 0, GGt_tile, false, 0, K_tile, false, 0); + + MM16x32x16(dW_raw, 0, ddelta_tile, false, 0, S_tile, false, kk / 16); + dW_raw = scale_rows(-dW_raw, row_exp); + + { + _M16x32 dKb_tile; + MM16x32x16(dKb_tile, 0, TWinv_tile, true, 0, dW_raw, false, 0); + KdKb = KdKb + dKb_tile * K_tile; + dk_acc = dk_acc + scale_rows(dKb_tile, beta_fm); + } + + MM16x32x16(t1, 0, delta_tile, false, 0, dS_tile, false, kk / 16); + t1 = scale_rows(t1, dec_exp); + dk_acc = dk_acc + t1; + + { + const _M16x32 KdKC = K_tile * t1; + dgam_row32 = dgam_row32 - KdKC; + STEEL_PRAGMA_UNROLL + for (short _i = 0; _i < decltype(KdKC)::kElemsPerTile; _i++) { + dgam_last += AT_NAX(KdKC, _i); + } + } + + STEEL_PRAGMA_UNROLL + for (short _i = 0; _i < mlx::steel::BaseNAXFrag::kElemsPerFrag; _i++) { + dgam_last += gamma_last_exp * + (S_tile.frag_at(0, kk / 16)[_i] * dS_tile.frag_at(0, kk / 16)[_i] + + S_tile.frag_at(0, kk / 16 + 1)[_i] * + dS_tile.frag_at(0, kk / 16 + 1)[_i]); + } + + { + _M16x32 TbK; + K_tile = scale_rows(K_tile, beta_fm); + MM16x32x16(TbK, 0, TWinv_tile, false, 0, K_tile, false, 0); + dgam_row32 = fmadd(TbK, dW_raw, dgam_row32); + } + + reduce_tile_tg( + dk_acc, + red_scratch, + dk_, + Hk * Dk, + kk, + valid_rows, + sg_id, + simd_lane_id); + } + + // dbeta = rowsum(V * dV_b) + rowsum(K * dK_b) + rowsum(tril_(dA) * KKt) + load_seq(V_tile, v_ + dv_idx, Dv * Hv); + + _M16x16 AKKt = dA_tile * KKt_tile; + AKKt = AKKt * tril_mask; + _M16x16 VdVb = fmadd(V_tile, TUddelta, AKKt); + row_sum(dbeta_acc, VdVb); + row_sum(dbeta_acc, KdKb); + + { + threadgroup float* db_st = db_stage + sg_id * C; + float _b0 = dbeta_acc[0]; + float _b1 = dbeta_acc[1]; + _b0 += simd_shuffle_xor(_b0, ushort(1)); + _b0 += simd_shuffle_xor(_b0, ushort(8)); + _b1 += simd_shuffle_xor(_b1, ushort(1)); + _b1 += simd_shuffle_xor(_b1, ushort(8)); + + if (fm < valid_rows) { + db_st[fm] = _b0; + } + if (fm1 < valid_rows) { + db_st[fm1] = _b1; + } + } + + dgam_row = reduce(dgam_row32); + + float dgam_acc[2] = {0.0f, 0.0f}; + dgam_row = dgam_row + dgam_pair; + row_sum(dgam_acc, dgam_row); + + _M16x16 cs_tile; + MM16x16x16(cs_tile, 0, dgam_pair, true, 0, Ones_tile, false, 0); + + const float dgam_last_red = simd_sum(dgam_last); + + { + threadgroup float* dg_st = dg_stage + sg_id * C; + float _d0 = dgam_acc[0]; + float _d1 = dgam_acc[1]; + _d0 += simd_shuffle_xor(_d0, ushort(1)); + _d0 += simd_shuffle_xor(_d0, ushort(8)); + _d1 += simd_shuffle_xor(_d1, ushort(1)); + _d1 += simd_shuffle_xor(_d1, ushort(8)); + + _d0 -= AT_NAX(cs_tile, 0); + _d1 -= AT_NAX(cs_tile, mlx::steel::BaseNAXFrag::kElemCols); + + if (fm == valid_rows - 1) { + _d0 += dgam_last_red; + } + if (fm < valid_rows) { + dg_st[fm] = _d0; + } + if (fm1 == valid_rows - 1) { + _d1 += dgam_last_red; + } + if (fm1 < valid_rows) { + dg_st[fm1] = _d1; + } + } + + threadgroup_barrier(mem_flags::mem_threadgroup); + if (sg_id == 0 && thread_index_in_simdgroup < (uint)valid_rows) { + float b_sum = 0.0f; + float g_sum = 0.0f; + STEEL_PRAGMA_UNROLL + for (short _s = 0; _s < kNSG; _s++) { + b_sum += db_stage[_s * C + thread_index_in_simdgroup]; + g_sum += dg_stage[_s * C + thread_index_in_simdgroup]; + } + const int idx = thread_index_in_simdgroup * Hv + hv_idx; + db_[idx] = b_sum; + dg_[idx] = g_sum; + } + + threadgroup_barrier(mem_flags::mem_threadgroup); + + // dS update + // dS = gamma_C * dS + dO.T @ (gamma * Q) - ddelta.T @ W + SCALE_NAX(dS_tile, gamma_last_exp); + for (int kk = 0; kk < Dk; kk += 32) { + _M16x32 term, W_raw; + load_seq(Q_tile, q_ + kk, Hk * Dk); + load_seq(K_tile, k_ + kk, Dk * Hk); + + // dO.T @ (gamma * Q) -> [dv, dk] + Q_tile = scale_rows(Q_tile, row_exp); + MM16x32x16(term, 0, dout_tile, true, 0, Q_tile, false, 0); + + // - ddelta.T @ W, with W = gamma * (TWinv @ beta*K) + K_tile = scale_rows(K_tile, beta_fm); + MM16x32x16(W_raw, 0, TWinv_tile, false, 0, K_tile, false, 0); + W_raw = scale_rows(-W_raw, row_exp); + MMA16x32x16(term, 0, ddelta_tile, true, 0, W_raw, false, 0); + + STEEL_PRAGMA_UNROLL + for (short _i = 0; _i < decltype(term)::kElemsPerFrag; _i++) { + dS_tile.frag_at(0, kk / 16)[_i] += AT_NAX(term, _i); + dS_tile.frag_at(0, kk / 16 + 1)[_i] += + AT_NAX(term, decltype(term)::kElemsPerFrag + _i); + } + } + }; + + // Walk segments in reverse. Each is replayed forward from its checkpoint to + // recover the entry states, then walked backwards. + // Backward over one segment: replay its entry states from the checkpoint, + // then walk it in reverse. HasTail is a compile-time tag so valid_rows is the + // constant C on every chunk of an interior segment, which folds away the + // row guards in the dq/dk stores and the load_rows switch in the replay. + auto process_segment = [&](const int seg, + auto tail_tag) __attribute__((always_inline)) { + constexpr bool HasTail = decltype(tail_tag)::value; + + const int seg_start = seg * Ckpt; + const int rem = n_chunks - seg_start; + const int n_in_seg = (rem < Ckpt) ? rem : Ckpt; + const int n_full = HasTail ? n_in_seg - 1 : n_in_seg; + + // delta is cached per chunk, so the replay is the state recurrence alone: + // S <- exp(gamma_C) S + K_dec^T delta + // No T^-1, no W, no U -- roughly a tenth of a full chunk's work, which is + // what makes trading Ckpt-fold memory for it worthwhile. + // + // Scoped so S is dead before process_chunk runs: it is a _M16xDk, and + // having it live alongside S_tile and dS_tile would put three of them in + // flight. + { + _M16xDk S; + S.load(ck_base + seg * Dv * Dk, Dk); + + auto replay_one = [&](const int j, + const short valid_rows, + auto bounded_tag) __attribute__((always_inline)) { + constexpr bool B = decltype(bounded_tag)::value; + + S.store(seg_base + j * Dv * Dk, Dk); + + const int c = seg_start + j; + auto k_c = k_base + c * C * Hk * Dk; + auto g_c = g_base + c * C * Hv; + auto cd_c = cd_base + c * C * Dv; + + float gv = (thread_index_in_simdgroup < (uint)valid_rows) + ? metal::fast::log( + metal::max( + g_c[thread_index_in_simdgroup * Hv + hv_idx], 1e-6)) + : 0.0f; + auto gs = simd_prefix_inclusive_sum(gv); + if (thread_index_in_simdgroup < C) { + gamma[thread_index_in_simdgroup] = static_cast(gs); + } + simdgroup_barrier(mem_flags::mem_threadgroup); + + const float g_last = gamma[C - 1]; + const float dec[2] = { + metal::fast::exp(g_last - gamma[fm]), + metal::fast::exp(g_last - gamma[fm1])}; + + _M16x16 delta_tile; + delta_tile.load(cd_c, Dv); + + S = S * metal::fast::exp(g_last); + for (int kk = 0; kk < Dk; kk += 32) { + _M16x32 Kd; + if constexpr (B) { + Kd.load_rows(k_c + kk, Dk * Hk, valid_rows); + } else { + Kd.load(k_c + kk, Dk * Hk); + } + Kd = scale_rows(Kd, dec); + MMA16x32x16(S, kk / 16, delta_tile, true, 0, Kd, false, 0); + } + }; + + for (int j = 0; j < n_full; j++) { + replay_one(j, C, metal::false_type{}); + } + if constexpr (HasTail) { + replay_one(n_in_seg - 1, tail, metal::true_type{}); + } + } + + auto set_chunk_ptrs = [&](const int j) __attribute__((always_inline)) { + const int c = seg_start + j; + q_ = q_base + c * C * Hk * Dk; + k_ = k_base + c * C * Hk * Dk; + dq_ = dq_base + c * C * Hk * Dk; + dk_ = dk_base + c * C * Hk * Dk; + v_ = v_base + c * C * Hv * Dv; + dv_ = dv_base + c * C * Hv * Dv; + co_ = co_base + c * C * Hv * Dv; + g_ = g_base + c * C * Hv; + beta_ = beta_base + c * C * Hv; + dg_ = dg_base + c * C * Hv; + db_ = db_base + c * C * Hv; + i_chunk_mats = cm_base + c * 3 * 256; + i_chunk_delta = cd_base + c * C * Dv; + i_seg = seg_base + j * Dv * Dk; + }; + + if constexpr (HasTail) { + set_chunk_ptrs(n_in_seg - 1); + process_chunk(tail, metal::true_type{}); + } + for (int j = n_full - 1; j >= 0; --j) { + set_chunk_ptrs(j); + process_chunk(C, metal::false_type{}); + } + }; + + // Only the last segment can hold a short chunk + int seg = n_ckpt - 1; + if (tail != C) { + process_segment(seg, metal::true_type{}); + --seg; + } + for (; seg >= 0; --seg) { + process_segment(seg, metal::false_type{}); + } + + dS_tile.store(o_dh, Dk); +} + +template +[[kernel]] void gated_delta_vjp_fused_nax1( + const device InT* q [[buffer(0)]], // [B, T, Hk, Dk] + const device InT* k [[buffer(1)]], // [B, T, Hk, Dk] + const device InT* v [[buffer(2)]], // [B, T, Hv, Dv] + const device InT* g [[buffer(3)]], // [B, T, Hv] + const device InT* beta [[buffer(4)]], // [B, T, Hv] + const device InT* cot_o [[buffer(5)]], // [B, T, Hv, Dv] + const device float* cot_h [[buffer(6)]], // [B, Hv, Dv, Dk] + const device float* state_cache [[buffer(7)]], // [B, Hv, n_chunks, Dv, Dk] + constant int& T [[buffer(8)]], + device mlx_atomic* dq [[buffer(9)]], + device mlx_atomic* dk [[buffer(10)]], + device float* dv [[buffer(11)]], + device float* dg [[buffer(12)]], + device float* db [[buffer(13)]], + device float* dh [[buffer(14)]], + // Cached, chunk-local forward intermediates produced by the forward-save + // pass (see gated_delta_fused_nax): avoids redoing the K@K^T contraction, + // the 15-step Neumann inversion, and the W/U/S contraction for delta. + const device float* chunk_mats [[buffer(15)]], + const device float* chunk_delta [[buffer(16)]], // [B, Hv, n_chunks, C, Dv] + uint3 thread_position_in_grid [[thread_position_in_grid]], + uint3 thread_position_in_threadgroup [[thread_position_in_threadgroup]], + uint thread_index_in_simdgroup [[thread_index_in_simdgroup]]) { + using _M16xDk = mlx::steel::NAXTile; + + auto n = thread_position_in_grid.z; + auto b_idx = n / Hv; + auto hv_idx = n % Hv; + auto hk_idx = hv_idx / (Hv / Hk); + + auto dv_idx = thread_position_in_grid.y * 16; + const short sg_id = thread_position_in_threadgroup.y; // 0..3 + + const ushort simd_lane_id = __metal_get_thread_index_in_simdgroup(ushort()); + const short qid = simd_lane_id >> 2; + const short fm = ((qid & 4) | ((simd_lane_id >> 1) & 3)); + + const int n_chunks = (T + C - 1) / C; + const int t_last = (n_chunks - 1) * C; + + // Pointers positioned at the final chunk + auto q_ = q + b_idx * T * Hk * Dk + hk_idx * Dk + t_last * Hk * Dk; + auto k_ = k + b_idx * T * Hk * Dk + hk_idx * Dk + t_last * Hk * Dk; + auto dq_ = dq + b_idx * T * Hk * Dk + hk_idx * Dk + t_last * Hk * Dk; + auto dk_ = dk + b_idx * T * Hk * Dk + hk_idx * Dk + t_last * Hk * Dk; + + auto v_ = v + b_idx * T * Hv * Dv + hv_idx * Dv + t_last * Hv * Dv; + auto dv_ = dv + b_idx * T * Hv * Dv + hv_idx * Dv + t_last * Hv * Dv; + auto co_ = cot_o + b_idx * T * Hv * Dv + hv_idx * Dv + t_last * Hv * Dv; + + auto g_ = g + b_idx * T * Hv + t_last * Hv; + auto beta_ = beta + b_idx * T * Hv + t_last * Hv; + auto dg_ = dg + b_idx * T * Hv + t_last * Hv; + auto db_ = db + b_idx * T * Hv + t_last * Hv; + + auto c_state = state_cache + (n * n_chunks * Dv + dv_idx) * Dk + + (n_chunks - 1) * Dv * Dk; + + // Cached forward intermediates, walked in reverse alongside c_state. + auto i_chunk_mats = + chunk_mats + n * n_chunks * 3 * 256 + (n_chunks - 1) * 3 * 256; + auto i_chunk_delta = + chunk_delta + n * n_chunks * C * Dv + dv_idx + (n_chunks - 1) * C * Dv; + + auto i_cot_h = cot_h + (n * Dv + dv_idx) * Dk; + auto o_dh = dh + (n * Dv + dv_idx) * Dk; + + // One threadgroup per (b, hv) now covers every Dv slice, so the Dv reduction + // never leaves threadgroup memory. + constexpr int kNSG = Dv / 16; + // Under GQA several threadgroups still share an hk, so dq/dk keep an atomic + // for that axis -- but with kNSG writers folded away first. + constexpr bool kGQA = (Hv != Hk); + + threadgroup float gamma_all[C * kNSG]; + threadgroup float* gamma = gamma_all + sg_id * C; + + // (kNSG/2) destinations x 32 lanes x kElemsPerTile(_M16x32) floats. + threadgroup float red_scratch[(kNSG / 1) * 32 * 16]; + threadgroup float db_stage[kNSG * C]; + threadgroup float dg_stage[kNSG * C]; + + float beta_fm[2]; + + // Carried state gradient, dL/dS. Same [dv, dk] orientation as S_tile. + _M16xDk dS_tile; + dS_tile.load(i_cot_h, Dk); + + // Forward state at chunk entry, reloaded from the checkpoint each chunk + _M16xDk S_tile; + + // Recomputed forward tiles + _M16x32 K_tile, Q_tile; + _M16xDk W_tile; + _M16x16 K16_tile, Q16_tile; + _M16x16 V_tile; + _M16x16 U_tile; + _M16x16 WS_tile; + _M16x16 delta_tile; + _M16x16 QKt_tile, QKt_raw; + _M16x16 KKt_tile; + _M16x16 TWinv_tile, TUinv_tile; + _M16x16 TMP_tile; + + // Backward tiles + _M16x16 dout_tile; + _M16x16 ddelta_tile; + _M16x16 dTU_tile, dTW_tile; + _M16x16 dTinv_tile; + _M16x16 dA_tile, G_tile; + _M16x16 tri_decay; + + _M16x16 I_tile; + STEEL_PRAGMA_UNROLL + for (short _i = 0; _i < decltype(I_tile)::kElemsPerFrag; _i++) { + const short2 _c = mlx::steel::BaseNAXFrag::get_coord(_i); /* {fn, fm} */ + AT_NAX(I_tile, _i) = (_c.x == _c.y) ? 1.0f : 0.0f; + } + + _M16x16 Ones_tile; + STEEL_PRAGMA_UNROLL + for (short _i = 0; _i < decltype(Ones_tile)::kElemsPerFrag; _i++) { + AT_NAX(Ones_tile, _i) = 1.0f; + } + + _M16x16 tril_mask; + STEEL_PRAGMA_UNROLL + for (short _i = 0; _i < decltype(tril_mask)::kElemsPerFrag; _i++) { + const short2 _c = mlx::steel::BaseNAXFrag::get_coord(_i); + AT_NAX(tril_mask, _i) = (_c.x >= _c.y) ? 0.0f : 1.0f; + } + + auto process_chunk = [&](const short valid_rows, + auto bounded_tag) __attribute__((always_inline)) { + constexpr bool B = decltype(bounded_tag)::value; + + auto load_seq = [&](thread auto& tile, auto src, int ld) { + if constexpr (B) { + tile.load_rows(src, ld, valid_rows); + } else { + tile.load(src, ld); + } + }; + + // reload the checkpoint + S_tile.load(c_state, Dk); + + float g_val = (thread_index_in_simdgroup < (uint)valid_rows) + ? metal::fast::log( + metal::max(g_[thread_index_in_simdgroup * Hv + hv_idx], 1e-6)) + : 0.0f; + + auto gamma_val = simd_prefix_inclusive_sum(g_val); + if (thread_index_in_simdgroup < C) { + gamma[thread_index_in_simdgroup] = static_cast(gamma_val); + } + simdgroup_barrier(mem_flags::mem_threadgroup); + + const float gamma_last = gamma[C - 1]; + const float gamma_last_exp = metal::fast::exp(gamma_last); + + beta_fm[0] = (fm < valid_rows) ? beta_[fm * Hv + hv_idx] : 0.0f; + const short fm1 = fm + mlx::steel::BaseNAXFrag::kElemRowsJump; + beta_fm[1] = (fm1 < valid_rows) ? beta_[fm1 * Hv + hv_idx] : 0.0f; + + // Per-lane decay factors. gamma[j] == gamma[valid_rows-1] for + // j >= valid_rows (the prefix sum of zeros), so the tail chunk needs no + // special case here. + const float row_exp[2] = { + metal::fast::exp(gamma[fm]), metal::fast::exp(gamma[fm1])}; + const float dec_exp[2] = { + metal::fast::exp(gamma_last - gamma[fm]), + metal::fast::exp(gamma_last - gamma[fm1])}; + + STEEL_PRAGMA_UNROLL + for (short _i = 0; _i < decltype(tri_decay)::kElemsPerFrag; _i++) { + const short2 _c = mlx::steel::BaseNAXFrag::get_coord(_i); // {fn, fm} + AT_NAX(tri_decay, _i) = + (_c.x > _c.y) ? 0.f : metal::fast::exp(gamma[_c.y] - gamma[_c.x]); + } + + KKt_tile.load(i_chunk_mats, 16); + TWinv_tile.load(i_chunk_mats + 256, 16); + QKt_raw.load(i_chunk_mats + 512, 16); + + QKt_tile = QKt_raw * tri_decay; + TUinv_tile = TWinv_tile * tri_decay; + + delta_tile.load(i_chunk_delta, Dv); + + // dgamma accumulators: + // dgam_row : per-row + // dgam_pair : row-minus-col + // dgam_last : the gamma_{C-1} + _M16x16 dgam_row, dgam_pair; + dgam_row.clear(); + dgam_pair.clear(); + float dgam_last = 0.0f; + float dbeta_acc[2] = {0.0f, 0.0f}; + + load_seq(dout_tile, co_ + dv_idx, Hv * Dv); + + // dODT = (dO @ delta.T * D) + _M16x16 dODT; + MM16x16x16(dODT, 0, dout_tile, false, 0, delta_tile, true, 0); + dODT = dODT * tri_decay; + + // cannot fuse this because ddelta is used by the next loop + MM16x16x16(ddelta_tile, 0, QKt_tile, true, 0, dout_tile, false, 0); + for (int kk = 0; kk < Dk; kk += 32) { + load_seq(K_tile, k_ + kk, Dk * Hk); + K_tile = scale_rows(K_tile, dec_exp); + MMA16x16x32(ddelta_tile, 0, K_tile, false, 0, dS_tile, true, kk / 16); + } + + // ddelta = M^T @ dO + K_dec @ dS^T + // dq = gamma * (dO @ S) + dODT @ K + dTW_tile.clear(); + for (int kk = 0; kk < Dk; kk += 32) { + _M16x32 dW_raw; + _M16x32 dq_acc; + load_seq(K_tile, k_ + kk, Dk * Hk); + + MM16x32x16(dW_raw, 0, ddelta_tile, false, 0, S_tile, false, kk / 16); + dW_raw = scale_rows(-dW_raw, row_exp); + _M16x32 Kb_tile = scale_rows(K_tile, beta_fm); + MMA16x16x32(dTW_tile, 0, dW_raw, false, 0, Kb_tile, true, 0); + + MM16x32x16(dq_acc, 0, dout_tile, false, 0, S_tile, false, kk / 16); + dq_acc = scale_rows(dq_acc, row_exp); + MMA16x32x16(dq_acc, 0, dODT, false, 0, K_tile, false, 0); + + // dq reduces over Dv, which is now entirely inside this threadgroup. + reduce_tile_tg( + dq_acc, + red_scratch, + dq_, + Hk * Dk, + kk, + valid_rows, + sg_id, + simd_lane_id); + } + + // dv = B(Tu.T @ ddelta). Indexed by dv_idx, so no reduction. + _M16x16 dVb_tile; + MM16x16x16(dVb_tile, 0, TUinv_tile, true, 0, ddelta_tile, false, 0); + _M16x16 TUddelta = dVb_tile; // pre-beta, reused by dbeta + dVb_tile = scale_rows(dVb_tile, beta_fm); + + dVb_tile.store_rows(dv_ + dv_idx, Hv * Dv, valid_rows); + + // dT_U = ddelta @ (beta * V).T + load_seq(V_tile, v_ + dv_idx, Dv * Hv); + V_tile = scale_rows(V_tile, beta_fm); + MM16x16x16(dTU_tile, 0, ddelta_tile, false, 0, V_tile, true, 0); + dTU_tile = dTU_tile * tri_decay; + + dTinv_tile = dTW_tile + dTU_tile; + + // dA = -T.T @ dTinv @ T.T + MM16x16x16(TMP_tile, 0, TWinv_tile, true, 0, dTinv_tile, false, 0); + MM16x16x16(dA_tile, 0, TMP_tile, false, 0, TWinv_tile, true, 0); + dA_tile = dA_tile * -1.0f; + + // G = tril_(dA) * beta, and GGt = G + G.T + G_tile = scale_rows(dA_tile, beta_fm) * tril_mask; + _M16x16 GGt_tile = G_tile; + MMA16x16x16(GGt_tile, 0, I_tile, true, 0, G_tile, true, 0); + + // dgamma pair sites + _M16x16 P_tile = QKt_raw * dODT; + _M16x16 R_tile = TWinv_tile * dTU_tile; + dgam_pair = P_tile + R_tile; + + _M16x32 KdKb; // rowsum(K * dK_b) source, needed for dbeta + KdKb.clear(); + _M16x32 dgam_row32; + dgam_row32.clear(); + for (int kk = 0; kk < Dk; kk += 32) { + _M16x32 dk_acc, dW_raw, t1; + + load_seq(Q_tile, q_ + kk, Hk * Dk); + load_seq(K_tile, k_ + kk, Dk * Hk); + + { + _M16x32 QgdOS; + MM16x32x16(QgdOS, 0, dout_tile, false, 0, S_tile, false, kk / 16); + dgam_row32 = dgam_row32 + scale_rows(QgdOS, row_exp) * Q_tile; + } + + MM16x32x16(dk_acc, 0, dODT, true, 0, Q_tile, false, 0); + MMA16x32x16(dk_acc, 0, GGt_tile, false, 0, K_tile, false, 0); + + MM16x32x16(dW_raw, 0, ddelta_tile, false, 0, S_tile, false, kk / 16); + dW_raw = scale_rows(-dW_raw, row_exp); + + { + _M16x32 dKb_tile; + MM16x32x16(dKb_tile, 0, TWinv_tile, true, 0, dW_raw, false, 0); + KdKb = KdKb + dKb_tile * K_tile; // KdKb_tmp folded away + dk_acc = dk_acc + scale_rows(dKb_tile, beta_fm); + } + + MM16x32x16(t1, 0, delta_tile, false, 0, dS_tile, false, kk / 16); + t1 = scale_rows(t1, dec_exp); + dk_acc = dk_acc + t1; + + { + const _M16x32 KdKC = K_tile * t1; + dgam_row32 = dgam_row32 - KdKC; + STEEL_PRAGMA_UNROLL + for (short _i = 0; _i < decltype(KdKC)::kElemsPerTile; _i++) { + dgam_last += AT_NAX(KdKC, _i); + } + } + + STEEL_PRAGMA_UNROLL + for (short _i = 0; _i < mlx::steel::BaseNAXFrag::kElemsPerFrag; _i++) { + dgam_last += gamma_last_exp * + (S_tile.frag_at(0, kk / 16)[_i] * dS_tile.frag_at(0, kk / 16)[_i] + + S_tile.frag_at(0, kk / 16 + 1)[_i] * + dS_tile.frag_at(0, kk / 16 + 1)[_i]); + } + + { + _M16x32 TbK; + K_tile = scale_rows(K_tile, beta_fm); + MM16x32x16(TbK, 0, TWinv_tile, false, 0, K_tile, false, 0); + dgam_row32 = fmadd(TbK, dW_raw, dgam_row32); + } + + reduce_tile_tg( + dk_acc, + red_scratch, + dk_, + Hk * Dk, + kk, + valid_rows, + sg_id, + simd_lane_id); + } + + // dbeta = rowsum(V * dV_b) + rowsum(K * dK_b) + rowsum(tril_(dA) * KKt) + + load_seq(V_tile, v_ + dv_idx, Dv * Hv); + + _M16x16 AKKt = dA_tile * KKt_tile; + AKKt = AKKt * tril_mask; + _M16x16 VdVb = fmadd(V_tile, TUddelta, AKKt); + row_sum(dbeta_acc, VdVb); + row_sum(dbeta_acc, KdKb); + + { + threadgroup float* db_st = db_stage + sg_id * C; + float _b0 = dbeta_acc[0]; + float _b1 = dbeta_acc[1]; + _b0 += simd_shuffle_xor(_b0, ushort(1)); + _b0 += simd_shuffle_xor(_b0, ushort(8)); + _b1 += simd_shuffle_xor(_b1, ushort(1)); + _b1 += simd_shuffle_xor(_b1, ushort(8)); + + const short _r0 = fm; + const short _r1 = fm + mlx::steel::BaseNAXFrag::kElemRowsJump; + if (_r0 < valid_rows) { + db_st[_r0] = _b0; + } + if (_r1 < valid_rows) { + db_st[_r1] = _b1; + } + } + + dgam_row = reduce(dgam_row32); + + float dgam_acc[2] = {0.0f, 0.0f}; + dgam_row = dgam_row + dgam_pair; + row_sum(dgam_acc, dgam_row); + + _M16x16 cs_tile; + MM16x16x16(cs_tile, 0, dgam_pair, true, 0, Ones_tile, false, 0); + + const float dgam_last_red = simd_sum(dgam_last); + + { + threadgroup float* dg_st = dg_stage + sg_id * C; + float _d0 = dgam_acc[0]; + float _d1 = dgam_acc[1]; + _d0 += simd_shuffle_xor(_d0, ushort(1)); + _d0 += simd_shuffle_xor(_d0, ushort(8)); + _d1 += simd_shuffle_xor(_d1, ushort(1)); + _d1 += simd_shuffle_xor(_d1, ushort(8)); + + const short _r0 = fm; + const short _r1 = fm + mlx::steel::BaseNAXFrag::kElemRowsJump; + + _d0 -= AT_NAX(cs_tile, 0); + _d1 -= AT_NAX(cs_tile, mlx::steel::BaseNAXFrag::kElemCols); + + if (_r0 == valid_rows - 1) { + _d0 += dgam_last_red; + } + if (_r0 < valid_rows) { + dg_st[_r0] = _d0; + } + if (_r1 == valid_rows - 1) { + _d1 += dgam_last_red; + } + if (_r1 < valid_rows) { + dg_st[_r1] = _d1; + } + } + + threadgroup_barrier(mem_flags::mem_threadgroup); + if (sg_id == 0 && thread_index_in_simdgroup < (uint)valid_rows) { + float b_sum = 0.0f; + float g_sum = 0.0f; + STEEL_PRAGMA_UNROLL + for (short _s = 0; _s < kNSG; _s++) { + b_sum += db_stage[_s * C + thread_index_in_simdgroup]; + g_sum += dg_stage[_s * C + thread_index_in_simdgroup]; + } + const int idx = thread_index_in_simdgroup * Hv + hv_idx; + db_[idx] = b_sum; + dg_[idx] = g_sum; + } + + threadgroup_barrier(mem_flags::mem_threadgroup); + + // dS update + // dS = gamma_C * dS + dO.T @ (gamma * Q) - ddelta.T @ W + SCALE_NAX(dS_tile, gamma_last_exp); + for (int kk = 0; kk < Dk; kk += 32) { + _M16x32 term, W_raw; + load_seq(Q_tile, q_ + kk, Hk * Dk); + load_seq(K_tile, k_ + kk, Dk * Hk); + + // dO.T @ (gamma * Q) -> [dv, dk] + Q_tile = scale_rows(Q_tile, row_exp); + MM16x32x16(term, 0, dout_tile, true, 0, Q_tile, false, 0); + + // - ddelta.T @ W, with W = gamma * (TWinv @ beta*K) + K_tile = scale_rows(K_tile, beta_fm); + MM16x32x16(W_raw, 0, TWinv_tile, false, 0, K_tile, false, 0); + W_raw = scale_rows(-W_raw, row_exp); + MMA16x32x16(term, 0, ddelta_tile, true, 0, W_raw, false, 0); + + STEEL_PRAGMA_UNROLL + for (short _i = 0; _i < decltype(term)::kElemsPerFrag; _i++) { + dS_tile.frag_at(0, kk / 16)[_i] += AT_NAX(term, _i); + dS_tile.frag_at(0, kk / 16 + 1)[_i] += + AT_NAX(term, decltype(term)::kElemsPerFrag + _i); + } + } + }; + + // Walk chunks in reverse. The tail chunk comes first. + int c = n_chunks - 1; + const short tail = short(T - c * C); + if (tail != C) { + process_chunk(tail, metal::true_type{}); + q_ -= C * Hk * Dk; + k_ -= C * Hk * Dk; + v_ -= C * Hv * Dv; + co_ -= C * Hv * Dv; + g_ -= C * Hv; + beta_ -= C * Hv; + dq_ -= C * Hk * Dk; + dk_ -= C * Hk * Dk; + dv_ -= C * Hv * Dv; + dg_ -= C * Hv; + db_ -= C * Hv; + c_state -= Dv * Dk; + i_chunk_mats -= 3 * 256; + i_chunk_delta -= C * Dv; + --c; + } + for (; c >= 0; --c) { + process_chunk(C, metal::false_type{}); + q_ -= C * Hk * Dk; + k_ -= C * Hk * Dk; + v_ -= C * Hv * Dv; + co_ -= C * Hv * Dv; + g_ -= C * Hv; + beta_ -= C * Hv; + dq_ -= C * Hk * Dk; + dk_ -= C * Hk * Dk; + dv_ -= C * Hv * Dv; + dg_ -= C * Hv; + db_ -= C * Hv; + c_state -= Dv * Dk; + i_chunk_mats -= 3 * 256; + i_chunk_delta -= C * Dv; + } + + dS_tile.store(o_dh, Dk); +} + +template +METAL_FUNC void reduce_tile_tg0( + thread TileT& acc, + threadgroup float* scratch, + const short sg_id, + const ushort simd_lane_id) { + constexpr short kE = TileT::kElemsPerTile; + + STEEL_PRAGMA_UNROLL + for (short lo = kNSG / 2; lo > 0; lo >>= 1) { + if (sg_id >= lo && sg_id < 2 * lo) { + const short dst = sg_id - lo; + STEEL_PRAGMA_UNROLL + for (short _i = 0; _i < kE; _i++) { + scratch[(dst * 32 + simd_lane_id) * kE + _i] = AT_NAX(acc, _i); + } + } + threadgroup_barrier(mem_flags::mem_threadgroup); + if (sg_id < lo) { + STEEL_PRAGMA_UNROLL + for (short _i = 0; _i < kE; _i++) { + AT_NAX(acc, _i) += scratch[(sg_id * 32 + simd_lane_id) * kE + _i]; + } + } + threadgroup_barrier(mem_flags::mem_threadgroup); + } +} + +template +[[kernel]] void gated_delta_vjp_fused_nax0( + const device InT* q [[buffer(0)]], // [B, T, Hk, Dk] + const device InT* k [[buffer(1)]], // [B, T, Hk, Dk] + const device InT* v [[buffer(2)]], // [B, T, Hv, Dv] + const device InT* g [[buffer(3)]], // [B, T, Hv] + const device InT* beta [[buffer(4)]], // [B, T, Hv] + const device InT* cot_o [[buffer(5)]], // [B, T, Hv, Dv] + const device float* cot_h [[buffer(6)]], // [B, Hv, Dv, Dk] + const device float* state_cache [[buffer(7)]], // [B, Hv, n_chunks, Dv, Dk] + constant int& T [[buffer(8)]], + device mlx_atomic* dq [[buffer(9)]], + device mlx_atomic* dk [[buffer(10)]], + device float* dv [[buffer(11)]], + device float* dg [[buffer(12)]], + device float* db [[buffer(13)]], + device float* dh [[buffer(14)]], + // Cached, chunk-local forward intermediates produced by the forward-save + // pass (see gated_delta_fused_nax): avoids redoing the K@K^T contraction, + // the 15-step Neumann inversion, and the W/U/S contraction for delta. + const device float* chunk_mats [[buffer(15)]], + const device float* chunk_delta [[buffer(16)]], // [B, Hv, n_chunks, C, Dv] + uint3 thread_position_in_grid [[thread_position_in_grid]], + uint3 thread_position_in_threadgroup [[thread_position_in_threadgroup]], + uint thread_index_in_simdgroup [[thread_index_in_simdgroup]]) { + using _M16xDk = mlx::steel::NAXTile; + + auto n = thread_position_in_grid.z; + auto b_idx = n / Hv; + auto hv_idx = n % Hv; + auto hk_idx = hv_idx / (Hv / Hk); + + auto dv_idx = thread_position_in_grid.y * 16; + const short sg_id = thread_position_in_threadgroup.y; // 0..3 + + const ushort simd_lane_id = __metal_get_thread_index_in_simdgroup(ushort()); + const short qid = simd_lane_id >> 2; + const short fm = ((qid & 4) | ((simd_lane_id >> 1) & 3)); + + const int n_chunks = (T + C - 1) / C; + const int t_last = (n_chunks - 1) * C; + + // Pointers positioned at the final chunk + auto q_ = q + b_idx * T * Hk * Dk + hk_idx * Dk + t_last * Hk * Dk; + auto k_ = k + b_idx * T * Hk * Dk + hk_idx * Dk + t_last * Hk * Dk; + auto dq_ = dq + b_idx * T * Hk * Dk + hk_idx * Dk + t_last * Hk * Dk; + auto dk_ = dk + b_idx * T * Hk * Dk + hk_idx * Dk + t_last * Hk * Dk; + + auto v_ = v + b_idx * T * Hv * Dv + hv_idx * Dv + t_last * Hv * Dv; + auto dv_ = dv + b_idx * T * Hv * Dv + hv_idx * Dv + t_last * Hv * Dv; + auto co_ = cot_o + b_idx * T * Hv * Dv + hv_idx * Dv + t_last * Hv * Dv; + + auto g_ = g + b_idx * T * Hv + t_last * Hv; + auto beta_ = beta + b_idx * T * Hv + t_last * Hv; + auto dg_ = dg + b_idx * T * Hv + t_last * Hv; + auto db_ = db + b_idx * T * Hv + t_last * Hv; + + auto c_state = state_cache + (n * n_chunks * Dv + dv_idx) * Dk + + (n_chunks - 1) * Dv * Dk; + + // Cached forward intermediates, walked in reverse alongside c_state. + auto i_chunk_mats = + chunk_mats + n * n_chunks * 3 * 256 + (n_chunks - 1) * 3 * 256; + auto i_chunk_delta = + chunk_delta + n * n_chunks * C * Dv + dv_idx + (n_chunks - 1) * C * Dv; + + auto i_cot_h = cot_h + (n * Dv + dv_idx) * Dk; + auto o_dh = dh + (n * Dv + dv_idx) * Dk; + + // One threadgroup per (b, hv) now covers every Dv slice, so the Dv reduction + // never leaves threadgroup memory. + constexpr int kNSG = Dv / 16; + // Under GQA several threadgroups still share an hk, so dq/dk keep an atomic + // for that axis -- but with kNSG writers folded away first. + constexpr bool kGQA = (Hv != Hk); + + threadgroup float gamma_all[C * kNSG]; + threadgroup float* gamma = gamma_all + sg_id * C; + + // (kNSG/2) destinations x 32 lanes x kElemsPerTile(_M16x32) floats. + threadgroup float red_scratch[(kNSG / 2) * 32 * 16]; + threadgroup float db_stage[kNSG * C]; + threadgroup float dg_stage[kNSG * C]; + + float beta_fm[2]; + + // Carried state gradient, dL/dS. Same [dv, dk] orientation as S_tile. + _M16xDk dS_tile; + dS_tile.load(i_cot_h, Dk); + + // Forward state at chunk entry, reloaded from the checkpoint each chunk + _M16xDk S_tile; + + // Recomputed forward tiles + _M16x32 K_tile, Q_tile; + _M16xDk W_tile; + _M16x16 K16_tile, Q16_tile; + _M16x16 V_tile; + _M16x16 U_tile; + _M16x16 WS_tile; + _M16x16 delta_tile; + _M16x16 QKt_tile, QKt_raw; + _M16x16 KKt_tile; + _M16x16 TWinv_tile, TUinv_tile; + _M16x16 TMP_tile; + + // Backward tiles + _M16x16 dout_tile; + _M16x16 ddelta_tile; + _M16x16 dTU_tile, dTW_tile; + _M16x16 dTinv_tile; + _M16x16 dA_tile, G_tile; + _M16x16 tri_decay; + + _M16x16 I_tile; + STEEL_PRAGMA_UNROLL + for (short _i = 0; _i < decltype(I_tile)::kElemsPerFrag; _i++) { + const short2 _c = mlx::steel::BaseNAXFrag::get_coord(_i); /* {fn, fm} */ + AT_NAX(I_tile, _i) = (_c.x == _c.y) ? 1.0f : 0.0f; + } + + _M16x16 NI_tile; + STEEL_PRAGMA_UNROLL + for (short _i = 0; _i < decltype(NI_tile)::kElemsPerFrag; _i++) { + const short2 _c = mlx::steel::BaseNAXFrag::get_coord(_i); /* {fn, fm} */ + AT_NAX(NI_tile, _i) = (_c.x == _c.y) ? -1.0f : 0.0f; + } + + _M16x16 Ones_tile; + STEEL_PRAGMA_UNROLL + for (short _i = 0; _i < decltype(Ones_tile)::kElemsPerFrag; _i++) { + AT_NAX(Ones_tile, _i) = 1.0f; + } + + auto process_chunk = [&](const short valid_rows, + auto bounded_tag) __attribute__((always_inline)) { + constexpr bool B = decltype(bounded_tag)::value; + + auto load_seq = [&](thread auto& tile, auto src, int ld) { + if constexpr (B) { + tile.load_rows(src, ld, valid_rows); + } else { + tile.load(src, ld); + } + }; + + // reload the checkpoint + S_tile.load(c_state, Dk); + + float g_val = (thread_index_in_simdgroup < (uint)valid_rows) + ? metal::fast::log( + metal::max(g_[thread_index_in_simdgroup * Hv + hv_idx], 1e-6)) + : 0.0f; + + auto gamma_val = simd_prefix_inclusive_sum(g_val); + if (thread_index_in_simdgroup < C) { + gamma[thread_index_in_simdgroup] = static_cast(gamma_val); + } + simdgroup_barrier(mem_flags::mem_threadgroup); + + const float gamma_last = gamma[C - 1]; + const float gamma_last_exp = metal::fast::exp(gamma_last); + + beta_fm[0] = (fm < valid_rows) ? beta_[fm * Hv + hv_idx] : 0.0f; + const short fm1 = fm + mlx::steel::BaseNAXFrag::kElemRowsJump; + beta_fm[1] = (fm1 < valid_rows) ? beta_[fm1 * Hv + hv_idx] : 0.0f; + + // Per-lane decay factors. gamma[j] == gamma[valid_rows-1] for + // j >= valid_rows (the prefix sum of zeros), so the tail chunk needs no + // special case here. + const float row_exp[2] = { + metal::fast::exp(gamma[fm]), metal::fast::exp(gamma[fm1])}; + const float dec_exp[2] = { + metal::fast::exp(gamma_last - gamma[fm]), + metal::fast::exp(gamma_last - gamma[fm1])}; + + STEEL_PRAGMA_UNROLL + for (short _i = 0; _i < decltype(tri_decay)::kElemsPerFrag; _i++) { + const short2 _c = mlx::steel::BaseNAXFrag::get_coord(_i); // {fn, fm} + AT_NAX(tri_decay, _i) = + (_c.x > _c.y) ? 0.f : metal::fast::exp(gamma[_c.y] - gamma[_c.x]); + } + + KKt_tile.load(i_chunk_mats, 16); + TWinv_tile.load(i_chunk_mats + 256, 16); + QKt_raw.load(i_chunk_mats + 512, 16); + + QKt_tile = QKt_raw; + MUL_NAX(QKt_tile, QKt_tile, tri_decay); + + TUinv_tile = TWinv_tile; + MUL_NAX(TUinv_tile, TUinv_tile, tri_decay); + + delta_tile.load(i_chunk_delta, Dv); + + // dgamma accumulators: + // dgam_row : per-row + // dgam_pair : row-minus-col + // dgam_last : the gamma_{C-1} + _M16x16 dgam_row, dgam_pair; + dgam_row.clear(); + dgam_pair.clear(); + float dgam_last = 0.0f; + float dbeta_acc[2] = {0.0f, 0.0f}; + + load_seq(dout_tile, co_ + dv_idx, Hv * Dv); + + // dODT = (dO @ delta.T * D) + _M16x16 dODT; + MM16x16x16(dODT, 0, dout_tile, false, 0, delta_tile, true, 0); + MUL_NAX(dODT, dODT, tri_decay); + + // cannot fuse this because ddelta is used by the next loop + MM16x16x16(ddelta_tile, 0, QKt_tile, true, 0, dout_tile, false, 0); + for (int kk = 0; kk < Dk; kk += 32) { + load_seq(K_tile, k_ + kk, Dk * Hk); + SCALE2_P(K_tile, dec_exp); + MMA16x16x32(ddelta_tile, 0, K_tile, false, 0, dS_tile, true, kk / 16); + } + + // ddelta = M^T @ dO + K_dec @ dS^T + // dq = gamma * (dO @ S) + dODT @ K + dTW_tile.clear(); + for (int kk = 0; kk < Dk; kk += 32) { + _M16x32 dW_raw; + _M16x32 dq_acc; + _M16x32 Kb_tile; + load_seq(K_tile, k_ + kk, Dk * Hk); + + MM16x32x16(dW_raw, 0, ddelta_tile, false, 0, S_tile, false, kk / 16); + NSCALE_ROW_P(dW_raw, row_exp); + + SCALE_BETA_NAX_O(Kb_tile, K_tile, beta_fm); + MMA16x16x32(dTW_tile, 0, dW_raw, false, 0, Kb_tile, true, 0); + + MM16x32x16(dq_acc, 0, dout_tile, false, 0, S_tile, false, kk / 16); + SCALE_ROW_P(dq_acc, row_exp); + MMA16x32x16(dq_acc, 0, dODT, false, 0, K_tile, false, 0); + + // dq reduces over Dv, which is now entirely inside this threadgroup. + reduce_tile_tg(dq_acc, red_scratch, sg_id, simd_lane_id); + + if (sg_id == 0) { + STEEL_PRAGMA_UNROLL + for (short _i = 0; _i < decltype(dq_acc)::kElemsPerTile; _i++) { + const short _f = _i / mlx::steel::BaseNAXFrag::kElemsPerFrag; + const short _w = _i % mlx::steel::BaseNAXFrag::kElemsPerFrag; + const short2 _c = mlx::steel::BaseNAXFrag::get_coord(_w); // {fn, fm} + const short _fn = _c.x + _f * 16; + const short _fm = _c.y; + if (_fm < valid_rows) { + const int idx = _fm * Hk * Dk + kk + _fn; + if (kGQA) { + mlx_atomic_fetch_add_explicit(dq_, AT_NAX(dq_acc, _i), idx); + } else { + // Dv reduction done in threadgroup memory: single writer. + mlx_atomic_store_explicit(dq_, AT_NAX(dq_acc, _i), idx); + } + } + } + } + } + + // dv = B(Tu.T @ ddelta). Indexed by dv_idx, so no reduction. + _M16x16 dVb_tile; + MM16x16x16(dVb_tile, 0, TUinv_tile, true, 0, ddelta_tile, false, 0); + _M16x16 TUddelta = dVb_tile; // pre-beta, reused by dbeta + SCALE_BETA_NAX(dVb_tile, beta_fm); + + STEEL_PRAGMA_UNROLL + for (short _i = 0; _i < decltype(dVb_tile)::kElemsPerFrag; _i++) { + const short2 _c = mlx::steel::BaseNAXFrag::get_coord(_i); // {fn, fm} + if (_c.y < valid_rows) { + dv_[_c.y * Hv * Dv + dv_idx + _c.x] = AT_NAX(dVb_tile, _i); + } + } + + // dT_U = ddelta @ (beta * V).T + load_seq(V_tile, v_ + dv_idx, Dv * Hv); + SCALE_BETA_NAX(V_tile, beta_fm); + MM16x16x16(dTU_tile, 0, ddelta_tile, false, 0, V_tile, true, 0); + MUL_NAX(dTU_tile, dTU_tile, tri_decay); // dTinv = dT_W + (dT_U * D). + + ADD_NAX(dTinv_tile, dTW_tile, dTU_tile); + + // dA = -T.T @ dTinv @ T.T + MM16x16x16(TMP_tile, 0, TWinv_tile, true, 0, dTinv_tile, false, 0); + MM16x16x16(dA_tile, 0, TMP_tile, false, 0, TWinv_tile, true, 0); + SCALE_NAX(dA_tile, -1.0f); + + // G = tril_(dA) * beta, and GGt = G + G.T + G_tile = dA_tile; + SCALE_TRIEQ_NAX1(G_tile, beta_fm); + _M16x16 GGt_tile; + GGt_tile = G_tile; + MMA16x16x16(GGt_tile, 0, I_tile, true, 0, G_tile, true, 0); + + // dgamma pair sites + _M16x16 P_tile; + _M16x16 R_tile; + + MUL_NAX(P_tile, QKt_raw, dODT); + MUL_NAX(R_tile, TWinv_tile, dTU_tile); + ADD_NAX(dgam_pair, P_tile, R_tile); + + _M16x32 KdKb; // rowsum(K * dK_b) source, needed for dbeta + KdKb.clear(); + _M16x32 dgam_row32; + dgam_row32.clear(); + for (int kk = 0; kk < Dk; kk += 32) { + _M16x32 dk_acc, dW_raw, dKb_tile, t1, KdKb_tmp; + _M16x32 QgdOS, TbK, KdKC_blk; + + load_seq(Q_tile, q_ + kk, Hk * Dk); // [16 x 32] + load_seq(K_tile, k_ + kk, Dk * Hk); // [16 x 32] + + // gamma: + rowsum(Q * gamma*(dO @ S)) + MM16x32x16(QgdOS, 0, dout_tile, false, 0, S_tile, false, kk / 16); + SCALE_ROW_P(QgdOS, row_exp); + MULA_NAX(dgam_row32, QgdOS, Q_tile); + + // dODT.T @ Q and (G + G.T) @ K + MM16x32x16(dk_acc, 0, dODT, true, 0, Q_tile, false, 0); + MMA16x32x16(dk_acc, 0, GGt_tile, false, 0, K_tile, false, 0); + + // B (T.T @ dW_raw), dW_raw = -gamma * (ddelta @ S) + MM16x32x16(dW_raw, 0, ddelta_tile, false, 0, S_tile, false, kk / 16); + NSCALE_ROW_P(dW_raw, row_exp); + MM16x32x16(dKb_tile, 0, TWinv_tile, true, 0, dW_raw, false, 0); + + // dbeta term uses dK_b before the beta scaling, with raw K + MUL_NAX(KdKb_tmp, dKb_tile, K_tile); + ADD_NAX(KdKb, KdKb, KdKb_tmp); + + SCALE_BETA_NAX(dKb_tile, beta_fm); + ADD_NAX(dk_acc, dk_acc, dKb_tile); + + // D_C * (delta @ dS) + MM16x32x16(t1, 0, delta_tile, false, 0, dS_tile, false, kk / 16); + SCALE2_P(t1, dec_exp); + ADD_NAX(dk_acc, dk_acc, t1); + + // gamma stuff + MUL_NAX(KdKC_blk, K_tile, t1); + SUB_NAX(dgam_row32, dgam_row32, KdKC_blk); + STEEL_PRAGMA_UNROLL + for (short _i = 0; _i < decltype(KdKC_blk)::kElemsPerTile; _i++) { + dgam_last += AT_NAX(KdKC_blk, _i); + } + + // gamma stuff: gamma_C * sum(S * dS) + STEEL_PRAGMA_UNROLL + for (short _i = 0; _i < mlx::steel::BaseNAXFrag::kElemsPerFrag; _i++) { + dgam_last += gamma_last_exp * S_tile.frag_at(0, kk / 16)[_i] * + dS_tile.frag_at(0, kk / 16)[_i]; + dgam_last += gamma_last_exp * S_tile.frag_at(0, kk / 16 + 1)[_i] * + dS_tile.frag_at(0, kk / 16 + 1)[_i]; + } + + // gamma stuff + SCALE_BETA_NAX(K_tile, beta_fm); + SCALE_NAX(dW_raw, -1.0f); + MM16x32x16(TbK, 0, TWinv_tile, false, 0, K_tile, false, 0); + MULS_NAX(dgam_row32, TbK, dW_raw); + + // dk reduces over Dv the same way dq does. + reduce_tile_tg(dk_acc, red_scratch, sg_id, simd_lane_id); + + if (sg_id == 0) { + STEEL_PRAGMA_UNROLL + for (short _i = 0; _i < decltype(dk_acc)::kElemsPerTile; _i++) { + const short _f = _i / mlx::steel::BaseNAXFrag::kElemsPerFrag; + const short _w = _i % mlx::steel::BaseNAXFrag::kElemsPerFrag; + const short2 _c = mlx::steel::BaseNAXFrag::get_coord(_w); // {fn, fm} + if (_c.y < valid_rows) { + const int idx = _c.y * Hk * Dk + kk + _c.x + _f * 16; + if (kGQA) { + mlx_atomic_fetch_add_explicit(dk_, AT_NAX(dk_acc, _i), idx); + } else { + mlx_atomic_store_explicit(dk_, AT_NAX(dk_acc, _i), idx); + } + } + } + } + } + + // dbeta = rowsum(V * dV_b) + rowsum(K * dK_b) + rowsum(tril_(dA) * KKt) + _M16x16 VdVb; + load_seq(V_tile, v_ + dv_idx, Dv * Hv); + MUL_NAX(VdVb, V_tile, TUddelta); + + _M16x16 AKKt; + TRIL_NAX(AKKt, dA_tile); + MUL_NAX(AKKt, AKKt, KKt_tile); + + ADD_NAX(VdVb, VdVb, AKKt); + ROWSUM1_NAX(dbeta_acc, VdVb); + ROWSUM1_NAX(dbeta_acc, KdKb); + + { + threadgroup float* db_st = db_stage + sg_id * C; + STEEL_PRAGMA_UNROLL + for (short _r = 0; _r < 2; _r++) { + float _b = dbeta_acc[_r]; + _b += simd_shuffle_xor(_b, ushort(1)); + _b += simd_shuffle_xor(_b, ushort(8)); + + const short _row = fm + _r * mlx::steel::BaseNAXFrag::kElemRowsJump; + if (_row < valid_rows) { + db_st[_row] = _b; + } + } + } + + // Fold the two fragments: only the row sum of dgam_row is used downstream, + // and row sums are additive across fragments. + STEEL_PRAGMA_UNROLL + for (short _i = 0; _i < mlx::steel::BaseNAXFrag::kElemsPerFrag; _i++) { + AT_NAX(dgam_row, _i) = AT_NAX(dgam_row32, _i) + + AT_NAX(dgam_row32, mlx::steel::BaseNAXFrag::kElemsPerFrag + _i); + } + + // reduce and store dgamma: rowsum(row + pair) - colsum(pair) + ADD_NAX(dgam_row, dgam_row, dgam_pair); + _M16x16 dgam_tile, cs_tile; + MM16x16x16(dgam_tile, 0, dgam_row, false, 0, Ones_tile, false, 0); + MM16x16x16(cs_tile, 0, dgam_pair, true, 0, Ones_tile, false, 0); + SUB_NAX(dgam_tile, dgam_tile, cs_tile); + + { + threadgroup float* dg_st = dg_stage + sg_id * C; + STEEL_PRAGMA_UNROLL + for (short _i = 0; _i < decltype(dgam_tile)::kElemsPerFrag; _i++) { + const short2 _c = mlx::steel::BaseNAXFrag::get_coord(_i); // {fn, fm} + if (_c.x == 0) { + dg_st[_c.y] = AT_NAX(dgam_tile, _i); + } + } + // The gamma_{C-1}-only term lands on the same slot the loop just wrote, + // so the lanes must be ordered before lane 0 accumulates into it. + simdgroup_barrier(mem_flags::mem_threadgroup); + const float dgam_last_red = simd_sum(dgam_last); + if (simd_lane_id == 0 && valid_rows > 0) { + dg_st[valid_rows - 1] += dgam_last_red; + } + } + + // One barrier serves both stages; simdgroup 0 commits. + threadgroup_barrier(mem_flags::mem_threadgroup); + if (sg_id == 0 && thread_index_in_simdgroup < (uint)valid_rows) { + float b_sum = 0.0f; + float g_sum = 0.0f; + STEEL_PRAGMA_UNROLL + for (short _s = 0; _s < kNSG; _s++) { + b_sum += db_stage[_s * C + thread_index_in_simdgroup]; + g_sum += dg_stage[_s * C + thread_index_in_simdgroup]; + } + const int idx = thread_index_in_simdgroup * Hv + hv_idx; + db_[idx] = b_sum; + dg_[idx] = g_sum; + } + + threadgroup_barrier(mem_flags::mem_threadgroup); + + // dS update + // dS = gamma_C * dS + dO.T @ (gamma * Q) - ddelta.T @ W + SCALE_NAX(dS_tile, gamma_last_exp); + for (int kk = 0; kk < Dk; kk += 32) { + _M16x32 term, W_raw; + load_seq(Q_tile, q_ + kk, Hk * Dk); + load_seq(K_tile, k_ + kk, Dk * Hk); + + // dO.T @ (gamma * Q) -> [dv, dk] + SCALE_ROW_P(Q_tile, row_exp); + MM16x32x16(term, 0, dout_tile, true, 0, Q_tile, false, 0); + + // - ddelta.T @ W, with W = gamma * (TWinv @ beta*K) + SCALE_BETA_NAX(K_tile, beta_fm); + MM16x32x16(W_raw, 0, TWinv_tile, false, 0, K_tile, false, 0); + NSCALE_ROW_P(W_raw, row_exp); + MMA16x32x16(term, 0, ddelta_tile, true, 0, W_raw, false, 0); + + STEEL_PRAGMA_UNROLL + for (short _i = 0; _i < decltype(term)::kElemsPerFrag; _i++) { + dS_tile.frag_at(0, kk / 16)[_i] += AT_NAX(term, _i); + dS_tile.frag_at(0, kk / 16 + 1)[_i] += + AT_NAX(term, decltype(term)::kElemsPerFrag + _i); + } + } + }; + + // Walk chunks in reverse. The tail chunk comes first. + int c = n_chunks - 1; + const short tail = short(T - c * C); + if (tail != C) { + process_chunk(tail, metal::true_type{}); + q_ -= C * Hk * Dk; + k_ -= C * Hk * Dk; + v_ -= C * Hv * Dv; + co_ -= C * Hv * Dv; + g_ -= C * Hv; + beta_ -= C * Hv; + dq_ -= C * Hk * Dk; + dk_ -= C * Hk * Dk; + dv_ -= C * Hv * Dv; + dg_ -= C * Hv; + db_ -= C * Hv; + c_state -= Dv * Dk; + i_chunk_mats -= 3 * 256; + i_chunk_delta -= C * Dv; + --c; + } + for (; c >= 0; --c) { + process_chunk(C, metal::false_type{}); + q_ -= C * Hk * Dk; + k_ -= C * Hk * Dk; + v_ -= C * Hv * Dv; + co_ -= C * Hv * Dv; + g_ -= C * Hv; + beta_ -= C * Hv; + dq_ -= C * Hk * Dk; + dk_ -= C * Hk * Dk; + dv_ -= C * Hv * Dv; + dg_ -= C * Hv; + db_ -= C * Hv; + c_state -= Dv * Dk; + i_chunk_mats -= 3 * 256; + i_chunk_delta -= C * Dv; + } + + dS_tile.store(o_dh, Dk); +} + +// Postprocessing dg. Similar idea to what is done in the triton kernel. +// The kernel computes the gradient with respect to the log. +template +[[kernel]] void gated_delta_dgamma_to_dg( + const device InT* g [[buffer(0)]], + device float* dg [[buffer(1)]], // in: dL/dgamma, out: dL/dg + constant int& T [[buffer(2)]], + constant int& Hv [[buffer(3)]], + constant int& n_total [[buffer(4)]], // B * Hv + uint2 pos [[thread_position_in_grid]]) { + const int n = int(pos.x); // b * Hv + hv + const int c = int(pos.y); // chunk index + + if (n >= n_total) { + return; + } + + const int t0 = c * C; + if (t0 >= T) { + return; + } + const int len = min(C, T - t0); + + const int b_idx = n / Hv; + const int hv_idx = n % Hv; + const int base = b_idx * T * Hv + hv_idx; + + float acc = 0.0f; + for (int j = len - 1; j >= 0; --j) { + const int idx = base + (t0 + j) * Hv; + acc += static_cast(dg[idx]); + float gv = metal::max(static_cast(g[idx]), 1e-6f); + dg[idx] = acc / gv; + } +} diff --git a/mlx/backend/metal/kernels/gated_delta_update_nax_vjp.metal b/mlx/backend/metal/kernels/gated_delta_update_nax_vjp.metal new file mode 100644 index 0000000000..954478d425 --- /dev/null +++ b/mlx/backend/metal/kernels/gated_delta_update_nax_vjp.metal @@ -0,0 +1,44 @@ +#include "mlx/backend/metal/kernels/gated_delta_update_nax_vjp.h" +#include "mlx/backend/metal/kernels/utils.h" + +using namespace metal; + +#define instantiate_gdu_vjp_nax(in_type, dk, dv, hk, hv, c, ckpt) \ + instantiate_kernel( \ + "gated_delta_vjp_fused_nax_" #in_type "_" #dk "_" #dv \ + "_" #hk "_" #hv "_" #c "_" #ckpt, \ + gated_delta_vjp_fused_nax, \ + in_type, \ + dk, \ + dv, \ + hk, \ + hv, \ + c, \ + ckpt) + +#define instantiate_gdu_vjp_dims(in_type, dk, dv, hk, hv) \ + instantiate_gdu_vjp_nax(in_type, dk, dv, hk, hv, 16, 1) \ + instantiate_gdu_vjp_nax(in_type, dk, dv, hk, hv, 16, 4) \ + instantiate_gdu_vjp_nax(in_type, dk, dv, hk, hv, 16, 8) \ + instantiate_gdu_vjp_nax(in_type, dk, dv, hk, hv, 16, 16) + +#define instantiate_gdu_vjp(in_type) \ + instantiate_gdu_vjp_dims(in_type, 128, 128, 24, 24) \ + instantiate_gdu_vjp_dims(in_type, 128, 128, 32, 32) \ + instantiate_gdu_vjp_dims(in_type, 128, 128, 16, 32) \ + instantiate_gdu_vjp_dims(in_type, 128, 128, 16, 16) \ + instantiate_gdu_vjp_dims(in_type, 128, 128, 16, 48) + +instantiate_gdu_vjp(float); +instantiate_gdu_vjp(bfloat16_t); + +// Postprocessing: converts dL/dgamma to dL/dg. Not templated on ckpt. +#define instantiate_gdu_dgamma(in_type, c) \ + instantiate_kernel( \ + "gated_delta_dgamma_to_dg_" #in_type "_" #c, \ + gated_delta_dgamma_to_dg, \ + in_type, \ + c) + +instantiate_gdu_dgamma(float, 16); +instantiate_gdu_dgamma(bfloat16_t, 16); diff --git a/mlx/backend/metal/kernels/gated_delta_update_vjp.h b/mlx/backend/metal/kernels/gated_delta_update_vjp.h new file mode 100644 index 0000000000..13b96416c7 --- /dev/null +++ b/mlx/backend/metal/kernels/gated_delta_update_vjp.h @@ -0,0 +1,184 @@ +// Copyright © 2024 Apple Inc. +#pragma once + +#include +#include + +#include "mlx/backend/metal/kernels/atomic.h" +#include "mlx/backend/metal/kernels/utils.h" + +using namespace metal; + +template +[[kernel]] void gated_delta_vjp_seq( + const device InT* q [[buffer(0)]], // [B, T, Hk, Dk] + const device InT* k [[buffer(1)]], // [B, T, Hk, Dk] + const device InT* v [[buffer(2)]], // [B, T, Hv, Dv] + const device InT* g [[buffer(3)]], // [B, T, Hv] + const device InT* b [[buffer(4)]], // [B, T, Hv] + const device InT* cot_o [[buffer(5)]], // [B, T, Hv, Dv] + const device float* cot_h [[buffer(6)]], // [B, Hv, Dv, Dk] + const device float* state_cache [[buffer(7)]], // [B*Hv, n_ckpt, Dv, Dk] + constant int& T [[buffer(8)]], + device mlx_atomic* dq [[buffer(9)]], + device mlx_atomic* dk [[buffer(10)]], + device float* dv [[buffer(11)]], + device mlx_atomic* dg [[buffer(12)]], + device mlx_atomic* db [[buffer(13)]], + device float* dh [[buffer(14)]], + uint3 thread_position_in_grid [[thread_position_in_grid]], + uint3 thread_position_in_threadgroup [[thread_position_in_threadgroup]], + uint thread_index_in_simdgroup [[thread_index_in_simdgroup]]) { + auto n = thread_position_in_grid.z; + auto b_idx = n / Hv; + auto hv_idx = n % Hv; + auto hk_idx = hv_idx / (Hv / Hk); + constexpr int n_per_t = Dk / 32; + + auto dk_idx = thread_position_in_threadgroup.x; + auto dv_idx = thread_position_in_grid.y; + + const int qk_stride = Hk * Dk; + const int v_stride = Hv * Dv; + const int g_stride = Hv; + + auto q_base = q + b_idx * T * qk_stride + hk_idx * Dk; + auto k_base = k + b_idx * T * qk_stride + hk_idx * Dk; + auto dq_base = dq + b_idx * T * qk_stride + hk_idx * Dk; + auto dk_base = dk + b_idx * T * qk_stride + hk_idx * Dk; + + auto v_base = v + b_idx * T * v_stride + hv_idx * Dv; + auto dv_base = dv + b_idx * T * v_stride + hv_idx * Dv; + auto co_base = cot_o + b_idx * T * v_stride + hv_idx * Dv; + + auto g_base = g + b_idx * T * g_stride + hv_idx; + auto b_base = b + b_idx * T * g_stride + hv_idx; + auto dg_base = dg + b_idx * T * g_stride + hv_idx; + auto db_base = db + b_idx * T * g_stride + hv_idx; + + const int n_ckpt = (T + Ckpt - 1) / Ckpt; + + float s_hat[n_per_t]; + auto base_state = cot_h + (n * Dv + dv_idx) * Dk; + for (int i = 0; i < n_per_t; i++) { + s_hat[i] = base_state[n_per_t * dk_idx + i]; + } + + // Only this thread's n_per_t slice of the state is ever needed, so a whole + // segment of entry states costs Ckpt * n_per_t registers and the replay needs + // no device buffer. + float seg[Ckpt][n_per_t]; + float s_prev[n_per_t]; + float s_dec[n_per_t]; + + for (int seg_idx = n_ckpt - 1; seg_idx >= 0; --seg_idx) { + const int t0 = seg_idx * Ckpt; + const int seg_len = metal::min(Ckpt, T - t0); + + // Replay forward from the checkpoint, recording the entry state of each + // step in the segment. + auto c_state = + state_cache + n * n_ckpt * Dv * Dk + seg_idx * Dv * Dk + dv_idx * Dk; + + float s[n_per_t]; + for (int i = 0; i < n_per_t; ++i) { + s[i] = c_state[n_per_t * dk_idx + i]; + } + + for (int j = 0; j < seg_len; ++j) { + const int t = t0 + j; + + for (int i = 0; i < n_per_t; ++i) { + seg[j][i] = s[i]; + } + + float gamma = static_cast(g_base[t * g_stride]); + float beta = static_cast(b_base[t * g_stride]); + + auto k_t = k_base + t * qk_stride; + + float kv_mem = 0.0f; + for (int i = 0; i < n_per_t; ++i) { + const int s_idx = n_per_t * dk_idx + i; + s[i] *= gamma; + kv_mem += s[i] * static_cast(k_t[s_idx]); + } + kv_mem = simd_sum(kv_mem); + + float delta = + beta * (static_cast(v_base[t * v_stride + dv_idx]) - kv_mem); + + for (int i = 0; i < n_per_t; ++i) { + s[i] += delta * static_cast(k_t[n_per_t * dk_idx + i]); + } + } + + // Backward over the same steps, newest first. + for (int j = seg_len - 1; j >= 0; --j) { + const int t = t0 + j; + + float gamma = static_cast(g_base[t * g_stride]); + float beta = static_cast(b_base[t * g_stride]); + + auto q_t = q_base + t * qk_stride; + auto k_t = k_base + t * qk_stride; + auto dq_t = dq_base + t * qk_stride; + auto dk_t = dk_base + t * qk_stride; + + float kv_mem = 0.0f; + float co = static_cast(co_base[t * v_stride + dv_idx]); + float w = 0.0f; + for (int i = 0; i < n_per_t; i++) { + const int s_idx = n_per_t * dk_idx + i; + s_prev[i] = seg[j][i]; + s_dec[i] = s_prev[i] * gamma; + kv_mem += s_dec[i] * static_cast(k_t[s_idx]); + + s_hat[i] += co * static_cast(q_t[s_idx]); + + w += s_hat[i] * static_cast(k_t[s_idx]); + } + kv_mem = simd_sum(kv_mem); + w = simd_sum(w); + + if (thread_index_in_simdgroup == 0) { + dv_base[t * v_stride + dv_idx] = beta * w; + } + + float u = static_cast(v_base[t * v_stride + dv_idx]) - kv_mem; + float delta = beta * u; + + if (thread_index_in_simdgroup == 0) { + mlx_atomic_fetch_add_explicit(db_base + t * g_stride, w * u, 0); + } + + float dgamma = 0.0f; + for (int i = 0; i < n_per_t; ++i) { + auto s_idx = n_per_t * dk_idx + i; + + float s_t = s_dec[i] + delta * static_cast(k_t[s_idx]); + mlx_atomic_fetch_add_explicit(dq_t, co * s_t, s_idx); + + float contrib = beta * (u * s_hat[i] - w * s_dec[i]); + mlx_atomic_fetch_add_explicit(dk_t, contrib, s_idx); + + s_hat[i] -= beta * w * static_cast(k_t[s_idx]); + + dgamma += s_hat[i] * s_prev[i]; + } + dgamma = simd_sum(dgamma); + if (thread_index_in_simdgroup == 0) { + mlx_atomic_fetch_add_explicit(dg_base + t * g_stride, dgamma, 0); + } + + for (int i = 0; i < n_per_t; ++i) { + s_hat[i] *= gamma; + } + } + } + + for (int i = 0; i < n_per_t; ++i) { + auto s_idx = n_per_t * dk_idx + i; + dh[(n * Dv + dv_idx) * Dk + s_idx] = s_hat[i]; + } +} diff --git a/mlx/backend/metal/kernels/gated_delta_update_vjp.metal b/mlx/backend/metal/kernels/gated_delta_update_vjp.metal new file mode 100644 index 0000000000..7112750d41 --- /dev/null +++ b/mlx/backend/metal/kernels/gated_delta_update_vjp.metal @@ -0,0 +1,35 @@ +#include "mlx/backend/metal/kernels/gated_delta_update_vjp.h" +#include "mlx/backend/metal/kernels/gated_delta_update.h" +#include "mlx/backend/metal/kernels/utils.h" + +using namespace metal; + +#define instantiate_gdn_vjp_seq(in_type, dk, dv, hk, hv, ckpt) \ + instantiate_kernel( \ + "seq_gated_delta_vjp_" #in_type "_" #dk "_" #dv "_" #hk "_" #hv \ + "_" #ckpt, \ + gated_delta_vjp_seq, \ + in_type, \ + dk, \ + dv, \ + hk, \ + hv, \ + ckpt) + +#define instantiate_gated_delta_vjp_dims(in_type, dk, dv, hk, hv) \ + instantiate_gdn_vjp_seq(in_type, dk, dv, hk, hv, 1) \ + instantiate_gdn_vjp_seq(in_type, dk, dv, hk, hv, 4) \ + instantiate_gdn_vjp_seq(in_type, dk, dv, hk, hv, 8) \ + instantiate_gdn_vjp_seq(in_type, dk, dv, hk, hv, 16) + +#define instantiate_gated_delta_vjp(in_type) \ + instantiate_gated_delta_vjp_dims(in_type, 128, 128, 24, 24) \ + instantiate_gated_delta_vjp_dims(in_type, 128, 128, 32, 32) \ + instantiate_gated_delta_vjp_dims(in_type, 128, 128, 16, 32) \ + instantiate_gated_delta_vjp_dims(in_type, 128, 128, 16, 48) \ + instantiate_gated_delta_vjp_dims(in_type, 128, 128, 16, 16) \ + instantiate_gated_delta_vjp_dims(in_type, 128, 128, 16, 64) + +instantiate_gated_delta_vjp(float); +instantiate_gated_delta_vjp(bfloat16_t); +instantiate_gated_delta_vjp(float16_t); \ No newline at end of file diff --git a/mlx/backend/metal/nojit_kernels.cpp b/mlx/backend/metal/nojit_kernels.cpp index 4f48a25b77..cd83b86bd6 100644 --- a/mlx/backend/metal/nojit_kernels.cpp +++ b/mlx/backend/metal/nojit_kernels.cpp @@ -526,4 +526,20 @@ MTL::ComputePipelineState* get_gated_delta_nax_kernel( return d.get_kernel(kernel_name, hash_name, func_consts); } +MTL::ComputePipelineState* get_gated_delta_vjp_kernel( + metal::Device& d, + const std::string& kernel_name, + const std::string& hash_name, + const metal::MTLFCList& func_consts) { + return d.get_kernel(kernel_name, hash_name, func_consts); +} + +MTL::ComputePipelineState* get_gated_delta_vjp_nax_kernel( + metal::Device& d, + const std::string& kernel_name, + const std::string& hash_name, + const metal::MTLFCList& func_consts) { + return d.get_kernel(kernel_name, hash_name, func_consts); +} + } // namespace mlx::core diff --git a/mlx/fast.cpp b/mlx/fast.cpp index add78fa7c1..030286c5cf 100644 --- a/mlx/fast.cpp +++ b/mlx/fast.cpp @@ -977,13 +977,16 @@ std::vector ScaledDotProductAttention::vjp( const std::vector& outputs) { assert(primals.size() >= 3); assert(cotangents.size() == outputs.size()); - auto s = stream(); if (ScaledDotProductAttentionVJP::use_fallback(primals[0], s)) { assert(outputs.size() == 1); return Custom::vjp(primals, cotangents, argnums, outputs); } + if (outputs.size() != 2) { + return Custom::vjp(primals, cotangents, argnums, outputs); + } + auto fallback = [sdpa = fallback_, s](const std::vector& inputs) { std::vector primals(inputs.begin(), std::prev(inputs.end())); auto [_, vjps] = mlx::core::vjp(sdpa, primals, {inputs.back()}); @@ -1158,6 +1161,54 @@ std::vector gated_delta_update( auto result = fallback({q, k, v, g, beta, h0, mask}); return result; } +std::vector GatedDeltaUpdate::vjp( + const std::vector& primals, + const std::vector& cotangents, + const std::vector& argnums, + const std::vector& outputs) { + const int Hk = primals[0].shape(2); + const int Dk = primals[0].shape(3); + const int Hv = primals[2].shape(2); + const int Dv = primals[2].shape(3); + + if (GatedDeltaUpdateVJP::use_fallback(Hk, Dk, Hv, Dv, stream())) { + return Custom::vjp(primals, cotangents, argnums, outputs); + } + + if (primals.size() != 6) { + throw std::runtime_error( + "[GatedDeltaUpdate::vjp] expected 6 primals (q,k,v,g,beta,h0), got " + + std::to_string(primals.size())); + } + if (cotangents.size() != 2) { + throw std::runtime_error( + "[GatedDeltaUpdate::vjp] expected 2 cotangents (out, state), got " + + std::to_string(cotangents.size())); + } + + std::vector inputs = primals; // q,k,v,g,beta,h0 (indices 0..5) + inputs.push_back(cotangents[0]); // cot_o -> index 6 + inputs.push_back(cotangents[1]); // cot_h -> index 7 + + std::vector shapes; + std::vector dtypes; + for (int i = 0; i < 6; ++i) { + shapes.push_back(primals[i].shape()); + dtypes.push_back(i == 5 ? float32 : primals[i].dtype()); // dh is float32 + } + + auto vjps = array::make_arrays( + std::move(shapes), + std::move(dtypes), + std::make_shared(stream(), fallback_), + std::move(inputs)); + + std::vector returned_vjps; + for (int arg : argnums) { + returned_vjps.push_back(std::move(vjps[arg])); + } + return returned_vjps; +} bool Quantize::is_equivalent(const Primitive& other) const { const Quantize& p_other = static_cast(other); diff --git a/mlx/fast_primitives.h b/mlx/fast_primitives.h index cef13bb0ff..42d66ea488 100644 --- a/mlx/fast_primitives.h +++ b/mlx/fast_primitives.h @@ -418,6 +418,12 @@ class GatedDeltaUpdate : public Custom { void eval_gpu(const std::vector& inputs, std::vector& outputs) override; + std::vector vjp( + const std::vector& primals, + const std::vector& cotangents, + const std::vector& argnums, + const std::vector& outputs) override; + DEFINE_NAME(GatedDeltaUpdate); DEFINE_INPUT_OUTPUT_SHAPE() auto state() const { @@ -425,6 +431,39 @@ class GatedDeltaUpdate : public Custom { } private: + bool is_training_; +}; + +class GatedDeltaUpdateVJP : public Custom { + public: + GatedDeltaUpdateVJP( + Stream stream, + std::function(std::vector)> fallback) + : Custom(stream, std::move(fallback)) {} + + static bool use_fallback( + const int Hk, + const int Dk, + const int Hv, + const int Dv, + Stream s); + + void eval_cpu(const std::vector& inputs, std::vector& outputs) + override { + throw std::runtime_error("NYI"); + } + + void eval_gpu(const std::vector& inputs, std::vector& outputs) + override; + + DEFINE_NAME(GatedDeltaUpdateVJP); + DEFINE_INPUT_OUTPUT_SHAPE() + auto state() const { + return std::make_tuple(nullptr); /* TODO */ + } + + private: + bool has_cache_; }; class Quantize : public Custom { From 8b7f81c0e093afb3651e4dc2e28df0176cad919a Mon Sep 17 00:00:00 2001 From: mlx-dev Date: Fri, 25 Sep 2026 07:30:24 -0700 Subject: [PATCH 02/11] add grad benchmark and tests --- benchmarks/python/gated_delta_bench.py | 138 ++++++++++++++++++++----- python/tests/test_fast_gated_delta.py | 53 ++++++++++ 2 files changed, 168 insertions(+), 23 deletions(-) diff --git a/benchmarks/python/gated_delta_bench.py b/benchmarks/python/gated_delta_bench.py index 277897699c..1e0dfcef97 100644 --- a/benchmarks/python/gated_delta_bench.py +++ b/benchmarks/python/gated_delta_bench.py @@ -3,23 +3,16 @@ import itertools import os import time -from datetime import datetime -from typing import Optional, Tuple import mlx.core as mx -import numpy as np -RED_BOLD = "\033[1;31m" -GREEN = "\033[0;32m" RESET = "\033[0m" - N_warmup = 8 N_iter_bench = 80 N_iter_func = 5 -# similar to ./blas/bench_gemm.py def bench(f, *args): for _ in range(N_warmup): f(*args) @@ -30,7 +23,7 @@ def bench(f, *args): f(*args) mx.synchronize() e = time.perf_counter_ns() - return (e - s) * 1e-9 # total seconds for N_iter_bench * N_iter_func calls + return (e - s) * 1e-9 def do_kernel_bench(f, *args): @@ -43,7 +36,23 @@ def do_kernel_bench(f, *args): return ys -def benchmark_shape(B, T, Hk, Hv, Dk, Dv, chunk_sizes): +def do_grad_bench(f, *args): + ys = [] + for _ in range(N_iter_func): + ys.extend(f(*args)) + mx.eval(ys) + return ys + + +def make_grad_fn(): + def f(q, k, v, g, b, h0): + out, state = mx.fast.gated_delta_update(q, k, v, g, b, h0) + return out.sum() + state.sum() + + return mx.grad(f, argnums=(0, 1, 2, 3, 4, 5)) + + +def make_inputs(B, T, Hk, Hv, Dk, Dv): mx.random.seed(42) q = mx.random.normal(shape=(B, T, Hk, Dk)) k = mx.random.normal(shape=(B, T, Hk, Dk)) @@ -51,12 +60,18 @@ def benchmark_shape(B, T, Hk, Hv, Dk, Dv, chunk_sizes): v = mx.random.normal(shape=(B, T, Hv, Dv)) g = mx.random.normal(shape=(B, T, Hv)) * 0.1 - 1.0 b = mx.sigmoid(mx.random.normal(shape=(B, T, Hv))) + h0 = mx.zeros((B, Hv, Dv, Dk), dtype=mx.float32) + mx.eval(q, k, v, g, b, h0) + return q, k, v, g, b, h0 + + +def benchmark_shape(B, T, Hk, Hv, Dk, Dv, chunk_sizes): + q, k, v, g, b, h0 = make_inputs(B, T, Hk, Hv, Dk, Dv) shape_str = f"B={B} T={T} Hk={Hk} Hv={Hv} Dk={Dk} Dv={Dv}" denom = N_iter_bench * N_iter_func os.environ["GATED_DELTA_CHUNK"] = "0" - h0 = mx.zeros((B, Hv, Dv, Dk), dtype=mx.float32) mx.eval(*mx.fast.gated_delta_update(q, k, v, g, b, initial_state=h0)) ms_seq = ( bench(do_kernel_bench, mx.fast.gated_delta_update, q, k, v, g, b, h0) @@ -68,7 +83,6 @@ def benchmark_shape(B, T, Hk, Hv, Dk, Dv, chunk_sizes): for C in (c for c in chunk_sizes if c != 0): try: os.environ["GATED_DELTA_CHUNK"] = str(C) - h0 = mx.zeros((B, Hv, Dv, Dk), dtype=mx.float32) mx.eval(*mx.fast.gated_delta_update(q, k, v, g, b, initial_state=h0)) ms_c = ( bench(do_kernel_bench, mx.fast.gated_delta_update, q, k, v, g, b, h0) @@ -86,18 +100,11 @@ def benchmark_shape(B, T, Hk, Hv, Dk, Dv, chunk_sizes): def run_benchmark(run_full, to_csv=False, csv_path="benchmark_results.csv"): if run_full: Bs = [1, 4, 8, 16] - Ts = [8, 64, 256, 512, 1024, 2048, 4096] - Hks = [16] - Hvs = [32] - Dks = [128] - Dvs = [128] + Ts = [8, 64, 256, 512, 1024, 2048] else: Bs = [1, 8, 16] Ts = [8, 512, 1024, 2048] - Hks = [16] - Hvs = [32] - Dks = [128] - Dvs = [128] + Hks, Hvs, Dks, Dvs = [16], [32], [128], [128] chunk_sizes = [0, 8, 16] non_zero_Cs = [C for C in chunk_sizes if C != 0] @@ -106,13 +113,13 @@ def run_benchmark(run_full, to_csv=False, csv_path="benchmark_results.csv"): f"C={C} (speedup)" for C in non_zero_Cs ] - col_widths = [6, 6, 6, 6, 6, 6, 15] + [25] * (len(non_zero_Cs)) + col_widths = [6, 6, 6, 6, 6, 6, 15] + [25] * len(non_zero_Cs) fmt = "".join(f"{{:<{w}}}" for w in col_widths) rows = [] print(fmt.format(*headers)) - print("-" * (sum(col_widths))) + print("-" * sum(col_widths)) for B, T, Hk, Hv, Dk, Dv in itertools.product(Bs, Ts, Hks, Hvs, Dks, Dvs): shapes_s, base_time_s, speedups, base_time = benchmark_shape( @@ -135,11 +142,96 @@ def run_benchmark(run_full, to_csv=False, csv_path="benchmark_results.csv"): print(f"\nResults also written to {csv_path}") +def benchmark_variants_shape(B, T, Hk, Hv, Dk, Dv, do_backward, variants): + q, k, v, g, b, h0 = make_inputs(B, T, Hk, Hv, Dk, Dv) + denom = N_iter_bench * N_iter_func + + if do_backward: + fn = make_grad_fn() + runner = do_grad_bench + else: + fn = mx.fast.gated_delta_update + runner = do_kernel_bench + + def time_one(variant): + os.environ["GATED_DELTA_VJP_FALLBACK"] = "1" if variant == "fallback" else "0" + C = "16" if variant == "nax" else "0" + if do_backward: + os.environ["GATED_DELTA_CHUNK"] = "16" + os.environ["GATED_DELTA_CHUNK_VJP"] = C + else: + os.environ["GATED_DELTA_CHUNK"] = C + mx.eval(*fn(q, k, v, g, b, h0)) + return bench(runner, fn, q, k, v, g, b, h0) / denom * 1e3 + + times = [] + for variant in variants: + try: + times.append(time_one(variant)) + except Exception as ex: + print(f" {variant} failed: {ex}") + times.append(float("nan")) + mx.clear_cache() + return times + + +def run_variants_benchmark(run_full, do_backward=False, do_fallback=False): + if run_full: + Bs = [1, 4, 8, 16] + Ts = [8, 32, 64, 128, 256, 512, 1024, 2048, 4096] + else: + Bs = [1, 8] + Ts = [8, 512, 1024] + if do_fallback: + Ts = [8, 32, 64, 128, 256, 512] + Hks, Hvs, Dks, Dvs = [16], [32], [128], [128] + + variants = ["seq", "nax"] + if do_fallback: + variants = ["fallback"] + variants + + headers = ["B", "T", "Hk", "Hv", "Dk", "Dv"] + headers += [f"{v} (ms)" for v in variants] + headers += [f"{v} (speedup)" for v in variants[1:]] + + col_widths = [6, 6, 6, 6, 6, 6] + [16] * len(variants) + [16] * (len(variants) - 1) + fmt = "".join(f"{{:<{w}}}" for w in col_widths) + + mode = "BACKWARD" if do_backward else "FORWARD" + print(f"\n=== {mode}: {' vs '.join(variants)} ===") + print(fmt.format(*headers)) + print("-" * sum(col_widths)) + + for B, T, Hk, Hv, Dk, Dv in itertools.product(Bs, Ts, Hks, Hvs, Dks, Dvs): + try: + times = benchmark_variants_shape( + B, T, Hk, Hv, Dk, Dv, do_backward, variants + ) + except Exception as ex: + print(f" B={B} T={T} failed: {ex}") + mx.clear_cache() + continue + + base = times[0] + row = [f"{B}", f"{T}", f"{Hk}", f"{Hv}", f"{Dk}", f"{Dv}"] + row += [f"{t:.3f}" for t in times] + row += [f"{base / t:.2f}x" if t > 0 else "nan" for t in times[1:]] + print(fmt.format(*row)) + print(RESET, end="") + + if __name__ == "__main__": parser = argparse.ArgumentParser(description="Gated delta benchmark") parser.add_argument("--full", "-f", action="store_true") parser.add_argument("--csv", "-c", action="store_true") parser.add_argument("--csv_out", "-co", default="benchmark_results.csv") + parser.add_argument("--fallback", "-fb", action="store_true") + parser.add_argument("--backward", "-bw", action="store_true") args = parser.parse_args() - run_benchmark(args.full, to_csv=args.csv, csv_path=args.csv_out) + if args.backward or args.fallback: + run_variants_benchmark( + args.full, do_backward=args.backward, do_fallback=args.fallback + ) + else: + run_benchmark(args.full, to_csv=args.csv, csv_path=args.csv_out) diff --git a/python/tests/test_fast_gated_delta.py b/python/tests/test_fast_gated_delta.py index d7a33c67a3..1e67347cc7 100644 --- a/python/tests/test_fast_gated_delta.py +++ b/python/tests/test_fast_gated_delta.py @@ -155,6 +155,7 @@ class TestGatedDelta(mlx_tests.MLXTestCase): diff_heads2, large_t_dims, ] + vjp_dims = [base_dims, unaligned_dims, diff_heads, diff_heads2] @unittest.skipIf(not has_torch, "requires Torch") def test_gated_delta_fallback(self): @@ -281,6 +282,58 @@ def test_gated_delta_nax(self): mx.allclose(hf_ref, hf, atol=1e-1, rtol=1e-4), msg="State " + msg ) + @unittest.skipIf(not mx.metal.is_available(), "Metal is not available") + def test_gated_delta_vjp(self): + for dims in self.vjp_dims: + B, Hk, Hv, T, Dk, Dv = dims + + mx.random.seed(0) + q = mx.random.normal(shape=(B, T, Hk, Dk)) + k = mx.random.normal(shape=(B, T, Hk, Dk)) + k = k / (mx.linalg.norm(k, axis=-1, keepdims=True) + 1e-6) + v = mx.random.normal(shape=(B, T, Hv, Dv)) + g = mx.sigmoid(mx.random.normal(shape=(B, T, Hv))) + b = mx.sigmoid(mx.random.normal(shape=(B, T, Hv))) + h0 = mx.random.normal((B, Hv, Dv, Dk), dtype=mx.float32) + primals = [q, k, v, g, b, h0] + + cotans = [ + mx.random.normal(shape=(B, T, Hv, Dv)) * 1e-2, + mx.random.normal(shape=(B, Hv, Dv, Dk)) * 1e-2, + ] + + def f(q, k, v, g, b, h0): + out, state = mx.fast.gated_delta_update(q, k, v, g, b, initial_state=h0) + return out, state + + def run(fallback, chunk): + with mlx_tests.scoped_env( + GATED_DELTA_VJP_FALLBACK="1" if fallback else "0", + GATED_DELTA_CHUNK_VJP=str(chunk), + ): + outs, vjps = mx.vjp(f, primals, cotans) + mx.eval(outs, vjps) + return outs, vjps + + o_ref, vjp_ref = run(True, 0) + + for chunk in (0, 16): + rtol = 1e-5 if chunk == 0 else 1e-2 + atol = 1e-5 if chunk == 0 else 1e-2 + with self.subTest(dims=dims, chunk=chunk): + o_out, vjp_out = run(False, chunk) + for i in range(len(o_ref)): + self.assertTrue( + mx.allclose(o_ref[i], o_out[i], rtol=rtol, atol=atol), + msg=f"Output {i}, dims={dims}, chunk={chunk}", + ) + for i in range(len(vjp_ref)): + self.assertTrue( + mx.allclose(vjp_ref[i], vjp_out[i], rtol=rtol, atol=atol), + msg=f"Grad {i}, dims={dims}, chunk={chunk}", + ) + mx.clear_cache() + if __name__ == "__main__": mlx_tests.MLXTestRunner(failfast=True) From 2f82935a68a8a0678115d2ce105c8a409ed7dc42 Mon Sep 17 00:00:00 2001 From: mlx-dev Date: Mon, 28 Sep 2026 02:21:17 -0700 Subject: [PATCH 03/11] add cuda gdn fallback --- mlx/backend/cuda/primitives.cpp | 9 +++++++++ 1 file changed, 9 insertions(+) diff --git a/mlx/backend/cuda/primitives.cpp b/mlx/backend/cuda/primitives.cpp index c02357b08f..f85e5cbae0 100644 --- a/mlx/backend/cuda/primitives.cpp +++ b/mlx/backend/cuda/primitives.cpp @@ -34,6 +34,15 @@ bool fast::GatedDeltaUpdate::use_fallback( return true; } +bool GatedDeltaUpdateVJP::use_fallback( + const int Hk, + const int Dk, + const int Hv, + const int Dv, + Stream s) { + return true; +} + NO_GPU_MULTI(LUF) NO_GPU_MULTI(QRF) NO_GPU_MULTI(SVD) From d21404b1368026016487cf14cf777176aae4a759 Mon Sep 17 00:00:00 2001 From: mlx-dev Date: Mon, 28 Sep 2026 02:25:47 -0700 Subject: [PATCH 04/11] fix no_gpu and cuda --- mlx/backend/cuda/primitives.cpp | 3 ++- mlx/backend/no_gpu/primitives.cpp | 11 +++++++++++ 2 files changed, 13 insertions(+), 1 deletion(-) diff --git a/mlx/backend/cuda/primitives.cpp b/mlx/backend/cuda/primitives.cpp index f85e5cbae0..a682edc47c 100644 --- a/mlx/backend/cuda/primitives.cpp +++ b/mlx/backend/cuda/primitives.cpp @@ -52,7 +52,8 @@ NO_GPU_MULTI(Eigh) namespace fast { NO_GPU_MULTI(GatedDeltaUpdate) -} +NO_GPU_MULTI(GatedDeltaUpdateVJP) +} // namespace fast namespace distributed { NO_GPU_MULTI(Send) diff --git a/mlx/backend/no_gpu/primitives.cpp b/mlx/backend/no_gpu/primitives.cpp index c37515f9ef..4d843808b2 100644 --- a/mlx/backend/no_gpu/primitives.cpp +++ b/mlx/backend/no_gpu/primitives.cpp @@ -62,6 +62,15 @@ bool fast::GatedDeltaUpdate::use_fallback( return true; } +bool GatedDeltaUpdateVJP::use_fallback( + const int Hk, + const int Dk, + const int Hv, + const int Dv, + Stream s) { + return true; +} + NO_GPU(Abs) NO_GPU(Add) NO_GPU(AddMM) @@ -190,6 +199,8 @@ NO_GPU_USE_FALLBACK(RoPE) NO_GPU_MULTI(ScaledDotProductAttention) NO_GPU_MULTI(ScaledDotProductAttentionVJP) NO_GPU_MULTI(GatedDeltaUpdate) +NO_GPU_MULTI(GatedDeltaUpdateVJP) +NO_GPU_MULTI(ConvertFP8) NO_GPU_MULTI(ConvertFP8) NO_GPU_MULTI(Quantize) NO_GPU_MULTI(CustomKernel) From efffc2536692c91e525cf8fc35c4a78c315d9959 Mon Sep 17 00:00:00 2001 From: mlx-dev Date: Mon, 28 Sep 2026 02:32:49 -0700 Subject: [PATCH 05/11] add fast namespace --- mlx/backend/cuda/primitives.cpp | 2 +- mlx/backend/no_gpu/primitives.cpp | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/mlx/backend/cuda/primitives.cpp b/mlx/backend/cuda/primitives.cpp index a682edc47c..104c7fd3fc 100644 --- a/mlx/backend/cuda/primitives.cpp +++ b/mlx/backend/cuda/primitives.cpp @@ -34,7 +34,7 @@ bool fast::GatedDeltaUpdate::use_fallback( return true; } -bool GatedDeltaUpdateVJP::use_fallback( +bool fast::GatedDeltaUpdateVJP::use_fallback( const int Hk, const int Dk, const int Hv, diff --git a/mlx/backend/no_gpu/primitives.cpp b/mlx/backend/no_gpu/primitives.cpp index 4d843808b2..ecf8b43a6d 100644 --- a/mlx/backend/no_gpu/primitives.cpp +++ b/mlx/backend/no_gpu/primitives.cpp @@ -62,7 +62,7 @@ bool fast::GatedDeltaUpdate::use_fallback( return true; } -bool GatedDeltaUpdateVJP::use_fallback( +bool fast::GatedDeltaUpdateVJP::use_fallback( const int Hk, const int Dk, const int Hv, From 1f3892c39ec67ddef82ab7b8fa7cc57e2331036f Mon Sep 17 00:00:00 2001 From: mlx-dev Date: Mon, 28 Sep 2026 02:38:50 -0700 Subject: [PATCH 06/11] remove duplicate line --- mlx/backend/no_gpu/primitives.cpp | 1 - 1 file changed, 1 deletion(-) diff --git a/mlx/backend/no_gpu/primitives.cpp b/mlx/backend/no_gpu/primitives.cpp index ecf8b43a6d..911f4e7673 100644 --- a/mlx/backend/no_gpu/primitives.cpp +++ b/mlx/backend/no_gpu/primitives.cpp @@ -201,7 +201,6 @@ NO_GPU_MULTI(ScaledDotProductAttentionVJP) NO_GPU_MULTI(GatedDeltaUpdate) NO_GPU_MULTI(GatedDeltaUpdateVJP) NO_GPU_MULTI(ConvertFP8) -NO_GPU_MULTI(ConvertFP8) NO_GPU_MULTI(Quantize) NO_GPU_MULTI(CustomKernel) } // namespace fast From b302f26145c088ef8e44b3d927ae828facd4ba2a Mon Sep 17 00:00:00 2001 From: mlx-dev Date: Mon, 28 Sep 2026 04:53:49 -0700 Subject: [PATCH 07/11] fix jit build --- mlx/backend/metal/gated_delta_update.cpp | 46 +++++++++++++++++++---- mlx/backend/metal/jit/includes.h | 1 + mlx/backend/metal/jit_kernels.cpp | 47 ++++++++++++++++++++++-- mlx/backend/metal/kernels.h | 33 +++++++++-------- mlx/backend/metal/nojit_kernels.cpp | 24 +++++++++--- 5 files changed, 121 insertions(+), 30 deletions(-) diff --git a/mlx/backend/metal/gated_delta_update.cpp b/mlx/backend/metal/gated_delta_update.cpp index ef9eca5fab..a089651b29 100644 --- a/mlx/backend/metal/gated_delta_update.cpp +++ b/mlx/backend/metal/gated_delta_update.cpp @@ -133,8 +133,8 @@ void GatedDeltaUpdate::eval_gpu( std::string base_name = "gated_delta_fused_nax_" + suffix + "_" + std::to_string(C) + "_" + std::to_string(ckpt); - auto delta_kernel = - get_gated_delta_nax_kernel(d, base_name, base_name, func_consts); + auto delta_kernel = get_gated_delta_nax_kernel( + d, base_name, base_name, func_consts, q, Dk, Dv, Hk, Hv, C, ckpt); compute_encoder.set_compute_pipeline_state(delta_kernel); compute_encoder.set_input_array(q, 0); @@ -333,7 +333,17 @@ void GatedDeltaUpdateVJP::eval_gpu( std::string base_name = "gated_delta_fused_nax_" + suffix + ckpt_suffix; auto delta_kernel = get_gated_delta_nax_kernel( - d, base_name, base_name + "_save", save_consts); + d, + base_name, + base_name + "_save", + save_consts, + q, + Dk, + Dv, + Hk, + Hv, + C, + ckpt); compute_encoder.set_compute_pipeline_state(delta_kernel); compute_encoder.set_input_array(q, 0); @@ -358,8 +368,19 @@ void GatedDeltaUpdateVJP::eval_gpu( std::string base_name = "gated_delta_vjp_fused_nax_" + suffix + ckpt_suffix; - auto delta_kernel = - get_gated_delta_vjp_nax_kernel(d, base_name, base_name, no_consts); + auto delta_kernel = get_gated_delta_vjp_nax_kernel( + d, + base_name, + base_name, + no_consts, + q, + Dk, + Dv, + Hk, + Hv, + C, + ckpt, + false); compute_encoder.set_compute_pipeline_state(delta_kernel); compute_encoder.set_input_array(q, 0); @@ -392,8 +413,19 @@ void GatedDeltaUpdateVJP::eval_gpu( std::string base_name = "gated_delta_dgamma_to_dg_" + get_type_string(q.dtype()) + "_" + std::to_string(C); - auto dgamma_kernel = - get_gated_delta_vjp_nax_kernel(d, base_name, base_name, no_consts); + auto dgamma_kernel = get_gated_delta_vjp_nax_kernel( + d, + base_name, + base_name, + no_consts, + q, + Dk, + Dv, + Hk, + Hv, + C, + ckpt, + true); compute_encoder.set_compute_pipeline_state(dgamma_kernel); compute_encoder.set_input_array(g, 0); diff --git a/mlx/backend/metal/jit/includes.h b/mlx/backend/metal/jit/includes.h index b4f72e660c..1f1a8c46bf 100644 --- a/mlx/backend/metal/jit/includes.h +++ b/mlx/backend/metal/jit/includes.h @@ -62,5 +62,6 @@ const char* steel_attention_nax(); const char* gated_delta_update(); const char* gated_delta_update_nax(); +const char* gated_delta_update_nax_vjp(); } // namespace mlx::core::metal diff --git a/mlx/backend/metal/jit_kernels.cpp b/mlx/backend/metal/jit_kernels.cpp index 386b8c2f60..1e0f3cb4be 100644 --- a/mlx/backend/metal/jit_kernels.cpp +++ b/mlx/backend/metal/jit_kernels.cpp @@ -1449,11 +1449,28 @@ MTL::ComputePipelineState* get_gated_delta_nax_kernel( metal::Device& d, const std::string& kernel_name, const std::string& hash_name, - const metal::MTLFCList& func_consts) { + const metal::MTLFCList& func_consts, + const array& q, + int dk, + int dv, + int hk, + int hv, + int c, + int ckpt) { const auto& lib_name = kernel_name; auto lib = d.get_library(lib_name, [&]() { std::string kernel_source; concatenate(kernel_source, metal::utils(), metal::gated_delta_update_nax()); + kernel_source += get_template_definition( + lib_name, + "gated_delta_fused_nax", + get_type_string(q.dtype()), + dk, + dv, + hk, + hv, + c, + ckpt); return kernel_source; }); return d.get_kernel(kernel_name, lib, hash_name, func_consts); @@ -1463,11 +1480,35 @@ MTL::ComputePipelineState* get_gated_delta_vjp_nax_kernel( metal::Device& d, const std::string& kernel_name, const std::string& hash_name, - const metal::MTLFCList& func_consts) { + const metal::MTLFCList& func_consts, + const array& q, + int dk, + int dv, + int hk, + int hv, + int c, + int ckpt, + bool dgamma) { const auto& lib_name = kernel_name; auto lib = d.get_library(lib_name, [&]() { std::string kernel_source; - concatenate(kernel_source, metal::utils(), metal::gated_delta_update_nax()); + concatenate( + kernel_source, metal::utils(), metal::gated_delta_update_nax_vjp()); + if (dgamma) { + kernel_source += get_template_definition( + lib_name, "gated_delta_dgamma_to_dg", get_type_string(q.dtype()), c); + } else { + kernel_source += get_template_definition( + lib_name, + "gated_delta_vjp_fused_nax", + get_type_string(q.dtype()), + dk, + dv, + hk, + hv, + c, + ckpt); + } return kernel_source; }); return d.get_kernel(kernel_name, lib, hash_name, func_consts); diff --git a/mlx/backend/metal/kernels.h b/mlx/backend/metal/kernels.h index da02c9f81e..41648b4810 100644 --- a/mlx/backend/metal/kernels.h +++ b/mlx/backend/metal/kernels.h @@ -444,13 +444,7 @@ MTL::ComputePipelineState* get_gated_delta_kernel( const std::string& hash_name, const metal::MTLFCList& func_consts); -MTL::ComputePipelineState* get_gated_delta_nax_kernel( - metal::Device& d, - const std::string& kernel_name, - const std::string& hash_name, - const metal::MTLFCList& func_consts); - -MTL::ComputePipelineState* get_gated_delta_kernel( +MTL::ComputePipelineState* get_gated_delta_vjp_kernel( metal::Device& d, const std::string& kernel_name, const std::string& hash_name, @@ -460,19 +454,28 @@ MTL::ComputePipelineState* get_gated_delta_nax_kernel( metal::Device& d, const std::string& kernel_name, const std::string& hash_name, - const metal::MTLFCList& func_consts); - -MTL::ComputePipelineState* get_gated_delta_vjp_kernel( - metal::Device& d, - const std::string& kernel_name, - const std::string& hash_name, - const metal::MTLFCList& func_consts); + const metal::MTLFCList& func_consts, + const array& q, + int dk, + int dv, + int hk, + int hv, + int c, + int ckpt); MTL::ComputePipelineState* get_gated_delta_vjp_nax_kernel( metal::Device& d, const std::string& kernel_name, const std::string& hash_name, - const metal::MTLFCList& func_consts); + const metal::MTLFCList& func_consts, + const array& q, + int dk, + int dv, + int hk, + int hv, + int c, + int ckpt, + bool dgamma); // Create a GPU kernel template definition for JIT compilation template diff --git a/mlx/backend/metal/nojit_kernels.cpp b/mlx/backend/metal/nojit_kernels.cpp index ce63dd72b8..c6a153c5b5 100644 --- a/mlx/backend/metal/nojit_kernels.cpp +++ b/mlx/backend/metal/nojit_kernels.cpp @@ -527,19 +527,25 @@ MTL::ComputePipelineState* get_gated_delta_kernel( return d.get_kernel(kernel_name, hash_name, func_consts); } -MTL::ComputePipelineState* get_gated_delta_nax_kernel( +MTL::ComputePipelineState* get_gated_delta_vjp_kernel( metal::Device& d, const std::string& kernel_name, const std::string& hash_name, const metal::MTLFCList& func_consts) { return d.get_kernel(kernel_name, hash_name, func_consts); } - -MTL::ComputePipelineState* get_gated_delta_vjp_kernel( +MTL::ComputePipelineState* get_gated_delta_nax_kernel( metal::Device& d, const std::string& kernel_name, const std::string& hash_name, - const metal::MTLFCList& func_consts) { + const metal::MTLFCList& func_consts, + const array&, + int, + int, + int, + int, + int, + int) { return d.get_kernel(kernel_name, hash_name, func_consts); } @@ -547,7 +553,15 @@ MTL::ComputePipelineState* get_gated_delta_vjp_nax_kernel( metal::Device& d, const std::string& kernel_name, const std::string& hash_name, - const metal::MTLFCList& func_consts) { + const metal::MTLFCList& func_consts, + const array&, + int, + int, + int, + int, + int, + int, + bool) { return d.get_kernel(kernel_name, hash_name, func_consts); } From 71a50adf0abde332b21128345098de4d26655f0c Mon Sep 17 00:00:00 2001 From: mlx-dev Date: Mon, 28 Sep 2026 05:10:46 -0700 Subject: [PATCH 08/11] fixing jit kernels again --- mlx/backend/metal/jit_kernels.cpp | 3 +++ 1 file changed, 3 insertions(+) diff --git a/mlx/backend/metal/jit_kernels.cpp b/mlx/backend/metal/jit_kernels.cpp index 1e0f3cb4be..6b104f5885 100644 --- a/mlx/backend/metal/jit_kernels.cpp +++ b/mlx/backend/metal/jit_kernels.cpp @@ -42,6 +42,9 @@ const char* steel_attention_nax() { const char* gated_delta_update_nax() { return ""; } +const char* gated_delta_update_nax_vjp() { + return ""; +} } // namespace metal #endif // MLX_METAL_NO_NAX From ce364d45dde1b2bd5a50687d9e1df39369015383 Mon Sep 17 00:00:00 2001 From: mlx-dev Date: Mon, 28 Sep 2026 05:51:13 -0700 Subject: [PATCH 09/11] add ops header --- mlx/backend/metal/CMakeLists.txt | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/mlx/backend/metal/CMakeLists.txt b/mlx/backend/metal/CMakeLists.txt index e05faf48f4..5871f2daf4 100644 --- a/mlx/backend/metal/CMakeLists.txt +++ b/mlx/backend/metal/CMakeLists.txt @@ -103,8 +103,8 @@ if(MLX_METAL_JIT) kernels/fp4.h) make_jit_source(steel/attn/kernels/steel_attention_nax) - make_jit_source(gated_delta_update_nax) - make_jit_source(gated_delta_update_nax_vjp) + make_jit_source(gated_delta_update_nax kernels/gated_delta_nax_ops.h) + make_jit_source(gated_delta_update_nax_vjp kernels/gated_delta_nax_ops.h) else() message( From ea4780c718e0f2dd7b91b9f250b016af94772eda Mon Sep 17 00:00:00 2001 From: mlx-dev Date: Mon, 28 Sep 2026 11:18:23 -0700 Subject: [PATCH 10/11] trying to fix CI --- mlx/backend/metal/kernels/CMakeLists.txt | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/mlx/backend/metal/kernels/CMakeLists.txt b/mlx/backend/metal/kernels/CMakeLists.txt index 94191a274a..b208b6df6b 100644 --- a/mlx/backend/metal/kernels/CMakeLists.txt +++ b/mlx/backend/metal/kernels/CMakeLists.txt @@ -174,9 +174,9 @@ if(NOT MLX_METAL_JIT) build_kernel(quantized_nax quantized_nax.h ${STEEL_NAX_HEADERS}) build_kernel(gated_delta_update_nax gated_delta_update_nax.h - ${STEEL_NAX_HEADERS}) + gated_delta_nax_ops.h ${STEEL_NAX_HEADERS}) build_kernel(gated_delta_update_nax_vjp gated_delta_update_nax_vjp.h - ${STEEL_NAX_HEADERS}) + gated_delta_nax_ops.h ${STEEL_NAX_HEADERS}) build_kernel(fp_quantized_nax fp4.h fp8.h fp_quantized_nax.h ${STEEL_NAX_HEADERS}) From 1533bb0b6c7835cef6b94a7247bef81f7a85c671 Mon Sep 17 00:00:00 2001 From: mlx-dev Date: Mon, 28 Sep 2026 12:23:41 -0700 Subject: [PATCH 11/11] gate nax --- mlx/backend/metal/gated_delta_update.cpp | 14 +++++++------- 1 file changed, 7 insertions(+), 7 deletions(-) diff --git a/mlx/backend/metal/gated_delta_update.cpp b/mlx/backend/metal/gated_delta_update.cpp index a089651b29..953c04578c 100644 --- a/mlx/backend/metal/gated_delta_update.cpp +++ b/mlx/backend/metal/gated_delta_update.cpp @@ -15,14 +15,14 @@ namespace mlx::core::fast { namespace { inline int gated_delta_chunk_size(int T) { - int C = metal::is_nax_available() ? 16 : 8; - C = T > 8 ? C : 1; - return env::get_var("GATED_DELTA_CHUNK", C); + const bool nax = metal::is_nax_available(); + const int C = env::get_var("GATED_DELTA_CHUNK", T > 8 ? (nax ? 16 : 8) : 1); + return (C == 16 && !nax) ? 8 : C; } -inline int gated_delta_chunk_size_vjp(int T) { - int C = metal::is_nax_available() ? 16 : 1; - return env::get_var("GATED_DELTA_CHUNK_VJP", C); +inline int gated_delta_chunk_size_vjp() { + const bool nax = metal::is_nax_available(); + return nax ? env::get_var("GATED_DELTA_CHUNK_VJP", 16) : 1; } inline int gated_delta_ckpt(int chunk) { @@ -257,7 +257,7 @@ void GatedDeltaUpdateVJP::eval_gpu( int Dv = v.shape(3); // 16 = chunked NAX, anything else = sequential. - int C = gated_delta_chunk_size_vjp(T); + int C = gated_delta_chunk_size_vjp(); const bool chunked = (C == 16); const int n_chunks = (T + 15) / 16;