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:
co-authored by
Awni Hannun
Angelos Katharopoulos
Angelos Katharopoulos
parent
7f4b7e553c
commit
b3825ac149
@@ -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))
|
||||
|
||||
Reference in New Issue
Block a user