From 5905395881e8deb6db77af6d689e4b330bc3f75d Mon Sep 17 00:00:00 2001 From: nicoloangileri Date: Wed, 23 Sep 2026 08:23:00 +0200 Subject: [PATCH] Use 64-bit offsets and sizes in reductions Co-Authored-By: Claude Opus 5.5 --- mlx/backend/cpu/reduce.cpp | 29 +++++++++++++++++------------ mlx/backend/metal/reduce.cpp | 12 ++++++------ python/tests/test_reduce.py | 28 ++++++++++++++++++++++++++++ 3 files changed, 51 insertions(+), 18 deletions(-) diff --git a/mlx/backend/cpu/reduce.cpp b/mlx/backend/cpu/reduce.cpp index 04965690be..b271d89a94 100644 --- a/mlx/backend/cpu/reduce.cpp +++ b/mlx/backend/cpu/reduce.cpp @@ -108,7 +108,12 @@ void strided_reduce( }; template -void contiguous_reduce(const T* x, U* accumulator, int size, Op op, U init) { +void contiguous_reduce( + const T* x, + U* accumulator, + int64_t size, + Op op, + U init) { constexpr int N = std::min(simd::max_size, simd::max_size); simd::Simd accumulator_v(init); while (size >= N) { @@ -125,11 +130,11 @@ void contiguous_reduce(const T* x, U* accumulator, int size, Op op, U init) { // Helper for the ndimensional strided loop void nd_loop( - std::function callback, + std::function callback, const Shape& shape, const Strides& strides) { - std::function loop_inner; - loop_inner = [&](int dim, int offset) { + std::function loop_inner; + loop_inner = [&](int dim, int64_t offset) { if (dim < shape.size() - 1) { auto size = shape[dim]; auto stride = strides[dim]; @@ -181,16 +186,16 @@ void reduction_op( auto [shape, strides] = shapes_without_reduction_axes(x, axes); if (plan.shape.size() == 0) { for (int i = 0; i < out.size(); i++, out_ptr++) { - int offset = elem_to_loc(i, shape, strides); + int64_t offset = elem_to_loc(i, shape, strides); *out_ptr = init; contiguous_reduce(in_ptr + offset, out_ptr, reduction_size, Op{}, init); } } else { for (int i = 0; i < out.size(); i++, out_ptr++) { - int offset = elem_to_loc(i, shape, strides); + int64_t offset = elem_to_loc(i, shape, strides); *out_ptr = init; nd_loop( - [&](int extra_offset) { + [&](int64_t extra_offset) { contiguous_reduce( in_ptr + offset + extra_offset, out_ptr, @@ -229,7 +234,7 @@ void reduction_op( if (plan.shape.size() == 0) { for (int i = 0; i < out.size(); i += reduction_stride) { - int offset = elem_to_loc(i, shape, strides); + int64_t offset = elem_to_loc(i, shape, strides); std::fill_n(out_ptr, reduction_stride, init); strided_reduce( in_ptr + offset, out_ptr, reduction_size, reduction_stride, Op{}); @@ -237,10 +242,10 @@ void reduction_op( } } else { for (int i = 0; i < out.size(); i += reduction_stride) { - int offset = elem_to_loc(i, shape, strides); + int64_t offset = elem_to_loc(i, shape, strides); std::fill_n(out_ptr, reduction_stride, init); nd_loop( - [&](int extra_offset) { + [&](int64_t extra_offset) { strided_reduce( in_ptr + offset + extra_offset, out_ptr, @@ -260,10 +265,10 @@ void reduction_op( auto [shape, strides] = shapes_without_reduction_axes(x, axes); for (int i = 0; i < out.size(); i++, out_ptr++) { - int offset = elem_to_loc(i, shape, strides); + int64_t offset = elem_to_loc(i, shape, strides); U val = init; nd_loop( - [&](int extra_offset) { + [&](int64_t extra_offset) { val = Op{}(val, *(in_ptr + offset + extra_offset)); }, plan.shape, diff --git a/mlx/backend/metal/reduce.cpp b/mlx/backend/metal/reduce.cpp index ac11ac6359..6d635c8a7c 100644 --- a/mlx/backend/metal/reduce.cpp +++ b/mlx/backend/metal/reduce.cpp @@ -421,7 +421,7 @@ void row_reduce_small( auto [in_type, out_type] = remap_reduce_types(in, op_name); const std::string func_name = "row_reduce_small"; std::string kname = func_name; - bool large = in.size() > INT32_MAX; + bool large = in.size() > INT32_MAX || in.data_size() > INT32_MAX; if (large) { kname += "_large"; } @@ -518,7 +518,7 @@ void row_reduce_looped( int n = get_kernel_reduce_ndim(args.reduce_ndim); const std::string func_name = "row_reduce_looped"; std::string kname = func_name; - bool large = in.size() > INT32_MAX; + bool large = in.size() > INT32_MAX || in.data_size() > INT32_MAX; if (large) { kname += "_large"; } @@ -602,7 +602,7 @@ void strided_reduce_small( int n = get_kernel_reduce_ndim(args.reduce_ndim); const std::string func_name = "col_reduce_small"; std::string kname = func_name; - bool large = in.size() > INT32_MAX; + bool large = in.size() > INT32_MAX || in.data_size() > INT32_MAX; if (large) { kname += "_large"; } @@ -693,7 +693,7 @@ void strided_reduce_longcolumn( int n = get_kernel_reduce_ndim(args.reduce_ndim); std::string func_name = "col_reduce_longcolumn"; std::string kname = func_name; - bool large = in.size() > INT32_MAX; + bool large = in.size() > INT32_MAX || in.data_size() > INT32_MAX; if (large) { kname += "_large"; } @@ -788,7 +788,7 @@ void strided_reduce_looped( int n = get_kernel_reduce_ndim(args.reduce_ndim); std::string func_name = "col_reduce_looped"; std::string kname = func_name; - bool large = in.size() > INT32_MAX; + bool large = in.size() > INT32_MAX || in.data_size() > INT32_MAX; if (large) { kname += "_large"; } @@ -865,7 +865,7 @@ void strided_reduce_2pass( int n = get_kernel_reduce_ndim(args.reduce_ndim); std::string func_name = "col_reduce_2pass"; std::string kname = func_name; - bool large = in.size() > INT32_MAX; + bool large = in.size() > INT32_MAX || in.data_size() > INT32_MAX; if (large) { kname += "_large"; } diff --git a/python/tests/test_reduce.py b/python/tests/test_reduce.py index 164e2dd803..d960015413 100644 --- a/python/tests/test_reduce.py +++ b/python/tests/test_reduce.py @@ -1,5 +1,7 @@ # Copyright © 2023 Apple Inc. +import os +import unittest from itertools import combinations, permutations import mlx.core as mx @@ -282,6 +284,32 @@ def test_and_or_negative_zero(self): getattr(np, op)(x_np, axis=1).tolist(), ) + @unittest.skipIf( + os.getenv("LOW_MEMORY", None) is not None, + "This test requires a lot of memory", + ) + def test_large_offsets(self): + # Row r holds r % 251 and row 2**15 starts at offset 2**31, so any + # reduction that narrows offsets or sizes to 32 bits reads the wrong + # rows. The views below have fewer than 2**31 elements themselves. + mx.clear_cache() + rows, cols = 2**15 + 1, 2**16 + row_max = (mx.arange(rows) % 251).astype(mx.uint8) + x = mx.contiguous(mx.broadcast_to(row_max[:, None], (rows, cols))) + + self.assertTrue(mx.array_equal(x[:, 7:].max(axis=-1), row_max)) + self.assertTrue(mx.array_equal(x[:, 7:].min(axis=-1), row_max)) + y = x.reshape(rows, 256, 256)[:, 1:, 1:] + self.assertTrue(mx.array_equal(y.max(axis=(1, 2)), row_max)) + y = x.reshape(rows, 16, 16, 256)[:, 1:, :, 1:] + expected = mx.broadcast_to(row_max[:, None], (rows, 16)) + self.assertTrue(mx.array_equal(y.max(axis=(1, 3)), expected)) + + # Reducing all 2**31 + 2**16 elements at once + self.assertEqual(x.max().item(), 250) + self.assertTrue(x.any().item()) + self.assertFalse(x.all().item()) + if __name__ == "__main__": mlx_tests.MLXTestRunner(failfast=True)