diff --git a/mlx/backend/metal/jit_kernels.cpp b/mlx/backend/metal/jit_kernels.cpp index d639719d81..f37bde56fa 100644 --- a/mlx/backend/metal/jit_kernels.cpp +++ b/mlx/backend/metal/jit_kernels.cpp @@ -319,6 +319,8 @@ MTL::ComputePipelineState* get_logsumexp_kernel( kernel_source += metal::logsumexp(); kernel_source += get_template_definition("block_" + lib_name, "logsumexp", t_str); + kernel_source += get_template_definition( + "simdrow_" + lib_name, "logsumexp_simd_row", t_str); kernel_source += get_template_definition( "looped_" + lib_name, "logsumexp_looped", t_str); return kernel_source; diff --git a/mlx/backend/metal/kernels/logsumexp.h b/mlx/backend/metal/kernels/logsumexp.h index c746050b3f..eef9248e21 100644 --- a/mlx/backend/metal/kernels/logsumexp.h +++ b/mlx/backend/metal/kernels/logsumexp.h @@ -72,6 +72,58 @@ template } } +template +[[kernel]] void logsumexp_simd_row( + const device T* in, + device T* out, + constant int& axis_size, + uint tid [[thread_position_in_grid]], + uint simd_lane_id [[thread_index_in_simdgroup]]) { + // One simdgroup per row: reductions stay within the simdgroup, so no + // threadgroup memory or barriers are needed (the block kernel uses 5). + // The grid has exactly n_rows * SIMD_SIZE threads, so every simdgroup is + // full and maps to one row. + constexpr int SIMD_SIZE = 32; + + uint row = tid / SIMD_SIZE; + in += row * size_t(axis_size); + + AccT prevmax; + AccT maxval = Limits::finite_min; + AccT normalizer = 0; + for (int r = 0; r < static_cast(ceildiv(axis_size, N_READS * SIMD_SIZE)); + r++) { + int offset = r * SIMD_SIZE * N_READS + simd_lane_id * 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); + + if (simd_lane_id == 0) { + out[row] = isinf(maxval) ? T(maxval) : T(log(normalizer) + maxval); + } +} + template [[kernel]] void logsumexp_looped( const device T* in, diff --git a/mlx/backend/metal/kernels/logsumexp.metal b/mlx/backend/metal/kernels/logsumexp.metal index eb76436cf0..2f7ac46f00 100644 --- a/mlx/backend/metal/kernels/logsumexp.metal +++ b/mlx/backend/metal/kernels/logsumexp.metal @@ -9,9 +9,10 @@ using namespace metal; #include "mlx/backend/metal/kernels/utils.h" #include "mlx/backend/metal/kernels/logsumexp.h" -#define instantiate_logsumexp(name, itype) \ - instantiate_kernel("block_logsumexp_" #name, logsumexp, itype) \ - instantiate_kernel("looped_logsumexp_" #name, logsumexp_looped, itype) \ +#define instantiate_logsumexp(name, itype) \ + instantiate_kernel("block_logsumexp_" #name, logsumexp, itype) \ + instantiate_kernel("simdrow_logsumexp_" #name, logsumexp_simd_row, itype) \ + instantiate_kernel("looped_logsumexp_" #name, logsumexp_looped, itype) \ instantiate_logsumexp(float32, float) instantiate_logsumexp(float16, half) diff --git a/mlx/backend/metal/logsumexp.cpp b/mlx/backend/metal/logsumexp.cpp index 8f7cbe3aff..6de98c6ab4 100644 --- a/mlx/backend/metal/logsumexp.cpp +++ b/mlx/backend/metal/logsumexp.cpp @@ -10,6 +10,9 @@ namespace mlx::core { constexpr int LOGSUMEXP_LOOPED_LIMIT = 4096; +// Rows up to this size go to the simdgroup-per-row kernel; on M4 it wins +// for short-to-medium rows while the block kernel stays ahead beyond it. +constexpr int LOGSUMEXP_SIMD_ROW_LIMIT = 2048; void LogSumExp::eval_gpu(const std::vector& inputs, array& out) { assert(inputs.size() == 1); @@ -61,15 +64,27 @@ void LogSumExp::eval_gpu(const std::vector& inputs, array& out) { const int simd_size = 32; const int n_reads = 4; const int looped_limit = LOGSUMEXP_LOOPED_LIMIT; + const int simd_row_limit = LOGSUMEXP_SIMD_ROW_LIMIT; - std::string kernel_name = (axis_size > looped_limit) ? "looped_" : "block_"; + std::string kernel_name = (axis_size > looped_limit) ? "looped_" + : (axis_size > simd_row_limit) ? "block_" + : "simdrow_"; kernel_name += "logsumexp_"; kernel_name += type_to_name(out); auto kernel = get_logsumexp_kernel(d, kernel_name, out); { MTL::Size grid_dims, group_dims; - if (axis_size <= looped_limit) { + if (axis_size <= simd_row_limit) { + // One simdgroup per row, eight rows per threadgroup. The grid has + // exactly 32 threads per row so simdgroups never straddle rows. + constexpr int simds_per_group = 8; + size_t threadgroup_size = simd_size * simds_per_group; + assert(threadgroup_size <= kernel->maxTotalThreadsPerThreadgroup()); + size_t n_threads = n_rows * size_t(simd_size); + grid_dims = MTL::Size(n_threads, 1, 1); + group_dims = MTL::Size(threadgroup_size, 1, 1); + } else if (axis_size <= looped_limit) { size_t threadgroup_needed = (axis_size + n_reads - 1) / n_reads; size_t simds_needed = (threadgroup_needed + simd_size - 1) / simd_size; size_t threadgroup_size = simd_size * simds_needed;