From fca3052f70eaa94a3226127d53b57b7ad5815b09 Mon Sep 17 00:00:00 2001 From: nicoloangileri Date: Wed, 23 Sep 2026 08:23:00 +0200 Subject: [PATCH 1/2] 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_ops.py | 4 ---- python/tests/test_reduce.py | 23 +++++++++++++++++++++++ 4 files changed, 46 insertions(+), 22 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_ops.py b/python/tests/test_ops.py index dd22c329f6..d27a2617b1 100644 --- a/python/tests/test_ops.py +++ b/python/tests/test_ops.py @@ -3469,10 +3469,6 @@ def all_outputs(x): self.assertEqual(inverse.shape, v_mx.shape) self.assertTrue(np.array_equal(np.array(values[inverse]), v_np)) - @unittest.skipIf( - os.getenv("LOW_MEMORY", None) is not None, - "This test requires a lot of memory", - ) def test_large_binary(self): a = mx.ones([1000, 2147484], mx.int8) b = mx.ones([2147484], mx.int8) diff --git a/python/tests/test_reduce.py b/python/tests/test_reduce.py index c155ad0a9f..91aa3fe323 100644 --- a/python/tests/test_reduce.py +++ b/python/tests/test_reduce.py @@ -1,6 +1,8 @@ # Copyright © 2023 Apple Inc. import math +import os +import unittest from itertools import combinations, permutations import mlx.core as mx @@ -339,6 +341,27 @@ def test_and_or_negative_zero(self): getattr(np, op)(x_np, axis=1).tolist(), ) + 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. + 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) From de91d8cfe6bc8d4ee4e61ca83c60b91bb897a621 Mon Sep 17 00:00:00 2001 From: Cheng Date: Thu, 1 Oct 2026 15:56:59 +0800 Subject: [PATCH 2/2] nit --- python/tests/test_reduce.py | 2 -- 1 file changed, 2 deletions(-) diff --git a/python/tests/test_reduce.py b/python/tests/test_reduce.py index 91aa3fe323..6c0138d7e3 100644 --- a/python/tests/test_reduce.py +++ b/python/tests/test_reduce.py @@ -1,8 +1,6 @@ # Copyright © 2023 Apple Inc. import math -import os -import unittest from itertools import combinations, permutations import mlx.core as mx