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
@@ -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);
|
||||
|
||||
@@ -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",
|
||||
}
|
||||
|
||||
@@ -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__()
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user