Add Masked Scatter (#2663)

Co-authored-by: Awni Hannun <[email protected]>
Co-authored-by: Angelos Katharopoulos <[email protected]>
Co-authored-by: Angelos Katharopoulos <[email protected]>
This commit is contained in:
CCYeh
2025-11-19 14:53:32 -08:00
committed by GitHub
co-authored by Awni Hannun Angelos Katharopoulos Angelos Katharopoulos
parent 7f4b7e553c
commit b3825ac149
26 changed files with 1099 additions and 51 deletions
+63 -1
View File
@@ -1260,7 +1260,6 @@ class TestOps(mlx_tests.MLXTestCase):
def test_put_along_axis(self):
for ax in [None, 0, 1, 2]:
a_np = np.arange(16).reshape(2, 2, 4).astype(np.int32)
a_mlx = mx.array(a_np)
@@ -3138,6 +3137,69 @@ class TestOps(mlx_tests.MLXTestCase):
out = mx.depends(b, c)
self.assertTrue(mx.array_equal(out, b))
def test_masked_scatter(self):
# boolean mask updates matching numpy semantics
a = mx.array([1.0, 2.0, 3.0])
mask = mx.array([True, False, True])
src = mx.array([5.0, 6.0])
expected = mx.array([5.0, 2.0, 6.0])
a[mask] = src
self.assertTrue(mx.array_equal(a, expected))
# non-boolean mask raises
b = mx.array([1.0, 2.0, 3.0])
bad_mask = mx.array([1, 0, 1])
src = mx.array([4.0, 5.0])
with self.assertRaises((TypeError, ValueError)):
b[bad_mask] = src
# mask matching leading dimension selects entire trailing slices
c = mx.array([[0.0, 0.0, 0.0], [1.0, 1.0, 1.0]])
mask = mx.array([True, False])
src = mx.array([2.0, 3.0, 4.0])
expected = mx.array([[2.0, 3.0, 4.0], [1.0, 1.0, 1.0]])
c[mask] = src
self.assertTrue(mx.array_equal(c, expected))
# scalar source applies to all selected entries
c = mx.array([[0.0, 0.0, 0.0], [1.0, 1.0, 1.0]])
mask = mx.array([True, False])
src = 2.0
expected = mx.array([[2.0, 2.0, 2.0], [1.0, 1.0, 1.0]])
c[mask] = src
self.assertTrue(mx.array_equal(c, expected))
# mask with no updates leaves values unchanged
d = mx.array([[7.0, 8.0], [9.0, 10.0]])
mask = mx.zeros_like(d).astype(mx.bool_)
src = mx.array([1.0])
d[mask] = src
self.assertTrue(mx.array_equal(d, mx.array([[7.0, 8.0], [9.0, 10.0]])))
# empty mask leaves array unchanged
e = mx.zeros((0,), dtype=mx.float32)
mask = mx.zeros((0,), dtype=mx.bool_)
src = mx.zeros((0,), dtype=mx.float32)
e[mask] = src
self.assertTrue(mx.array_equal(e, mx.zeros((0,), dtype=mx.float32)))
# strided target, mask, and source derived from slices
target = mx.arange(10.0, dtype=mx.float32)[1::2]
mask = mx.array(
[False, True, False, False, True, False, False, True, False, False],
dtype=mx.bool_,
)[1::2]
src = mx.arange(-4.0, 0.0, dtype=mx.float32)[::2]
target[mask] = src
self.assertTrue(
mx.array_equal(
target, mx.array([-4.0, 3.0, 5.0, -2.0, 9.0], dtype=mx.float32)
)
)
class TestBroadcast(mlx_tests.MLXTestCase):
def test_broadcast_shapes(self):
# Basic broadcasting
self.assertEqual(mx.broadcast_shapes((1, 2, 3), (3,)), (1, 2, 3))