Skip to content
Open
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
29 changes: 17 additions & 12 deletions mlx/backend/cpu/reduce.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -108,7 +108,12 @@ void strided_reduce(
};

template <typename T, typename U, typename Op>
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<T>, simd::max_size<U>);
simd::Simd<U, N> accumulator_v(init);
while (size >= N) {
Expand All @@ -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<void(int)> callback,
std::function<void(int64_t)> callback,
const Shape& shape,
const Strides& strides) {
std::function<void(int, int)> loop_inner;
loop_inner = [&](int dim, int offset) {
std::function<void(int, int64_t)> loop_inner;
loop_inner = [&](int dim, int64_t offset) {
if (dim < shape.size() - 1) {
auto size = shape[dim];
auto stride = strides[dim];
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -229,18 +234,18 @@ 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{});
out_ptr += reduction_stride;
}
} 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,
Expand All @@ -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,
Expand Down
12 changes: 6 additions & 6 deletions mlx/backend/metal/reduce.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -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";
}
Expand Down Expand Up @@ -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";
}
Expand Down Expand Up @@ -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";
}
Expand Down Expand Up @@ -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";
}
Expand Down Expand Up @@ -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";
}
Expand Down Expand Up @@ -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";
}
Expand Down
28 changes: 28 additions & 0 deletions python/tests/test_reduce.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,7 @@
# Copyright © 2023 Apple Inc.

import os
import unittest
from itertools import combinations, permutations

import mlx.core as mx
Expand Down Expand Up @@ -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)