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
+24
View File
@@ -1,5 +1,6 @@
// Copyright © 2023-2024 Apple Inc.
#include <numeric>
#include <optional>
#include <sstream>
#include "python/src/convert.h"
@@ -885,6 +886,22 @@ auto mlx_slice_update(
return std::make_pair(true, out);
}
std::optional<mx::array> extract_boolean_mask(const nb::object& obj) {
using NDArray = nb::ndarray<nb::ro, nb::c_contig, nb::device::cpu>;
if (nb::isinstance<mx::array>(obj)) {
auto mask = nb::cast<mx::array>(obj);
if (mask.dtype() == mx::bool_) {
return mask;
}
} else if (nb::isinstance<NDArray>(obj)) {
auto mask = nb::cast<NDArray>(obj);
if (mask.dtype() == nb::dtype<bool>()) {
return nd_array_to_mlx(mask, mx::bool_);
}
}
return std::nullopt;
}
void mlx_set_item(
mx::array& src,
const nb::object& obj,
@@ -895,6 +912,13 @@ void mlx_set_item(
return;
}
if (auto mask = extract_boolean_mask(obj)) {
auto updates = to_array(v, src.dtype());
auto result = masked_scatter(src, *mask, updates);
src.overwrite_descriptor(result);
return;
}
auto [indices, updates, axes] = mlx_compute_scatter_args(src, obj, v);
if (indices.size() > 0) {
auto out = scatter(src, indices, updates, axes);
+4
View File
@@ -56,4 +56,8 @@ cuda_skip = {
"TestQuantized.test_throw",
"TestQuantized.test_vjp_scales_biases",
"TestExportImport.test_export_quantized_model",
# Masked scatter
"TestOps.test_masked_scatter",
"TestVmap.test_vmap_masked_scatter",
"TestArray.test_setitem_with_boolean_mask",
}
+12
View File
@@ -1928,6 +1928,18 @@ class TestArray(mlx_tests.MLXTestCase):
anp[:, idx] = 4
self.assertTrue(np.array_equal(a, anp))
def test_setitem_with_boolean_mask(self):
mask_np = np.zeros((10, 10), dtype=bool)
mx.arange(1000).reshape(10, 10, 10)[mask_np] = 0
mask_np = np.zeros((1, 10, 10), dtype=bool)
with self.assertRaises(ValueError):
mx.arange(1000).reshape(10, 10, 10)[mask_np] = 0
mask_np = np.zeros((10, 10, 1), dtype=bool)
with self.assertRaises(ValueError):
mx.arange(1000).reshape(10, 10, 10)[mask_np] = 0
def test_array_namespace(self):
a = mx.array(1.0)
api = a.__array_namespace__()
+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))
+87
View File
@@ -723,6 +723,93 @@ class TestVmap(mlx_tests.MLXTestCase):
out = mx.vmap(gconv, in_axes=(0, 0))(x, w)
self.assertTrue(mx.allclose(expected, out))
def test_vmap_masked_scatter(self):
def scatter_fn(x, m, src):
x[m] = src
return x
# Batched sources
a = mx.array([[10, 20, 30, 40], [50, 60, 70, 80]])
mask = mx.array([[False, True, True, True], [True, False, True, True]])
src = mx.array([[1, 2, 3], [4, 5, 6]])
expected = mx.array([[10, 1, 2, 3], [4, 60, 5, 6]])
vmap_scatter = mx.vmap(scatter_fn, in_axes=(0, 0, 0))
out = vmap_scatter(a, mask, src)
self.assertTrue(mx.array_equal(expected, out))
# Shared source across batch (matching mask populations)
a = mx.array([[0, 0, 0], [5, 5, 5]])
mask = mx.array([[True, False, True], [False, True, True]])
src = mx.array([9, 8])
expected = mx.array([[9, 0, 8], [5, 9, 8]])
vmap_scatter = mx.vmap(scatter_fn, in_axes=(0, 0, None))
out = vmap_scatter(a, mask, src)
self.assertTrue(mx.array_equal(expected, out))
# Shared destination with batched mask and sources
a = mx.array([10, 20, 30, 40])
mask = mx.array([[True, False, False, True], [False, True, True, False]])
src = mx.array([[1, 2], [3, 4]])
expected = mx.array([[1, 20, 30, 2], [10, 3, 4, 40]])
vmap_scatter = mx.vmap(scatter_fn, in_axes=(None, 0, 0))
out = vmap_scatter(a, mask, src)
self.assertTrue(mx.array_equal(expected, out))
# Shared mask across batch with batched sources
a = mx.array([[0, 0, 0, 0], [10, 20, 30, 40]])
mask = mx.array([True, False, True, False])
src = mx.array([[7, 8], [9, 10]])
expected = mx.array([[7, 0, 8, 0], [9, 20, 10, 40]])
vmap_scatter = mx.vmap(scatter_fn, in_axes=(0, None, 0))
out = vmap_scatter(a, mask, src)
self.assertTrue(mx.array_equal(expected, out))
# Uneven mask populations with scalar broadcast
a = mx.array([[0.0, 0.0, 0.0, 0.0], [10.0, 20.0, 30.0, 40.0]])
mask = mx.array([[True, False, True, True], [False, True, False, False]])
shared_src = mx.array(1.5)
expected = mx.array(
[[1.5, 0.0, 1.5, 1.5], [10.0, 1.5, 30.0, 40.0]], dtype=a.dtype
)
vmap_scatter = mx.vmap(scatter_fn, in_axes=(0, 0, None))
out = vmap_scatter(a, mask, shared_src)
self.assertTrue(mx.array_equal(expected, out))
# Shared src with identical masks must restart for each batch
a = mx.array([[0, 0, 0, 0, 0], [10, 20, 30, 40, 50]])
mask = mx.array(
[[True, True, True, False, False], [True, True, True, False, False]]
)
src = mx.array([1, 2, 3, 4, 5])
expected = mx.array([[1, 2, 3, 0, 0], [1, 2, 3, 40, 50]])
vmap_scatter = mx.vmap(scatter_fn, in_axes=(0, 0, None))
out = vmap_scatter(a, mask, src)
self.assertTrue(mx.array_equal(expected, out))
# Double vmap
a = mx.zeros((8, 8, 8))
mask = mx.random.normal((8, 8, 8)) > 0
src = mx.random.normal((8, 8))
expected = mx.stack(
[
mx.stack(
[scatter_fn(a[i, j] + 0, mask[i, j], src[i]) for j in range(8)]
)
for i in range(8)
]
)
double_scatter = mx.vmap(
mx.vmap(scatter_fn, in_axes=(0, 0, None)), in_axes=(0, 0, 0)
)
out = double_scatter(a + 0, mask, src)
self.assertTrue(mx.array_equal(expected, out))
if __name__ == "__main__":
mlx_tests.MLXTestRunner()