From 4e9f20b428beebe0f4d6a5b9de56c1edc52f0dd7 Mon Sep 17 00:00:00 2001 From: Cursor Agent Date: Tue, 15 Sep 2026 10:30:51 +0000 Subject: [PATCH 1/3] cpu: Use size_t for contiguous binary functor counts The VectorScalar/ScalarVector/VectorVector operators took int size, so counts past 2^31 wrapped and the op wrote nothing. Co-authored-by: Jasim Kareem --- mlx/backend/cpu/binary.h | 6 +++--- python/tests/test_ops.py | 12 ++++++++++++ 2 files changed, 15 insertions(+), 3 deletions(-) diff --git a/mlx/backend/cpu/binary.h b/mlx/backend/cpu/binary.h index 5333aff9a1..fb7db76f93 100644 --- a/mlx/backend/cpu/binary.h +++ b/mlx/backend/cpu/binary.h @@ -16,7 +16,7 @@ namespace mlx::core { template struct VectorScalar { template - void operator()(const T* a, const T* b, U* dst, int size) { + void operator()(const T* a, const T* b, U* dst, size_t size) { T scalar = *b; constexpr int N = simd::max_size; while (size >= N) { @@ -36,7 +36,7 @@ struct VectorScalar { template struct ScalarVector { template - void operator()(const T* a, const T* b, U* dst, int size) { + void operator()(const T* a, const T* b, U* dst, size_t size) { T scalar = *a; constexpr int N = simd::max_size; while (size >= N) { @@ -56,7 +56,7 @@ struct ScalarVector { template struct VectorVector { template - void operator()(const T* a, const T* b, U* dst, int size) { + void operator()(const T* a, const T* b, U* dst, size_t size) { constexpr int N = simd::max_size; while (size >= N) { simd::store(dst, Op{}(simd::load(a), simd::load(b))); diff --git a/python/tests/test_ops.py b/python/tests/test_ops.py index 151eddea8e..f388c59fbe 100644 --- a/python/tests/test_ops.py +++ b/python/tests/test_ops.py @@ -3211,6 +3211,18 @@ def test_large_binary(self): b = mx.ones([2147484], mx.int8) self.assertEqual((a + b)[0, 0].item(), 2) + @unittest.skipIf( + os.getenv("LOW_MEMORY", None) is not None, + "This test requires a lot of memory", + ) + def test_large_contiguous_binary(self): + # VectorScalar used int size, so n >= 2**31 wrote nothing. + n = 2**31 + a = mx.full((n,), 3, dtype=mx.uint8, stream=mx.cpu) + out = mx.add(a, mx.array(1, mx.uint8), stream=mx.cpu) + self.assertEqual(out[0].item(), 4) + self.assertEqual(out[-1].item(), 4) + def test_eye(self): self.assertCmpNumpy([3], mx.eye, np.eye) # Test for zero rows and columns From acaf23317dbaefa7214e829908c46d19fbd87a4f Mon Sep 17 00:00:00 2001 From: Cursor Agent Date: Tue, 15 Sep 2026 10:31:58 +0000 Subject: [PATCH 2/3] python: Clear cache before large contiguous binary test Co-authored-by: Jasim Kareem --- python/tests/test_ops.py | 1 + 1 file changed, 1 insertion(+) diff --git a/python/tests/test_ops.py b/python/tests/test_ops.py index f388c59fbe..937e4c468d 100644 --- a/python/tests/test_ops.py +++ b/python/tests/test_ops.py @@ -3217,6 +3217,7 @@ def test_large_binary(self): ) def test_large_contiguous_binary(self): # VectorScalar used int size, so n >= 2**31 wrote nothing. + mx.clear_cache() n = 2**31 a = mx.full((n,), 3, dtype=mx.uint8, stream=mx.cpu) out = mx.add(a, mx.array(1, mx.uint8), stream=mx.cpu) From b26f53df1e1cbeeb85b8c333c6608a0f45ca11a0 Mon Sep 17 00:00:00 2001 From: Cheng Date: Mon, 28 Sep 2026 17:29:41 +0800 Subject: [PATCH 3/3] nit --- python/tests/test_ops.py | 13 ------------- 1 file changed, 13 deletions(-) diff --git a/python/tests/test_ops.py b/python/tests/test_ops.py index 937e4c468d..151eddea8e 100644 --- a/python/tests/test_ops.py +++ b/python/tests/test_ops.py @@ -3211,19 +3211,6 @@ def test_large_binary(self): b = mx.ones([2147484], mx.int8) self.assertEqual((a + b)[0, 0].item(), 2) - @unittest.skipIf( - os.getenv("LOW_MEMORY", None) is not None, - "This test requires a lot of memory", - ) - def test_large_contiguous_binary(self): - # VectorScalar used int size, so n >= 2**31 wrote nothing. - mx.clear_cache() - n = 2**31 - a = mx.full((n,), 3, dtype=mx.uint8, stream=mx.cpu) - out = mx.add(a, mx.array(1, mx.uint8), stream=mx.cpu) - self.assertEqual(out[0].item(), 4) - self.assertEqual(out[-1].item(), 4) - def test_eye(self): self.assertCmpNumpy([3], mx.eye, np.eye) # Test for zero rows and columns