From 7fba477d165335343170652b3981c87ee19ca499 Mon Sep 17 00:00:00 2001 From: Zhongrui Sun Date: Sat, 12 Sep 2026 20:22:54 -0700 Subject: [PATCH 1/2] Add fused Metal kernels for fast.cross_entropy On Metal, CrossEntropy::use_fallback returned true and both eval_gpu overloads threw NYI, so mx.fast.cross_entropy always ran the unfused logsumexp - take_along_axis graph. Add the forward and VJP kernels and enable them, following the CUDA implementation from #3947. Forward: one threadgroup per row, single-pass online logsumexp with float32 accumulation for every input dtype. The loss is formed as (max - x_t) + log(normalizer) so the two close values are subtracted first. The host shrinks the threadgroup to ceil(V / N_READS) rounded to a SIMD multiple, so short rows take one iteration of the same looped kernel; the cross-SIMD reduction only reads the slots that were written. VJP: exp((x - x_t) - loss) is softmax(x), so the backward pass needs no reduction and the one-hot target is never materialized. When the logits buffer can be donated the gradient is written in place, with a device memory barrier between the reads of x_t and the first write. Negative targets wrap, matching the take_along_axis fallback this replaces. The JIT library name is passed explicitly because deriving it from the kernel name would drop the "cross_" prefix. The float16/bfloat16 tolerances in test_cross_entropy tighten to 1e-3 on the GPU. The fallback fails this because its logsumexp runs in the input dtype; the fused kernels are within ~4e-6 for all three dtypes. The CPU path keeps the old tolerances. --- mlx/backend/metal/CMakeLists.txt | 2 + mlx/backend/metal/cross_entropy.cpp | 125 ++++++++++++++++ mlx/backend/metal/jit/includes.h | 1 + mlx/backend/metal/jit_kernels.cpp | 20 +++ mlx/backend/metal/kernels.h | 5 + mlx/backend/metal/kernels/CMakeLists.txt | 1 + mlx/backend/metal/kernels/cross_entropy.h | 133 ++++++++++++++++++ mlx/backend/metal/kernels/cross_entropy.metal | 18 +++ mlx/backend/metal/kernels/defines.h | 1 + mlx/backend/metal/nojit_kernels.cpp | 7 + mlx/backend/metal/primitives.cpp | 23 --- python/src/fast.cpp | 4 +- python/tests/test_fast.py | 3 + 13 files changed, 318 insertions(+), 25 deletions(-) create mode 100644 mlx/backend/metal/cross_entropy.cpp create mode 100644 mlx/backend/metal/kernels/cross_entropy.h create mode 100644 mlx/backend/metal/kernels/cross_entropy.metal diff --git a/mlx/backend/metal/CMakeLists.txt b/mlx/backend/metal/CMakeLists.txt index ea4a995ade..f202d4304e 100644 --- a/mlx/backend/metal/CMakeLists.txt +++ b/mlx/backend/metal/CMakeLists.txt @@ -44,6 +44,7 @@ if(MLX_METAL_JIT) make_jit_source(binary_two) make_jit_source(fft kernels/fft/radix.h kernels/fft/readwrite.h) make_jit_source(logsumexp) + make_jit_source(cross_entropy) make_jit_source(ternary) make_jit_source(softmax) make_jit_source(scan) @@ -134,6 +135,7 @@ target_sources( ${CMAKE_CURRENT_SOURCE_DIR}/hadamard.cpp ${CMAKE_CURRENT_SOURCE_DIR}/indexing.cpp ${CMAKE_CURRENT_SOURCE_DIR}/logsumexp.cpp + ${CMAKE_CURRENT_SOURCE_DIR}/cross_entropy.cpp ${CMAKE_CURRENT_SOURCE_DIR}/matmul.cpp ${CMAKE_CURRENT_SOURCE_DIR}/scaled_dot_product_attention.cpp ${CMAKE_CURRENT_SOURCE_DIR}/metal.cpp diff --git a/mlx/backend/metal/cross_entropy.cpp b/mlx/backend/metal/cross_entropy.cpp new file mode 100644 index 0000000000..cc6f0e1711 --- /dev/null +++ b/mlx/backend/metal/cross_entropy.cpp @@ -0,0 +1,125 @@ +// Copyright © 2026 Apple Inc. +#include + +#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" + +namespace mlx::core::fast { + +namespace { + +// Threadgroup shrinks with the axis so a small vocabulary needs one iteration. +inline std::pair cross_entropy_dims( + MTL::ComputePipelineState* kernel, + int axis_size, + int n_rows) { + constexpr size_t simd_size = 32; + size_t tg_max = + (kernel->maxTotalThreadsPerThreadgroup() / simd_size) * simd_size; + size_t needed = + (axis_size + CROSS_ENTROPY_N_READS - 1) / CROSS_ENTROPY_N_READS; + size_t threadgroup_size = std::clamp( + ((needed + simd_size - 1) / simd_size) * simd_size, simd_size, tg_max); + return { + MTL::Size(n_rows * threadgroup_size, 1, 1), + MTL::Size(threadgroup_size, 1, 1)}; +} + +array ensure_row_contiguous( + const array& x, + metal::CommandEncoder& encoder, + const Stream& s) { + if (x.flags().row_contiguous) { + return x; + } + array x_copy = contiguous_copy_gpu(x, s); + encoder.add_temporary(x_copy); + return x_copy; +} + +} // namespace + +bool CrossEntropy::use_fallback(Stream s) { + return s.device == Device::cpu; +} + +void CrossEntropy::eval_gpu( + const std::vector& inputs, + std::vector& outputs) { + assert(inputs.size() == 2); + auto& in_pre = inputs[0]; + auto& out = outputs[0]; + if (in_pre.size() == 0) { + throw std::invalid_argument("[cross_entropy] Received empty array."); + } + + auto& s = stream(); + auto& d = metal::device(s.device); + auto& compute_encoder = metal::get_command_encoder(s); + + auto in = ensure_row_contiguous(in_pre, compute_encoder, s); + auto targets = ensure_row_contiguous(inputs[1], compute_encoder, s); + out.set_data(allocator::malloc(out.nbytes())); + + int axis_size = in.shape().back(); + int n_rows = in.data_size() / axis_size; + + std::string kernel_name = "cross_entropy_" + type_to_name(in); + auto kernel = get_cross_entropy_kernel(d, kernel_name, in); + auto [grid_dims, group_dims] = cross_entropy_dims(kernel, axis_size, n_rows); + + compute_encoder.set_compute_pipeline_state(kernel); + compute_encoder.set_input_array(in, 0); + compute_encoder.set_input_array(targets, 1); + compute_encoder.set_output_array(out, 2); + compute_encoder.set_bytes(axis_size, 3); + compute_encoder.dispatch_threads(grid_dims, group_dims); +} + +void CrossEntropyVJP::eval_gpu( + const std::vector& inputs, + std::vector& outputs) { + assert(inputs.size() == 4); + auto& in_pre = inputs[0]; + auto& out = outputs[0]; + if (in_pre.size() == 0) { + throw std::invalid_argument("[cross_entropy] Received empty array."); + } + + auto& s = stream(); + auto& d = metal::device(s.device); + auto& compute_encoder = metal::get_command_encoder(s); + + bool donate_in = in_pre.is_donatable() || !in_pre.flags().row_contiguous; + auto in = ensure_row_contiguous(in_pre, compute_encoder, s); + auto targets = ensure_row_contiguous(inputs[1], compute_encoder, s); + auto loss = ensure_row_contiguous(inputs[2], compute_encoder, s); + auto cotan = ensure_row_contiguous(inputs[3], compute_encoder, s); + if (donate_in) { + out.copy_shared_buffer(in); + } else { + out.set_data(allocator::malloc(out.nbytes())); + } + + int axis_size = in.shape().back(); + int n_rows = in.data_size() / axis_size; + + std::string kernel_name = "cross_entropy_vjp_" + type_to_name(in); + auto kernel = get_cross_entropy_kernel(d, kernel_name, in); + auto [grid_dims, group_dims] = cross_entropy_dims(kernel, axis_size, n_rows); + + compute_encoder.set_compute_pipeline_state(kernel); + compute_encoder.set_input_array(in, 0); + compute_encoder.set_input_array(targets, 1); + compute_encoder.set_input_array(loss, 2); + compute_encoder.set_input_array(cotan, 3); + compute_encoder.set_output_array(out, 4); + compute_encoder.set_bytes(axis_size, 5); + compute_encoder.dispatch_threads(grid_dims, group_dims); +} + +} // namespace mlx::core::fast diff --git a/mlx/backend/metal/jit/includes.h b/mlx/backend/metal/jit/includes.h index 4fb1be1110..9fe1923e6f 100644 --- a/mlx/backend/metal/jit/includes.h +++ b/mlx/backend/metal/jit/includes.h @@ -18,6 +18,7 @@ const char* unary(); const char* binary(); const char* binary_two(); const char* copy(); +const char* cross_entropy(); const char* fft(); const char* gather_axis(); const char* gather_front(); diff --git a/mlx/backend/metal/jit_kernels.cpp b/mlx/backend/metal/jit_kernels.cpp index 807a6550b7..aa943e89ad 100644 --- a/mlx/backend/metal/jit_kernels.cpp +++ b/mlx/backend/metal/jit_kernels.cpp @@ -360,6 +360,26 @@ MTL::ComputePipelineState* get_logsumexp_kernel( return d.get_kernel(kernel_name, lib); } +MTL::ComputePipelineState* get_cross_entropy_kernel( + metal::Device& d, + const std::string& kernel_name, + const array& in) { + // Both kernels live in one library, so the name is not derived from + // kernel_name. + std::string lib_name = "cross_entropy_" + type_to_name(in); + auto lib = d.get_library(lib_name, [&] { + auto t_str = get_type_string(in.dtype()); + std::string kernel_source = metal::utils(); + kernel_source += metal::cross_entropy(); + kernel_source += get_template_definition( + "cross_entropy_" + type_to_name(in), "cross_entropy", t_str); + kernel_source += get_template_definition( + "cross_entropy_vjp_" + type_to_name(in), "cross_entropy_vjp", t_str); + return kernel_source; + }); + return d.get_kernel(kernel_name, lib); +} + MTL::ComputePipelineState* get_scan_kernel( metal::Device& d, const std::string& kernel_name, diff --git a/mlx/backend/metal/kernels.h b/mlx/backend/metal/kernels.h index 42888f0a78..d9589e7fc8 100644 --- a/mlx/backend/metal/kernels.h +++ b/mlx/backend/metal/kernels.h @@ -64,6 +64,11 @@ MTL::ComputePipelineState* get_logsumexp_kernel( const std::string& kernel_name, const array& out); +MTL::ComputePipelineState* get_cross_entropy_kernel( + metal::Device& d, + const std::string& kernel_name, + const array& in); + MTL::ComputePipelineState* get_scan_kernel( metal::Device& d, const std::string& kernel_name, diff --git a/mlx/backend/metal/kernels/CMakeLists.txt b/mlx/backend/metal/kernels/CMakeLists.txt index ecabbda550..8a60470b5e 100644 --- a/mlx/backend/metal/kernels/CMakeLists.txt +++ b/mlx/backend/metal/kernels/CMakeLists.txt @@ -142,6 +142,7 @@ if(NOT MLX_METAL_JIT) build_kernel(scan scan.h) build_kernel(softmax softmax.h) build_kernel(logsumexp logsumexp.h) + build_kernel(cross_entropy cross_entropy.h) build_kernel(searchsorted searchsorted.h sort.h) build_kernel(sort sort.h) build_kernel(ternary ternary.h ternary_ops.h) diff --git a/mlx/backend/metal/kernels/cross_entropy.h b/mlx/backend/metal/kernels/cross_entropy.h new file mode 100644 index 0000000000..bef0bce5f8 --- /dev/null +++ b/mlx/backend/metal/kernels/cross_entropy.h @@ -0,0 +1,133 @@ +// Copyright © 2026 Apple Inc. + +template < + typename T, + typename AccT = float, + int N_READS = CROSS_ENTROPY_N_READS> +[[kernel]] void cross_entropy( + const device T* in, + const device int32_t* targets, + device float* out, + constant int& axis_size, + uint gid [[threadgroup_position_in_grid]], + uint lid [[thread_position_in_threadgroup]], + uint lsize [[threads_per_threadgroup]], + uint simd_lane_id [[thread_index_in_simdgroup]], + uint simd_group_id [[simdgroup_index_in_threadgroup]]) { + in += gid * size_t(axis_size); + + constexpr int SIMD_SIZE = 32; + + threadgroup AccT local_max[SIMD_SIZE]; + threadgroup AccT local_normalizer[SIMD_SIZE]; + + int y_n = targets[gid]; + y_n = (y_n < 0) ? y_n + axis_size : y_n; + AccT x_t = AccT(in[y_n]); + + AccT prevmax; + AccT maxval = Limits::finite_min; + AccT normalizer = 0; + for (int r = 0; r < static_cast(ceildiv(axis_size, N_READS * lsize)); + r++) { + int offset = r * lsize * N_READS + lid * N_READS; + AccT vals[N_READS]; + if (offset + N_READS <= axis_size) { + for (int i = 0; i < N_READS; i++) { + vals[i] = AccT(in[offset + i]); + } + } else { + for (int i = 0; i < N_READS; i++) { + vals[i] = + (offset + i < axis_size) ? AccT(in[offset + i]) : Limits::min; + } + } + prevmax = maxval; + for (int i = 0; i < N_READS; i++) { + maxval = (maxval < vals[i]) ? vals[i] : maxval; + } + normalizer *= fast::exp(prevmax - maxval); + for (int i = 0; i < N_READS; i++) { + normalizer += fast::exp(vals[i] - maxval); + } + } + + prevmax = maxval; + maxval = simd_max(maxval); + normalizer *= fast::exp(prevmax - maxval); + normalizer = simd_sum(normalizer); + + uint n_simdgroups = ceildiv(lsize, uint(SIMD_SIZE)); + prevmax = maxval; + if (simd_lane_id == 0) { + local_max[simd_group_id] = maxval; + } + threadgroup_barrier(mem_flags::mem_threadgroup); + maxval = (simd_lane_id < n_simdgroups) ? local_max[simd_lane_id] + : Limits::finite_min; + maxval = simd_max(maxval); + normalizer *= fast::exp(prevmax - maxval); + if (simd_lane_id == 0) { + local_normalizer[simd_group_id] = normalizer; + } + threadgroup_barrier(mem_flags::mem_threadgroup); + normalizer = + (simd_lane_id < n_simdgroups) ? local_normalizer[simd_lane_id] : AccT(0); + normalizer = simd_sum(normalizer); + + if (lid == 0) { + // Subtract the two logits first; they are the same magnitude, so less is + // lost than in logsumexp - x_t. + AccT gap = maxval - x_t; + out[gid] = isinf(maxval) ? float(gap) : float(log(normalizer) + gap); + } +} + +template < + typename T, + typename AccT = float, + int N_READS = CROSS_ENTROPY_N_READS> +[[kernel]] void cross_entropy_vjp( + const device T* in, + const device int32_t* targets, + const device float* loss, + const device float* cotan, + device T* out, + constant int& axis_size, + uint gid [[threadgroup_position_in_grid]], + uint lid [[thread_position_in_threadgroup]], + uint lsize [[threads_per_threadgroup]]) { + size_t row_offset = gid * size_t(axis_size); + in += row_offset; + out += row_offset; + + int y_n = targets[gid]; + y_n = (y_n < 0) ? y_n + axis_size : y_n; + AccT x_t = AccT(in[y_n]); + AccT g = AccT(cotan[gid]); + AccT l = AccT(loss[gid]); + + // out aliases in when the input is donated, so every thread must read the + // target column before any thread writes. + threadgroup_barrier(mem_flags::mem_device); + + for (int r = 0; r < static_cast(ceildiv(axis_size, N_READS * lsize)); + r++) { + int offset = r * lsize * N_READS + lid * N_READS; + if (offset + N_READS <= axis_size) { + for (int i = 0; i < N_READS; i++) { + int col = offset + i; + AccT p = fast::exp((AccT(in[col]) - x_t) - l); + out[col] = T(g * (p - ((col == y_n) ? AccT(1) : AccT(0)))); + } + } else { + for (int i = 0; i < N_READS; i++) { + int col = offset + i; + if (col < axis_size) { + AccT p = fast::exp((AccT(in[col]) - x_t) - l); + out[col] = T(g * (p - ((col == y_n) ? AccT(1) : AccT(0)))); + } + } + } + } +} diff --git a/mlx/backend/metal/kernels/cross_entropy.metal b/mlx/backend/metal/kernels/cross_entropy.metal new file mode 100644 index 0000000000..f35fbe0609 --- /dev/null +++ b/mlx/backend/metal/kernels/cross_entropy.metal @@ -0,0 +1,18 @@ +// Copyright © 2026 Apple Inc. + +#include +#include + +using namespace metal; + +// clang-format off +#include "mlx/backend/metal/kernels/utils.h" +#include "mlx/backend/metal/kernels/cross_entropy.h" + +#define instantiate_cross_entropy(name, itype) \ + instantiate_kernel("cross_entropy_" #name, cross_entropy, itype) \ + instantiate_kernel("cross_entropy_vjp_" #name, cross_entropy_vjp, itype) + +instantiate_cross_entropy(float32, float) +instantiate_cross_entropy(float16, half) +instantiate_cross_entropy(bfloat16, bfloat16_t) // clang-format on diff --git a/mlx/backend/metal/kernels/defines.h b/mlx/backend/metal/kernels/defines.h index c369adb7e8..0e65906164 100644 --- a/mlx/backend/metal/kernels/defines.h +++ b/mlx/backend/metal/kernels/defines.h @@ -14,6 +14,7 @@ static MTL_CONST constexpr int REDUCE_N_WRITES = 4; static MTL_CONST constexpr int SOFTMAX_N_READS = 4; static MTL_CONST constexpr int RMS_N_READS = 4; static MTL_CONST constexpr int RMS_LOOPED_LIMIT = 4096; +static MTL_CONST constexpr int CROSS_ENTROPY_N_READS = 4; // Instantiate a templated kernel. // Extra args are used as template parameters: diff --git a/mlx/backend/metal/nojit_kernels.cpp b/mlx/backend/metal/nojit_kernels.cpp index f1141e9792..e07a06f441 100644 --- a/mlx/backend/metal/nojit_kernels.cpp +++ b/mlx/backend/metal/nojit_kernels.cpp @@ -79,6 +79,13 @@ MTL::ComputePipelineState* get_logsumexp_kernel( return d.get_kernel(kernel_name); } +MTL::ComputePipelineState* get_cross_entropy_kernel( + metal::Device& d, + const std::string& kernel_name, + const array&) { + return d.get_kernel(kernel_name); +} + MTL::ComputePipelineState* get_scan_kernel( metal::Device& d, const std::string& kernel_name, diff --git a/mlx/backend/metal/primitives.cpp b/mlx/backend/metal/primitives.cpp index d1d0e781cc..45929e27dd 100644 --- a/mlx/backend/metal/primitives.cpp +++ b/mlx/backend/metal/primitives.cpp @@ -13,7 +13,6 @@ #include "mlx/backend/metal/kernels.h" #include "mlx/backend/metal/utils.h" #include "mlx/dtype_utils.h" -#include "mlx/fast_primitives.h" #include "mlx/primitives.h" #include "mlx/scheduler.h" #include "mlx/utils.h" @@ -215,26 +214,4 @@ void LUF::eval_gpu( throw std::runtime_error("[LUF::eval_gpu] Metal LU factorization NYI."); } -namespace fast { - -// There is no fused Metal cross entropy kernel yet -bool CrossEntropy::use_fallback(Stream s) { - return true; -} - -void CrossEntropy::eval_gpu( - const std::vector& inputs, - std::vector& outputs) { - throw std::runtime_error("[CrossEntropy::eval_gpu] Metal cross entropy NYI."); -} - -void CrossEntropyVJP::eval_gpu( - const std::vector& inputs, - std::vector& outputs) { - throw std::runtime_error( - "[CrossEntropyVJP::eval_gpu] Metal cross entropy NYI."); -} - -} // namespace fast - } // namespace mlx::core diff --git a/python/src/fast.cpp b/python/src/fast.cpp index 67c3442cff..58bfe3ec4c 100644 --- a/python/src/fast.cpp +++ b/python/src/fast.cpp @@ -189,8 +189,8 @@ void init_fast(nb::module_& parent_module) { Computes ``logsumexp(logits, axis=-1) - logits[..., target]`` in a fused kernel with accumulation in float32. - Note: Currently is implemented only on CUDA, fallback to unfused version with - manual casting on Metal and CPU. + Note: The fused kernel is available on Metal and CUDA. The CPU falls + back to the unfused version, which reduces in the dtype of the logits. Args: logits (array): The unnormalized logits. The loss is computed over diff --git a/python/tests/test_fast.py b/python/tests/test_fast.py index 200781d372..56f6211ef9 100644 --- a/python/tests/test_fast.py +++ b/python/tests/test_fast.py @@ -535,6 +535,9 @@ def cross_entropy_ref(logits, targets): ) tolerances = {mx.float32: 1e-5, mx.float16: 3e-2, mx.bfloat16: 3e-1} + # The fused GPU kernels accumulate in float32 for every input dtype. + if mx.default_device() == mx.gpu: + tolerances = {mx.float32: 1e-5, mx.float16: 1e-3, mx.bfloat16: 1e-3} for V in [7, 32, 128, 255, 256, 1000, 4096, 8192]: for dtype in [mx.float32, mx.float16, mx.bfloat16]: From d69dacd341d484786464eeee53e9007a584cedc1 Mon Sep 17 00:00:00 2001 From: Zhongrui Sun Date: Sat, 12 Sep 2026 20:22:54 -0700 Subject: [PATCH 2/2] Use the fused cross entropy kernel in nn.losses on Metal nn.losses.cross_entropy only took the mx.fast.cross_entropy path when CUDA was available. The Metal kernel exists now, so gate on the default device being the GPU alone. The half precision example in the docstring now shows what the fast path returns: the loss is cast back to the logits dtype, so it is bfloat16, not float32. --- python/mlx/nn/losses.py | 11 +++++------ 1 file changed, 5 insertions(+), 6 deletions(-) diff --git a/python/mlx/nn/losses.py b/python/mlx/nn/losses.py index b98d2765d6..b72b91816f 100644 --- a/python/mlx/nn/losses.py +++ b/python/mlx/nn/losses.py @@ -64,15 +64,15 @@ def cross_entropy( >>> nn.losses.cross_entropy(logits, targets) array([0.348587, 0.348587], dtype=float32) >>> - >>> # Half precision logits with class indices as targets. On CUDA a + >>> # Half precision logits with class indices as targets. On the GPU a >>> # fused kernel accumulates the reduction in float32: >>> logits = mx.array([[2.0, -1.0], [-1.0, 2.0]], mx.bfloat16) >>> targets = mx.array([0, 1]) >>> nn.losses.cross_entropy(logits, targets) - array([0.0485873, 0.0485873], dtype=float32) + array([0.048584, 0.048584], dtype=bfloat16) >>> - >>> # Metal and the CPU reduce in the dtype of the logits, so upcast - >>> # them to get the same accuracy: + >>> # The CPU reduces in the dtype of the logits, so upcast them to get + >>> # the same accuracy: >>> nn.losses.cross_entropy(logits.astype(mx.float32), targets) array([0.0485873, 0.0485873], dtype=float32) """ @@ -96,8 +96,7 @@ def _drop_dim(shape, axis): ) use_fast = ( - mx.cuda.is_available() - and mx.default_device() == mx.gpu + mx.default_device() == mx.gpu and not targets_as_probs and label_smoothing == 0 and axis in (-1, logits.ndim - 1)