Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions mlx/backend/metal/CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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
Expand Down
125 changes: 125 additions & 0 deletions mlx/backend/metal/cross_entropy.cpp
Original file line number Diff line number Diff line change
@@ -0,0 +1,125 @@
// Copyright © 2026 Apple Inc.
#include <algorithm>

#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<MTL::Size, MTL::Size> 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<array>& inputs,
std::vector<array>& 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<array>& inputs,
std::vector<array>& 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
1 change: 1 addition & 0 deletions mlx/backend/metal/jit/includes.h
Original file line number Diff line number Diff line change
Expand Up @@ -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();
Expand Down
20 changes: 20 additions & 0 deletions mlx/backend/metal/jit_kernels.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
5 changes: 5 additions & 0 deletions mlx/backend/metal/kernels.h
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
1 change: 1 addition & 0 deletions mlx/backend/metal/kernels/CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
133 changes: 133 additions & 0 deletions mlx/backend/metal/kernels/cross_entropy.h
Original file line number Diff line number Diff line change
@@ -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<AccT>::finite_min;
AccT normalizer = 0;
for (int r = 0; r < static_cast<int>(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<AccT>::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<AccT>::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<int>(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))));
}
}
}
}
}
18 changes: 18 additions & 0 deletions mlx/backend/metal/kernels/cross_entropy.metal
Original file line number Diff line number Diff line change
@@ -0,0 +1,18 @@
// Copyright © 2026 Apple Inc.

#include <metal_common>
#include <metal_simdgroup>

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
1 change: 1 addition & 0 deletions mlx/backend/metal/kernels/defines.h
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
7 changes: 7 additions & 0 deletions mlx/backend/metal/nojit_kernels.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
Loading
Loading