diff --git a/benchmarks/python/gated_delta_bench.py b/benchmarks/python/gated_delta_bench.py new file mode 100644 index 0000000000..277897699c --- /dev/null +++ b/benchmarks/python/gated_delta_bench.py @@ -0,0 +1,145 @@ +import argparse +import csv +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) + mx.synchronize() + + s = time.perf_counter_ns() + for _ in range(N_iter_bench): + f(*args) + mx.synchronize() + e = time.perf_counter_ns() + return (e - s) * 1e-9 # total seconds for N_iter_bench * N_iter_func calls + + +def do_kernel_bench(f, *args): + ys = [] + for _ in range(N_iter_func): + out, hf = f(*args) + ys.append(out) + ys.append(hf) + mx.eval(ys) + return ys + + +def benchmark_shape(B, T, Hk, Hv, Dk, Dv, chunk_sizes): + mx.random.seed(42) + 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.random.normal(shape=(B, T, Hv)) * 0.1 - 1.0 + b = mx.sigmoid(mx.random.normal(shape=(B, T, Hv))) + + 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) + / denom + * 1e3 + ) + + speedups = [] + 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) + / denom + * 1e3 + ) + speedups.append(ms_seq / ms_c if ms_c > 0 else float("nan")) + except Exception as ex: + print(f" chunk {C} failed: {ex}") + speedups.append(float("nan")) + + return shape_str, f"{ms_seq:.3f}", speedups, ms_seq + + +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] + else: + Bs = [1, 8, 16] + Ts = [8, 512, 1024, 2048] + Hks = [16] + Hvs = [32] + Dks = [128] + Dvs = [128] + + chunk_sizes = [0, 8, 16] + non_zero_Cs = [C for C in chunk_sizes if C != 0] + + headers = ["B", "T", "Hk", "Hv", "Dk", "Dv", "time_seq (ms)"] + [ + f"C={C} (speedup)" for C in 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))) + + 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( + B, T, Hk, Hv, Dk, Dv, chunk_sizes + ) + row = [f"{B}", f"{T}", f"{Hk}", f"{Hv}", f"{Dk}", f"{Dv}", base_time_s] + for speed in speedups: + row.append(f"{(base_time / speed):<8.2f} ({speed:<5.2f}x)") + + print(fmt.format(*row), end="") + print(f"{RESET}") + + rows.append(row) + + if to_csv: + with open(csv_path, "w", newline="") as f: + writer = csv.writer(f) + writer.writerow(headers) + writer.writerows(rows) + print(f"\nResults also written to {csv_path}") + + +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") + args = parser.parse_args() + + run_benchmark(args.full, to_csv=args.csv, csv_path=args.csv_out) diff --git a/mlx/backend/cuda/primitives.cpp b/mlx/backend/cuda/primitives.cpp index 94f260767f..ad3991a7e5 100644 --- a/mlx/backend/cuda/primitives.cpp +++ b/mlx/backend/cuda/primitives.cpp @@ -24,6 +24,16 @@ namespace mlx::core { throw std::runtime_error(#func " has no CUDA implementation."); \ } +bool fast::GatedDeltaUpdate::use_fallback( + const int Hk, + const int Dk, + const int Hv, + const int Dv, + const bool has_mask, + Stream s) { + return true; +} + NO_GPU_MULTI(LUF) NO_GPU_MULTI(QRF) NO_GPU_MULTI(SVD) @@ -32,6 +42,10 @@ NO_GPU(Cholesky) NO_GPU_MULTI(Eig) NO_GPU_MULTI(Eigh) +namespace fast { +NO_GPU_MULTI(GatedDeltaUpdate) +} + namespace distributed { NO_GPU_MULTI(Send) NO_GPU_MULTI(Recv) diff --git a/mlx/backend/metal/CMakeLists.txt b/mlx/backend/metal/CMakeLists.txt index e7a4d9d2af..8b1266f36f 100644 --- a/mlx/backend/metal/CMakeLists.txt +++ b/mlx/backend/metal/CMakeLists.txt @@ -101,6 +101,7 @@ if(MLX_METAL_JIT) kernels/fp4.h) make_jit_source(steel/attn/kernels/steel_attention_nax) + make_jit_source(gated_delta_update_nax) else() message( @@ -135,6 +136,7 @@ target_sources( ${CMAKE_CURRENT_SOURCE_DIR}/logsumexp.cpp ${CMAKE_CURRENT_SOURCE_DIR}/matmul.cpp ${CMAKE_CURRENT_SOURCE_DIR}/scaled_dot_product_attention.cpp + ${CMAKE_CURRENT_SOURCE_DIR}/gated_delta_update.cpp ${CMAKE_CURRENT_SOURCE_DIR}/metal.cpp ${CMAKE_CURRENT_SOURCE_DIR}/primitives.cpp ${CMAKE_CURRENT_SOURCE_DIR}/quantized.cpp diff --git a/mlx/backend/metal/gated_delta_update.cpp b/mlx/backend/metal/gated_delta_update.cpp new file mode 100644 index 0000000000..ae223c0379 --- /dev/null +++ b/mlx/backend/metal/gated_delta_update.cpp @@ -0,0 +1,200 @@ +// Copyright © 2024 Apple Inc. +#include + +#include "mlx/backend/common/compiled.h" +#include "mlx/backend/gpu/copy.h" +#include "mlx/backend/metal/device.h" +#include "mlx/backend/metal/kernels.h" +#include "mlx/backend/metal/kernels/defines.h" +#include "mlx/backend/metal/utils.h" +#include "mlx/fast_primitives.h" +#include "mlx/utils.h" + +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; + } + + if (has_mask) { + return true; + } + + if (Dk != 128 || Dv != 128) { + return true; + } + + const bool supported_heads = (Hk == 24 && Hv == 24) || + (Hk == 32 && Hv == 32) || (Hk == 16 && Hv == 32) || + (Hk == 16 && Hv == 48); + if (!supported_heads) { + return true; + } + + return false; + + return false; +} + +inline array +ensure_row_contiguous(const array& x, metal::Device& d, const Stream& s) { + if (!x.flags().row_contiguous) { + array x_copy = contiguous_copy_gpu(x, s); + metal::get_command_encoder(s).add_temporary(x_copy); + return x_copy; + } else { + return x; + } +} + +void GatedDeltaUpdate::eval_gpu( + const std::vector& inputs, + std::vector& outputs) { + 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& out = outputs[0]; + auto& hf = outputs[1]; + + 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); + + int C = 1; + const char* threashold_env = std::getenv("GATED_DELTA_THRESH"); + int threshold = threashold_env ? std::stoi(threashold_env) : 16; + if (T > threshold) { + if (metal::is_nax_available()) + C = 16; + else + C = 8; + } + const char* chunk_env = std::getenv("GATED_DELTA_CHUNK"); + C = chunk_env ? std::stoi(chunk_env) : C; + + if (!metal::is_nax_available()) + C = std::min(C, 8); // override in case nax is not available. + + 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); + + switch (C) { + case 16: { + std::string kernel_name = "gated_delta_fused_nax_"; + std::string base_name = kernel_name + suffix; + + base_name += "_" + std::to_string(C); + + std::string hash_name = base_name; + + metal::MTLFCList func_consts = {}; + + auto delta_kernel = + get_gated_delta_nax_kernel(d, base_name, hash_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 / 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 = kernel_name + suffix; + + base_name += "_" + std::to_string(C); + + std::string hash_name = base_name; + + metal::MTLFCList func_consts = {}; + + auto delta_kernel = + get_gated_delta_kernel(d, base_name, hash_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 1: + case 0: { + std::string kernel_name = "seq_gated_delta_"; + std::string base_name = kernel_name + suffix; + std::string hash_name = base_name; + + metal::MTLFCList func_consts = {}; + + auto delta_kernel = + get_gated_delta_kernel(d, base_name, hash_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(g, 3); + compute_encoder.set_input_array(beta, 4); + compute_encoder.set_input_array(h0, 5); + compute_encoder.set_bytes(T, 6); + compute_encoder.set_output_array(out, 7); + compute_encoder.set_output_array(hf, 8); + + 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 8,16 are supported"); + } + } +} + +} // namespace mlx::core::fast diff --git a/mlx/backend/metal/jit/includes.h b/mlx/backend/metal/jit/includes.h index ac9fb81e26..3001a31f1d 100644 --- a/mlx/backend/metal/jit/includes.h +++ b/mlx/backend/metal/jit/includes.h @@ -58,4 +58,7 @@ const char* fp_quantized_nax(); const char* steel_attention_nax(); +const char* gated_delta_update(); +const char* gated_delta_update_nax(); + } // namespace mlx::core::metal diff --git a/mlx/backend/metal/jit_kernels.cpp b/mlx/backend/metal/jit_kernels.cpp index 1384e06c50..b75ab7a874 100644 --- a/mlx/backend/metal/jit_kernels.cpp +++ b/mlx/backend/metal/jit_kernels.cpp @@ -39,6 +39,9 @@ const char* fp_quantized_nax() { const char* steel_attention_nax() { return ""; } +const char* gated_delta_update_nax() { + return ""; +} } // namespace metal #endif // MLX_METAL_NO_NAX @@ -1322,4 +1325,26 @@ MTL::ComputePipelineState* get_steel_attention_nax_kernel( return d.get_kernel(kernel_name, lib, hash_name, 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) { + 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, + 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 973041932a..a3cfc2e6ba 100644 --- a/mlx/backend/metal/kernels.h +++ b/mlx/backend/metal/kernels.h @@ -417,6 +417,18 @@ MTL::ComputePipelineState* get_steel_attention_nax_kernel( int wn, const array& m); +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); + // 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 6d9a0883f0..c899458f64 100644 --- a/mlx/backend/metal/kernels/CMakeLists.txt +++ b/mlx/backend/metal/kernels/CMakeLists.txt @@ -54,6 +54,7 @@ build_kernel(random) build_kernel(rms_norm) build_kernel(rope) build_kernel(scaled_dot_product_attention sdpa_vector.h) +build_kernel(gated_delta_update gated_delta_update_impl.h) if(MLX_METAL_VERSION GREATER_EQUAL 320) build_kernel(fence) endif() @@ -167,6 +168,8 @@ if(NOT MLX_METAL_JIT) ${STEEL_NAX_HEADERS}) 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(fp_quantized_nax fp4.h fp8.h fp_quantized_nax.h ${STEEL_NAX_HEADERS}) diff --git a/mlx/backend/metal/kernels/gated_delta_update.metal b/mlx/backend/metal/kernels/gated_delta_update.metal new file mode 100644 index 0000000000..fba60ed104 --- /dev/null +++ b/mlx/backend/metal/kernels/gated_delta_update.metal @@ -0,0 +1,39 @@ +#include "mlx/backend/metal/kernels/gated_delta_update_impl.h" +#include "mlx/backend/metal/kernels/utils.h" + +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_fused_chunk(in_type, dk, dv, hk, hv, c) \ + instantiate_kernel( \ + "gated_delta_fused_chunk_" #in_type "_" #dk "_" #dv "_" #hk "_" #hv \ + "_" #c, \ + gated_delta_fused_chunk, \ + in_type, \ + dk, \ + dv, \ + hk, \ + 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) + +#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(float); +instantiate_gated_delta(bfloat16_t); \ No newline at end of file diff --git a/mlx/backend/metal/kernels/gated_delta_update_impl.h b/mlx/backend/metal/kernels/gated_delta_update_impl.h new file mode 100644 index 0000000000..27558cdf48 --- /dev/null +++ b/mlx/backend/metal/kernels/gated_delta_update_impl.h @@ -0,0 +1,373 @@ +#pragma once + +#include +#include "mlx/backend/metal/kernels/utils.h" + +#include +#include + +#define FULL_UNROLL _Pragma("clang loop unroll(full)") + +#define AT(TILE, IDX) TILE.thread_elements()[IDX] +#define SUB(TILE0, TILE1, TILE2) \ + { \ + AT(TILE0, 0) = AT(TILE1, 0) - AT(TILE2, 0); \ + AT(TILE0, 1) = AT(TILE1, 1) - AT(TILE2, 1); \ + } +#define ADD(TILE0, TILE1, TILE2) \ + { \ + AT(TILE0, 0) = AT(TILE1, 0) + AT(TILE2, 0); \ + AT(TILE0, 1) = AT(TILE1, 1) + AT(TILE2, 1); \ + } +#define FMA(TILE0, S, TILE1, TILE2) \ + { \ + AT(TILE0, 0) = S * AT(TILE1, 0) + AT(TILE2, 0); \ + AT(TILE0, 1) = S * AT(TILE1, 1) + AT(TILE2, 1); \ + } + +#define SCALE(TILE0, S) \ + { \ + AT(TILE0, 0) *= S; \ + AT(TILE0, 1) *= S; \ + } +#define SCALE2(TILE0, S0, S1) \ + { \ + AT(TILE0, 0) *= S0; \ + AT(TILE0, 1) *= S1; \ + } +#define SCALE_TRI(TILE0, S0, S1) \ + { \ + AT(TILE0, 0) *= fn > fm ? 0.f : S0; \ + AT(TILE0, 1) *= fn + 1 > fm ? 0.f : S1; \ + } +#define SCALE_TRIEQ(TILE0, S0, S1) \ + { \ + AT(TILE0, 0) *= fn >= fm ? 0.f : S0; \ + AT(TILE0, 1) *= fn + 1 >= fm ? 0.f : S1; \ + } + +// lambdas are not supported in metal 14 so porting to macros. + +// non transposed +#define LOAD_M(M, SRC, LD, B) \ + if constexpr (B) { \ + AT(M, 0) = \ + static_cast((fm < valid_rows) ? ((SRC)[fm * (LD) + fn]) : 0.f); \ + AT(M, 1) = static_cast( \ + (fm < valid_rows) ? ((SRC)[fm * (LD) + fn + 1]) : 0.f); \ + } else { \ + AT(M, 0) = static_cast((SRC)[fm * (LD) + fn]); \ + AT(M, 1) = static_cast((SRC)[fm * (LD) + fn + 1]); \ + } + +// transposed load: sequence is the column -> mask fn / fn+1 +#define LOAD_MT(M, SRC, LD, B) \ + if constexpr (B) { \ + AT(M, 0) = \ + static_cast((fn < valid_rows) ? ((SRC)[fn * (LD) + fm]) : 0.f); \ + AT(M, 1) = static_cast( \ + (fn + 1 < valid_rows) ? ((SRC)[(fn + 1) * (LD) + fm]) : 0.f); \ + } else { \ + AT(M, 0) = static_cast((SRC)[fn * (LD) + fm]); \ + AT(M, 1) = static_cast((SRC)[(fn + 1) * (LD) + fm]); \ + } + +// non-transposed load +#define LOAD_M(M, SRC, LD, B) \ + if constexpr (B) { \ + AT(M, 0) = \ + static_cast((fm < valid_rows) ? ((SRC)[fm * (LD) + fn]) : 0.f); \ + AT(M, 1) = static_cast( \ + (fm < valid_rows) ? ((SRC)[fm * (LD) + fn + 1]) : 0.f); \ + } else { \ + AT(M, 0) = static_cast((SRC)[fm * (LD) + fn]); \ + AT(M, 1) = static_cast((SRC)[fm * (LD) + fn + 1]); \ + } + +// transposed load +#define LOAD_MT(M, SRC, LD, B) \ + if constexpr (B) { \ + AT(M, 0) = \ + static_cast((fn < valid_rows) ? ((SRC)[fn * (LD) + fm]) : 0.f); \ + AT(M, 1) = static_cast( \ + (fn + 1 < valid_rows) ? ((SRC)[(fn + 1) * (LD) + fm]) : 0.f); \ + } else { \ + AT(M, 0) = static_cast((SRC)[fn * (LD) + fm]); \ + AT(M, 1) = static_cast((SRC)[(fn + 1) * (LD) + fm]); \ + } + +#define PROCESS_CHUNK_SG(B, S_tile, VALID) \ + { \ + const short valid_rows = (VALID); \ + \ + 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; \ + \ + float gamma_val = simd_prefix_inclusive_sum(g_val); \ + \ + if (thread_index_in_simdgroup < C) { \ + gamma[thread_index_in_simdgroup] = gamma_val; \ + } \ + \ + float gamma_fm = metal::fast::exp(gamma[fm]); \ + float gamma_fmdfn = metal::fast::exp(gamma[fm] - gamma[fn]); \ + float gamma_fmdfn1 = metal::fast::exp(gamma[fm] - gamma[fn + 1]); \ + float gamma_Cdfn = metal::fast::exp(gamma[C - 1] - gamma[fn]); \ + float gamma_Cdfn1 = metal::fast::exp(gamma[C - 1] - gamma[fn + 1]); \ + float gamma_C = metal::fast::exp(gamma[C - 1]); \ + \ + float beta_fm = (fm < valid_rows) ? beta_[fm * Hv + hv_idx] : 0.0f; \ + \ + KKt_tile = make_filled_simdgroup_matrix(0.f); \ + FULL_UNROLL \ + for (int kk = 0; kk < Dk; kk += 8) { \ + LOAD_M(K_tile, k_ + kk, Dk * Hk, B) \ + LOAD_MT(KT_tile, k_ + kk, Dk * Hk, B) \ + simdgroup_multiply_accumulate(KKt_tile, K_tile, KT_tile, KKt_tile); \ + } \ + \ + KKtK_tile = KKt_tile; \ + SCALE_TRIEQ(KKtK_tile, beta_fm, beta_fm) \ + \ + simdgroup_float8x8 Tinv, P; \ + AT(P, 0) = AT(KKtK_tile, 0); \ + AT(P, 1) = AT(KKtK_tile, 1); \ + SUB(Tinv, I_tile, KKtK_tile) \ + \ + FULL_UNROLL \ + for (int step = 1; (1 << step) < C; step++) { \ + simdgroup_multiply(P, P, P); \ + simdgroup_multiply_accumulate(Tinv, Tinv, P, Tinv); \ + } \ + \ + WS_tile = make_filled_simdgroup_matrix(0.f); \ + FULL_UNROLL \ + for (int kk = 0; kk < Dk; kk += 8) { \ + LOAD_M(K_tile, k_ + kk, Dk * Hk, B) \ + SCALE(K_tile, beta_fm) \ + simdgroup_multiply(W_tile, Tinv, K_tile); \ + SCALE(W_tile, gamma_fm) \ + simdgroup_multiply_accumulate(WS_tile, W_tile, S_tile[kk / 8], WS_tile); \ + } \ + \ + SCALE_TRI(Tinv, gamma_fmdfn, gamma_fmdfn1) \ + \ + LOAD_M(V_tile, v_ + dv_idx, Dv * Hv, B) \ + SCALE(V_tile, beta_fm) \ + simdgroup_multiply(U_tile, Tinv, V_tile); \ + SUB(delta_tile, U_tile, WS_tile) \ + \ + tmp_tile = make_filled_simdgroup_matrix(0.f); \ + QKt_tile = make_filled_simdgroup_matrix(0.f); \ + FULL_UNROLL \ + for (int kk = 0; kk < Dk; kk += 8) { \ + LOAD_M(Q_tile, q_ + kk, Hk * Dk, B) \ + LOAD_MT(K_tile, k_ + kk, Hk * Dk, B) \ + simdgroup_multiply_accumulate(QKt_tile, Q_tile, K_tile, QKt_tile); \ + SCALE(Q_tile, gamma_fm) \ + simdgroup_multiply_accumulate( \ + tmp_tile, Q_tile, S_tile[kk / 8], tmp_tile); \ + } \ + \ + SCALE_TRI(QKt_tile, gamma_fmdfn, gamma_fmdfn1) \ + \ + simdgroup_multiply_accumulate(out_tile, QKt_tile, delta_tile, tmp_tile); \ + \ + if (fm < valid_rows) { \ + y[fm * Hv * Dv + dv_idx + fn] = static_cast(AT(out_tile, 0)); \ + y[fm * Hv * Dv + dv_idx + fn + 1] = static_cast(AT(out_tile, 1)); \ + } \ + \ + FULL_UNROLL \ + for (int kk = 0; kk < Dk; kk += 8) { \ + LOAD_MT(K_tile, k_ + kk, Hk * Dk, B) \ + SCALE2(K_tile, gamma_Cdfn, gamma_Cdfn1) \ + simdgroup_multiply(KD_tile, K_tile, delta_tile); \ + FMA(S_tile[kk / 8], gamma_C, S_tile[kk / 8], KD_tile) \ + } \ + } + +template +[[kernel]] void gated_delta_fused_chunk( + 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)]], + const device InT* beta [[buffer(5)]], + device InT* y [[buffer(6)]], + device float* state_out [[buffer(7)]], + constant int& T [[buffer(8)]], + 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); + + const short qid = thread_index_in_simdgroup / 4; + const short fm = (qid & 4) + + ((thread_index_in_simdgroup / 2) % 4); // row coordinate of the held tile + const short fn = (qid & 2) * 2 + + (thread_index_in_simdgroup % 2) * 2; // column coordinate of the held tile + + auto dv_idx = thread_position_in_grid.y * 8; + const short sg_id = thread_position_in_threadgroup.y; // 0..3 + +#define OUTPUT(T) \ + if (true) { \ + simdgroup_store(T, y); \ + return; \ + } + // set up pointers + // g: [B, T, Hv] + auto g_ = g + b_idx * T * Hv; + + // q, k: [B, T, Hk, Dk] + auto q_ = q + b_idx * T * Hk * Dk + hk_idx * Dk; + auto k_ = k + b_idx * T * Hk * Dk + hk_idx * Dk; + + // v, y: [B, T, Hv, Dv] + y += b_idx * T * Hv * Dv + hv_idx * Dv; + auto v_ = v + b_idx * T * Hv * Dv + hv_idx * Dv; + auto beta_ = beta + b_idx * T * Hv; + + // state_in, state_out: [B, Hv, Dv, Dk] + auto i_state = state_in + (n * Dv + dv_idx) * Dk; + auto o_state = state_out + (n * Dv + dv_idx) * Dk; + + simdgroup_float8x8 S_tile[Dk / 8]; + + // simdgroup matrices + simdgroup_float8x8 V_tile, K_tile, KT_tile, Q_tile; + simdgroup_float8x8 W_tile, U_tile; + simdgroup_float8x8 WS_tile; + simdgroup_float8x8 delta_tile; + simdgroup_float8x8 tmp_tile; + simdgroup_float8x8 QKt_tile; + simdgroup_float8x8 out_tile; + simdgroup_float8x8 KD_tile; + + // tiles for WY form computation + simdgroup_float8x8 KKtK_tile, KKt_tile; + + threadgroup float gamma_all[C * 4]; + threadgroup float* gamma = gamma_all + sg_id * C; + + simdgroup_float8x8 I_tile = make_filled_simdgroup_matrix(0.f); + AT(I_tile, 0) = (fm == fn) ? 1.0f : 0.0f; + AT(I_tile, 1) = (fm == fn + 1) ? 1.0f : 0.0f; + + // load initial state into registers + for (int kk = 0; kk < Dk; kk += 8) { + simdgroup_load(S_tile[kk / 8], i_state + kk, Dk, ulong2(0, 0), true); + } + + int t = 0; + FULL_UNROLL + for (; t + C <= T; t += C) { + PROCESS_CHUNK_SG(false, S_tile, C); + q_ += C * Hk * Dk; + k_ += C * Hk * Dk; + v_ += C * Hv * Dv; + beta_ += C * Hv; + y += C * Hv * Dv; + g_ += C * Hv; + } + if (t < T) { + PROCESS_CHUNK_SG(true, S_tile, short(T - t)); + } + + FULL_UNROLL + for (int kk = 0; kk < Dk; kk += 8) { + simdgroup_store(S_tile[kk / 8], o_state + kk, Dk, ulong2(0, 0), true); + } +} + +/* + auto grid = MTL::Size(32, Dv, B * Hv); + auto threads = MTL::Size(32, 4, 1); + */ +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 InT* g [[buffer(3)]], // [B, T, Hv] or [B, T, Hv, Dk] + const device InT* beta [[buffer(4)]], // [B, T, Hv] + const device float* state_in [[buffer(5)]], // [B, Hv, Dv, Dk] + constant int& T [[buffer(6)]], + device InT* y [[buffer(7)]], // [B, T, Hv, Dv] + device float* state_out [[buffer(8)]], // [B, Hv, Dv, Dk] + 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; + auto hk_idx = hv_idx / (Hv / Hk); + constexpr int n_per_t = Dk / 32; + + // q, k: [B, T, Hk, Dk] + auto q_ = q + b_idx * T * Hk * Dk + hk_idx * Dk; + auto k_ = k + b_idx * T * Hk * Dk + hk_idx * Dk; + + // v, y: [B, T, Hv, Dv] + auto v_ = v + b_idx * T * Hv * Dv + hv_idx * Dv; + y += b_idx * T * Hv * Dv + hv_idx * Dv; + + auto dk_idx = thread_position_in_threadgroup.x; + auto dv_idx = thread_position_in_grid.y; + + // state_in, state_out: [B, Hv, Dv, Dk] + auto i_state = state_in + (n * Dv + dv_idx) * Dk; + auto o_state = state_out + (n * Dv + dv_idx) * Dk; + + float state[n_per_t]; + for (int i = 0; i < n_per_t; ++i) { + auto s_idx = n_per_t * dk_idx + i; + state[i] = static_cast(i_state[s_idx]); + } + + // g: [B, T, Hv] + auto g_ = g + b_idx * T * Hv; + auto beta_ = beta + b_idx * T * Hv; + + for (int t = 0; t < T; ++t) { + float kv_mem = 0.0f; + for (int i = 0; i < n_per_t; ++i) { + auto s_idx = n_per_t * dk_idx + i; + state[i] = state[i] * g_[hv_idx]; + kv_mem += state[i] * k_[s_idx]; + } + kv_mem = simd_sum(kv_mem); + + auto delta = (v_[dv_idx] - kv_mem) * beta_[hv_idx]; + + float out = 0.0f; + for (int i = 0; i < n_per_t; ++i) { + auto s_idx = n_per_t * dk_idx + i; + state[i] = state[i] + k_[s_idx] * delta; + out += state[i] * q_[s_idx]; + } + out = simd_sum(out); + if (thread_index_in_simdgroup == 0) { + y[dv_idx] = static_cast(out); + } + // Increment data pointers to next time step + q_ += Hk * Dk; + k_ += Hk * Dk; + v_ += Hv * Dv; + y += Hv * Dv; + g_ += Hv; + beta_ += Hv; + } + for (int i = 0; i < n_per_t; ++i) { + auto s_idx = n_per_t * dk_idx + i; + o_state[s_idx] = static_cast(state[i]); + } +} \ No newline at end of file diff --git a/mlx/backend/metal/kernels/gated_delta_update_nax.h b/mlx/backend/metal/kernels/gated_delta_update_nax.h new file mode 100644 index 0000000000..72392bcba6 --- /dev/null +++ b/mlx/backend/metal/kernels/gated_delta_update_nax.h @@ -0,0 +1,523 @@ +#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; + +// 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; + } + + 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 + +#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 +[[kernel]] void gated_delta_fused_nax( + 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)]], + const device InT* beta [[buffer(5)]], + device InT* y [[buffer(6)]], + device float* state_out [[buffer(7)]], + constant int& T [[buffer(8)]], + 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); + + 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)); + + // set up pointers + // g: [B, T, Hv] + auto g_ = g + b_idx * T * Hv; + + // q, k: [B, T, Hk, Dk] + auto q_ = q + b_idx * T * Hk * Dk + hk_idx * Dk; + auto k_ = k + b_idx * T * Hk * Dk + hk_idx * Dk; + + // v, y: [B, T, Hv, Dv] + y += b_idx * T * Hv * Dv + hv_idx * Dv; + auto v_ = v + b_idx * T * Hv * Dv + hv_idx * Dv; + auto beta_ = beta + b_idx * T * Hv; + + // state_in, state_out: [B, Hv, Dv, Dk] + auto i_state = state_in + (n * Dv + dv_idx) * Dk; + auto o_state = state_out + (n * Dv + dv_idx) * Dk; + + threadgroup float gamma_all[C * 4]; + threadgroup float* gamma = gamma_all + sg_id * C; + + float beta_fm[2]; + + mlx::steel::NAXTile S_tile; + S_tile.load(i_state, Dk); + + mlx::steel::NAXTile K_tile, Q_tile; + mlx::steel::NAXTile W_tile; // panel + + mlx::steel::NAXTile V_tile; + mlx::steel::NAXTile U_tile; + mlx::steel::NAXTile WS_tile; + mlx::steel::NAXTile delta_tile; + mlx::steel::NAXTile tmp_tile; + mlx::steel::NAXTile QKt_tile; + mlx::steel::NAXTile out_tile; + mlx::steel::NAXTile Tinv_tile, P; + + mlx::steel::NAXTile KKtK_tile, KKt_tile; + + mlx::steel::NAXTile 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} */ + const short _fn = _c.x; + const short _fm = _c.y; + AT_NAX(I_tile, _i) = (_fn == _fm) ? 1.0f : 0.0f; + } + mlx::steel::NAXTile TMP_tile; + + 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); + } + }; + + 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); + } + + 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; + + KKt_tile.clear(); + for (int kk = 0; kk < Dk; kk += 32) { + load_seq(K_tile, k_ + kk, Dk * Hk); + MMA16x16x32(KKt_tile, 0, K_tile, false, 0, K_tile, true, 0); + } + + KKtK_tile = KKt_tile; + + SCALE_TRIEQ_NAX1(KKtK_tile, beta_fm); + SUB_NAX(Tinv_tile, I_tile, KKtK_tile); + STEEL_PRAGMA_UNROLL + for (int step = 0; step < 15; step++) { + MM16x16x16(TMP_tile, 0, KKtK_tile, false, 0, Tinv_tile, false, 0); + SUB_NAX(Tinv_tile, I_tile, TMP_tile); + } + + STEEL_PRAGMA_UNROLL + for (short nn = 0; nn < Dk / 16; nn += 2) { + load_seq(K_tile, k_ + nn * 16, Dk * Hk); + SCALE_BETA_NAX(K_tile, beta_fm); + MM16x32x16(W_tile, nn, Tinv_tile, false, 0, K_tile, false, 0); + } + SCALE_ROW_NAX(W_tile, gamma) + + SCALE_TRI_NAX(Tinv_tile, gamma) + load_seq(V_tile, v_ + dv_idx, Dv * Hv); + SCALE_BETA_NAX(V_tile, beta_fm); + MM16x16x16(U_tile, 0, Tinv_tile, false, 0, V_tile, false, 0) + + WS_tile.clear(); + STEEL_PRAGMA_UNROLL + for (short kk = 0; kk < Dk / 16; kk += 2) { + MMA16x16x32(WS_tile, 0, W_tile, false, kk, S_tile, true, kk) + } + + SUB_NAX(delta_tile, U_tile, WS_tile) + + tmp_tile.clear(); + QKt_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) + + out_tile = tmp_tile; + MMA16x16x16(out_tile, 0, QKt_tile, false, 0, delta_tile, false, 0); + + STEEL_PRAGMA_UNROLL + for (short _i = 0; _i < decltype(out_tile)::kElemsPerFrag; _i++) { + const short2 _c = mlx::steel::BaseNAXFrag::get_coord(_i); // {fn, fm} + const short _fn = _c.x; + const short _fm = _c.y; + if (_fm < valid_rows) { + y[_fm * Hv * Dv + dv_idx + _fn] = + static_cast(AT_NAX(out_tile, _i)); + } + } + + SCALE_NAX(S_tile, metal::fast::exp(gamma[C - 1])); + + for (int kk = 0; kk < Dk; kk += 32) { + load_seq(K_tile, k_ + kk, Hk * Dk); + SCALE2_NAX(K_tile, gamma); + MMA16x32x16(S_tile, kk / 16, delta_tile, true, 0, K_tile, false, 0); + } + }; + + int t = 0; + for (; t + C <= T; t += C) { + process_chunk(C, metal::false_type{}); + q_ += C * Hk * Dk; + k_ += C * Hk * Dk; + v_ += C * Hv * Dv; + beta_ += C * Hv; + y += C * Hv * Dv; + g_ += C * Hv; + } + if (t < T) { + process_chunk(short(T - t), metal::true_type{}); + } + + 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 new file mode 100644 index 0000000000..b17707fb40 --- /dev/null +++ b/mlx/backend/metal/kernels/gated_delta_update_nax.metal @@ -0,0 +1,28 @@ +#include "mlx/backend/metal/kernels/gated_delta_update_nax.h" +#include "mlx/backend/metal/kernels/utils.h" + +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_dims(in_type, dk, dv, hk, hv) \ + instantiate_gated_delta_update_fused_nax(in_type, dk, dv, hk, hv, 16) + +#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(float); +instantiate_gated_delta(bfloat16_t); \ No newline at end of file diff --git a/mlx/backend/metal/nojit_kernels.cpp b/mlx/backend/metal/nojit_kernels.cpp index 9f6f8782f5..2ea601863c 100644 --- a/mlx/backend/metal/nojit_kernels.cpp +++ b/mlx/backend/metal/nojit_kernels.cpp @@ -492,4 +492,20 @@ MTL::ComputePipelineState* get_steel_attention_nax_kernel( return d.get_kernel(kernel_name, hash_name, 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) { + 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, + 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/backend/no_gpu/primitives.cpp b/mlx/backend/no_gpu/primitives.cpp index 0e05e9d19f..c4067bcd8f 100644 --- a/mlx/backend/no_gpu/primitives.cpp +++ b/mlx/backend/no_gpu/primitives.cpp @@ -36,6 +36,16 @@ bool fast::ScaledDotProductAttention::use_fallback( return true; } +bool fast::GatedDeltaUpdate::use_fallback( + const int Hk, + const int Dk, + const int Hv, + const int Dv, + const bool has_mask, + Stream s) { + return true; +} + bool fast::ScaledDotProductAttention::supports_bool_mask() { return false; } @@ -170,6 +180,7 @@ NO_GPU_MULTI(RMSNormVJP) NO_GPU_USE_FALLBACK(RoPE) NO_GPU_MULTI(ScaledDotProductAttention) NO_GPU_MULTI(ScaledDotProductAttentionVJP) +NO_GPU_MULTI(GatedDeltaUpdate) NO_GPU_MULTI(ConvertFP8) NO_GPU_MULTI(Quantize) NO_GPU_MULTI(CustomKernel) diff --git a/mlx/fast.cpp b/mlx/fast.cpp index a668fe9abd..e552868799 100644 --- a/mlx/fast.cpp +++ b/mlx/fast.cpp @@ -922,6 +922,132 @@ bool ScaledDotProductAttentionVJP::is_equivalent(const Primitive& other) const { has_sinks_ == a_other.has_sinks_; } +std::vector gated_delta_update( + const array& queries, + const array& keys, + const array& values, + const array& gates, + const array& beta_, + const std::optional& initial_state, /* = std::nullopt */ + const std::optional& mask_, /* = std::nullopt */ + StreamOrDevice s_ /* = {} */) { + // determine output dtype + auto s = to_stream(s_); + + auto promoted = promote_types(queries.dtype(), keys.dtype()); + auto out_dtype = issubdtype(promoted, floating) + ? promoted + : promote_types(promoted, float32); + + // cast all inputs + auto q = astype(queries, out_dtype, s); + auto k = astype(keys, out_dtype, s); + auto v = astype(values, out_dtype, s); + auto g = astype(gates, out_dtype, s); + auto beta = astype(beta_, out_dtype, 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); + + auto h0 = initial_state.has_value() ? astype(*initial_state, float32, s) + : zeros({B, Hv, Dv, Dk}, float32, s); + + bool has_mask = mask_.has_value(); + auto mask = has_mask ? astype(*mask_, bool_, s) : array(false); + + auto fallback = [B, T, Hk, Dk, Hv, Dv, has_mask, s]( + std::vector inputs) { + auto q = astype(inputs[0], float32, s); + auto k = astype(inputs[1], float32, s); + auto v = astype(inputs[2], float32, s); + auto g = astype(inputs[3], float32, s); + auto beta = astype(inputs[4], float32, s); + auto state = astype(inputs[5], float32, s); + + if (Hv != Hk) { + int repeat_factor = Hv / Hk; + q = repeat(q, repeat_factor, 2, s); + k = repeat(k, repeat_factor, 2, s); + } + + array mask = has_mask ? astype(inputs[6], bool_, s) : array(false); + const array zero = array(0.0f, float32); + + std::vector outputs; + for (int t = 0; t < T; t++) { + auto get_t = [&](const array& a, int t) { + Shape start(a.ndim(), 0), stop = a.shape(); + start[1] = t; + stop[1] = t + 1; + return squeeze(slice(a, start, stop, s), 1, s); + }; + auto q_t = get_t(q, t); + auto k_t = get_t(k, t); + auto v_t = get_t(v, t); + auto g_t = get_t(g, t); + auto beta_t = get_t(beta, t); + array mask_t = has_mask ? get_t(mask, t) : array(true); + + auto state_prev = state; + + auto decay = (g_t.ndim() == 2) + ? expand_dims(g_t, {-1, -2}, s) // [B,H,1,1] + : expand_dims(g_t, -2, s); // [B,H,1,Dk] + + auto state_next = multiply(state_prev, decay, s); + + auto kv = + sum(multiply(state_next, expand_dims(k_t, -2, s), s), + -1, + false, + s); // [B,H,Dv] + + auto delta = subtract(v_t, kv, s); + delta = multiply(delta, expand_dims(beta_t, -1, s), s); + + state_next = + add(state_next, + multiply(expand_dims(delta, -1, s), expand_dims(k_t, -2, s), s), + s); + + if (has_mask) { + auto state_mask = expand_dims(mask_t, {-1, -2, -3}, s); + state = where(state_mask, state_next, state_prev, s); + } else { + state = state_next; + } + + auto o_t = sum(multiply(state, expand_dims(q_t, -2, s), s), -1, false, s); + + if (has_mask) { + auto out_mask = expand_dims(mask_t, {-1, -2}, s); + o_t = where(out_mask, o_t, zero, s); + } + outputs.push_back(o_t); + } + auto out = stack(outputs, 1, s); + return std::vector{out, state}; + }; + + if (!GatedDeltaUpdate::use_fallback(Hk, Dk, Hv, Dv, has_mask, s)) { + auto result = array::make_arrays( + /* output shapes */ {{B, T, Hv, Dv}, {B, Hv, Dv, Dk}}, + /* dtypes */ {out_dtype, float32}, + /* primitive */ + std::make_shared(s, fallback), + /* inputs */ {q, k, v, g, beta, h0}); + + return result; + } + + auto result = fallback({q, k, v, g, beta, h0, mask}); + return result; +} + bool Quantize::is_equivalent(const Primitive& other) const { const Quantize& p_other = static_cast(other); return ( diff --git a/mlx/fast.h b/mlx/fast.h index 934fadc2b7..052a23811d 100644 --- a/mlx/fast.h +++ b/mlx/fast.h @@ -55,6 +55,16 @@ MLX_API array scaled_dot_product_attention( const std::optional& sinks = {}, StreamOrDevice s = {}); +MLX_API std::vector gated_delta_update( + const array& queries, + const array& keys, + const array& values, + const array& gates, + const array& beta_, + const std::optional& initial_state = std::nullopt, + const std::optional& mask = std::nullopt, + StreamOrDevice s = {}); + using TemplateArg = std::variant; using ScalarArg = std::variant; diff --git a/mlx/fast_primitives.h b/mlx/fast_primitives.h index 0d2f861045..9bd3954b2a 100644 --- a/mlx/fast_primitives.h +++ b/mlx/fast_primitives.h @@ -325,6 +325,38 @@ class ConvertFP8 : public Primitive { bool to_fp8_; }; +class GatedDeltaUpdate : public Custom { + public: + GatedDeltaUpdate( + 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, + const bool has_mask, + 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(GatedDeltaUpdate); + DEFINE_INPUT_OUTPUT_SHAPE() + auto state() const { + return std::make_tuple(nullptr); /* TODO */ + } + + private: +}; + class Quantize : public Custom { public: explicit Quantize( diff --git a/python/src/fast.cpp b/python/src/fast.cpp index cd30b0bacd..7021e361a9 100644 --- a/python/src/fast.cpp +++ b/python/src/fast.cpp @@ -333,6 +333,32 @@ void init_fast(nb::module_& parent_module) { out = mx.fast.scaled_dot_product_attention(q, k, v, scale=scale, mask="causal") )pbdoc"); + m.def( + "gated_delta_update", + &mlx::core::fast::gated_delta_update, + "q"_a, + "k"_a, + "v"_a, + "gamma"_a, + "beta"_a, + "initial_state"_a = nb::none(), // optional, defaults to None + "mask"_a = nb::none(), // optional, defaults to None + "stream"_a = nb::none(), // optional, defaults to None + R"( + Chunked gated delta network forward pass. + + Args: + q: Queries [B, H, T, Dk] + k: Keys [B, H, T, Dk] + v: Values [B, H, T, Dv] + gamma: + beta: Delta update rates [B, H, T] + initial_state: Optional initial hidden state [B, H, Dk, Dv] + mask: Optional + Returns: + Tuple of (output [B, H, T, Dv], final_state [B, H, Dk, Dv]) + )"); + m.def( "metal_kernel", [](const std::string& name, diff --git a/python/tests/test_fast_gated_delta.py b/python/tests/test_fast_gated_delta.py new file mode 100644 index 0000000000..b2334dcdbb --- /dev/null +++ b/python/tests/test_fast_gated_delta.py @@ -0,0 +1,237 @@ +import os +import unittest + +import mlx.core as mx +import mlx_tests +import numpy as np + +try: + import torch + + has_torch = True +except ImportError as e: + has_torch = False + + +def gated_delta_oracle( + q, + k, + v, + beta, + g, + scale=None, + initial_state=None, + output_final_state=False, +): + """ + Reference PyTorch implementation of recurrent gated delta rule. + Taken from: https://github.com/fla-org/flash-linear-attention/blob/main/fla/ops/gated_delta_rule/naive.py + + Args: + q: [B, T, H, K] + k: [B, T, H, K] + v: [B, T, H, V] + beta: [B, T, H] + g: [B, T, H] <--- Difference: with out kernel: this is expected as a log. + scale: float, optional <--- Difference: This is done by the qwen3_5 model. + initial_state: [B, H, K, V], optional <- Difference: last two dimensions are transposed. + output_final_state: bool + + Returns: + o: [B, T, H, V] + final_state: [B, H, K, V] if output_final_state else None + """ + q, k, v, beta, g = map( + lambda x: x.transpose(1, 2).contiguous().to(torch.float32), [q, k, v, beta, g] + ) + B, H, T, K, V = *k.shape, v.shape[-1] + o = torch.zeros(B, H, T, V).to(v) + h = torch.zeros(B, H, K, V).to(v) + if initial_state is not None: + h = initial_state.to(torch.float32) + if scale is None: + scale = 1 / (q.shape[-1] ** 0.5) + q = q * scale + + for i in range(T): + b_q = q[:, :, i] + b_k = k[:, :, i] + b_v = v[:, :, i].clone() + h = h.clone() * g[:, :, i].exp()[..., None, None] + b_beta = beta[:, :, i] + b_v = b_v - (h.clone() * b_k[..., None]).sum(-2) + b_v = b_v * b_beta[..., None] + h = h.clone() + b_k.unsqueeze(-1) * b_v.unsqueeze(-2) + o[:, :, i] = torch.einsum("bhd,bhdm->bhm", b_q, h) + + if not output_final_state: + h = None + o = o.transpose(1, 2).contiguous() + return o, h + + +def runner(dims, stream=mx.gpu, reference=True): + B, Hk, Hv, T, Dk, Dv = dims + + assert Hv % Hk == 0 + repeat_factor = Hv // Hk + + 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.random.uniform(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) + + if reference: + # Prepare reference inputs + qpt = torch.from_numpy(np.array(q)) + kpt = torch.from_numpy(np.array(k)) + vpt = torch.from_numpy(np.array(v)) + bpt = torch.from_numpy(np.array(b)) + gpt = torch.from_numpy(np.array(g)) + h0pt = torch.from_numpy(np.array(h0)).transpose(-1, -2).contiguous() + + if repeat_factor > 1: + qpt = qpt.repeat_interleave(repeat_factor, dim=2) + kpt = kpt.repeat_interleave(repeat_factor, dim=2) + + out_on_py, hf_on_py = gated_delta_oracle( + qpt, + kpt, + vpt, + bpt, + torch.log(gpt), + scale=1.0, + initial_state=h0pt, + output_final_state=True, + ) + + out_on = mx.array(out_on_py.detach().cpu().numpy()) # [B, T, Hv, Dv] + hf_on = mx.swapaxes( + mx.array(hf_on_py.detach().cpu().numpy()), -1, -2 + ) # -> [B, Hv, Dv, Dk] + out_ref = mx.array(out_on) + hf_ref = mx.array(hf_on) + else: + # use fallback for tests once fallback is validated + out_ref, hf_ref = mx.fast.gated_delta_update( + q, k, v, g, b, initial_state=h0, stream=mx.cpu + ) + + mx.eval(out_ref, hf_ref) + + out, hf = mx.fast.gated_delta_update(q, k, v, g, b, initial_state=h0, stream=stream) + + mx.eval(out, hf) + return (out, hf), (out_ref, hf_ref) + + +class TestGatedDelta(mlx_tests.MLXTestCase): + base_dims = (1, 32, 32, 1, 128, 128) + unaligned_dims = (1, 32, 32, 33, 128, 128) + big_batch_dims = (128, 32, 32, 16, 128, 128) + large_t_dims = (2, 32, 32, 1111, 128, 128) + diff_heads = (1, 16, 32, 33, 128, 128) + diff_heads2 = (1, 16, 48, 33, 128, 128) + + fallback_dims = [base_dims, unaligned_dims, big_batch_dims, diff_heads, diff_heads2] + gpu_dims = fallback_dims + [large_t_dims] + + @unittest.skipIf(not has_torch, "requires Torch") + def test_gated_delta_fallback(self): + for dims in self.fallback_dims: + (out, hf), (out_ref, hf_ref) = runner(dims, mx.cpu) + msg = f"Failed on Dimensions: {dims}" + self.assertTrue( + mx.allclose(out_ref, out, atol=1e-4, rtol=1e-4), msg="Out " + msg + ) + self.assertTrue( + mx.allclose(hf_ref, hf, atol=1e-4, rtol=1e-4), msg="State " + msg + ) + + def test_gated_delta_fallback_masked(self): + for dims in self.fallback_dims: + + B, Hk, Hv, T, Dk, Dv = dims + + 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.random.uniform(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) + + # make a mask + lengths = mx.random.randint(1, T + 1, shape=(B,)) + mask = mx.arange(T)[None, :] < lengths[:, None] + # mask one input in python + mask_float = mask.astype(q.dtype) + km = k * mask_float[..., None, None] + vm = v * mask_float[..., None, None] + qm = q * mask_float[..., None, None] + bm = b * mask_float[..., None] + gm = mx.where(mask[..., None], g, 1.0) + + out_ref, hf_ref = mx.fast.gated_delta_update( + qm, km, vm, gm, bm, initial_state=h0, stream=mx.cpu + ) + + mx.eval(out_ref, hf_ref) + out, hf = mx.fast.gated_delta_update( + q, k, v, g, b, initial_state=h0, mask=mask, stream=mx.cpu + ) + mx.eval(out, hf) + + msg = f"Failed on Dimensions: {dims}" + self.assertTrue( + mx.allclose(out_ref, out, atol=1e-4, rtol=1e-4), msg="Out " + msg + ) + self.assertTrue( + mx.allclose(hf_ref, hf, atol=1e-4, rtol=1e-4), msg="State " + msg + ) + + @unittest.skipIf(not mx.is_available(mx.gpu), "No GPU available") + def test_gated_delta_sequential(self): + os.environ["GATED_DELTA_CHUNK"] = "0" + for dims in self.gpu_dims: + (out, hf), (out_ref, hf_ref) = runner(dims, reference=False) + msg = f"Failed on Dimensions: {dims}" + self.assertTrue( + mx.allclose(out_ref, out, atol=1e-4, rtol=1e-4), msg="Out " + msg + ) + self.assertTrue( + mx.allclose(hf_ref, hf, atol=1e-4, rtol=1e-4), msg="State " + msg + ) + + @unittest.skipIf(not mx.is_available(mx.gpu), "No GPU available") + def test_gated_delta_simdgroup(self): + os.environ["GATED_DELTA_CHUNK"] = "8" + for dims in self.gpu_dims: + (out, hf), (out_ref, hf_ref) = runner(dims, reference=False) + msg = f"Failed on Dimensions: {dims}" + self.assertTrue( + mx.allclose(out_ref, out, atol=1e-4, rtol=1e-4), msg="Out " + msg + ) + self.assertTrue( + mx.allclose(hf_ref, hf, atol=1e-4, rtol=1e-4), msg="State " + msg + ) + + @unittest.skipIf(not mx.is_available(mx.gpu), "No GPU available") + def test_gated_delta_nax(self): + os.environ["GATED_DELTA_CHUNK"] = "16" + for dims in self.gpu_dims: + (out, hf), (out_ref, hf_ref) = runner(dims, reference=False) + msg = f"Failed on Dimensions: {dims}" + self.assertTrue( + mx.allclose(out_ref, out, atol=1e-1, rtol=1e-4), msg="Out " + msg + ) + self.assertTrue( + mx.allclose(hf_ref, hf, atol=1e-1, rtol=1e-4), msg="State " + msg + ) + + +if __name__ == "__main__": + mlx_tests.MLXTestRunner(failfast=True)