diff --git a/mlx/backend/metal/CMakeLists.txt b/mlx/backend/metal/CMakeLists.txt index c86cfdeb07..bc5c2fcf70 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) @@ -135,6 +136,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}/gated_delta_update.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 f5a5fe45ef..b4f72e660c 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 1c052e6314..71122d2ed9 100644 --- a/mlx/backend/metal/jit_kernels.cpp +++ b/mlx/backend/metal/jit_kernels.cpp @@ -363,6 +363,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 28ac3ba940..ca321b96c1 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 2414337509..d5513dd030 100644 --- a/mlx/backend/metal/kernels/CMakeLists.txt +++ b/mlx/backend/metal/kernels/CMakeLists.txt @@ -144,6 +144,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 4f48a25b77..879a74f378 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/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) diff --git a/python/src/fast.cpp b/python/src/fast.cpp index c388980f3e..f9c472ab99 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]: