From bf178938b5b9eb7afb1def50c917e738f27045a3 Mon Sep 17 00:00:00 2001 From: Valentin Roussellet Date: Sun, 13 Sep 2026 23:44:21 -0700 Subject: [PATCH 1/8] initial claude pass at unique --- docs/src/python/ops.rst | 1 + mlx/ops.cpp | 97 ++++++++++++++++++++ mlx/ops.h | 13 +++ python/src/ops.cpp | 75 ++++++++++++++++ python/tests/test_export_import.py | 18 ++++ python/tests/test_ops.py | 140 +++++++++++++++++++++++++++++ 6 files changed, 344 insertions(+) diff --git a/docs/src/python/ops.rst b/docs/src/python/ops.rst index f55ba1b4d4..4f2d41e7d6 100644 --- a/docs/src/python/ops.rst +++ b/docs/src/python/ops.rst @@ -203,6 +203,7 @@ Operations triu trunc unflatten + unique unstack vecdot var diff --git a/mlx/ops.cpp b/mlx/ops.cpp index 7ec5369820..76d98b5a3c 100644 --- a/mlx/ops.cpp +++ b/mlx/ops.cpp @@ -2944,6 +2944,103 @@ array topk(const array& a, int k, int axis, StreamOrDevice s /* = {}*/) { return slice(a_partitioned, slice_starts, slice_ends, s); } +std::vector unique( + const array& a, + int size, + bool return_inverse /* = false */, + bool return_counts /* = false */, + const std::optional& fill_value /* = std::nullopt */, + StreamOrDevice s /* = {} */) { + if (size < 0) { + std::ostringstream msg; + msg << "[unique] Received negative size " << size << "."; + throw std::invalid_argument(msg.str()); + } + if (fill_value && fill_value->size() != 1) { + std::ostringstream msg; + msg << "[unique] Fill value must be a scalar but got shape " + << fill_value->shape() << "."; + throw std::invalid_argument(msg.str()); + } + auto flat = flatten(a, s); + int n = flat.size(); + + if (n == 0) { + // There is no smallest element to take the default fill value from. + if (size > 0 && !fill_value) { + throw std::invalid_argument( + "[unique] A fill value is required for an empty input with a" + " non-zero size."); + } + std::vector out; + out.push_back( + size == 0 ? flat + : full( + {size}, + astype(*fill_value, flat.dtype(), s), + flat.dtype(), + s)); + if (return_inverse) { + out.push_back(zeros(a.shape(), uint32, s)); + } + if (return_counts) { + out.push_back(zeros({size}, uint32, s)); + } + return out; + } + + auto order = argsort(flat, 0, s); + auto sorted = take(flat, order, 0, s); + + // True where a new unique value starts in the sorted array. + auto boundary = concatenate( + {array({true}), + not_equal( + slice(sorted, {1}, {n}, s), slice(sorted, {0}, {n - 1}, s), s)}, + 0, + s); + + // Index of the unique value each sorted position belongs to. + auto group = subtract( + cumsum(astype(boundary, uint32, s), 0, false, true, s), + array(1, uint32), + s); + + // A buffer that fits every group index keeps the scatter in bounds. + int buffer_size = std::max(n, size); + auto fill = fill_value ? astype(*fill_value, flat.dtype(), s) : min(flat, s); + + std::vector out; + out.push_back(slice( + scatter( + full({buffer_size}, fill, flat.dtype(), s), + group, + expand_dims(sorted, 1, s), + 0, + s), + {0}, + {size}, + s)); + if (return_inverse) { + auto inverse = + scatter(zeros({n}, uint32, s), order, expand_dims(group, 1, s), 0, s); + out.push_back(reshape(inverse, a.shape(), s)); + } + if (return_counts) { + out.push_back(slice( + scatter_add( + zeros({buffer_size}, uint32, s), + group, + ones({n, 1}, uint32, s), + 0, + s), + {0}, + {size}, + s)); + } + return out; +} + array logsumexp(const array& a, bool keepdims, StreamOrDevice s /* = {}*/) { std::vector axes(a.ndim()); std::iota(axes.begin(), axes.end(), 0); diff --git a/mlx/ops.h b/mlx/ops.h index 63c84e6cf4..2f9e88df8d 100644 --- a/mlx/ops.h +++ b/mlx/ops.h @@ -878,6 +878,19 @@ MLX_API array topk(const array& a, int k, StreamOrDevice s = {}); /** Returns topk elements of the array along a given axis. */ MLX_API array topk(const array& a, int k, int axis, StreamOrDevice s = {}); +/** + * Returns the sorted unique elements of the flattened array, and optionally + * the inverse indices and the counts. The output has the given size and the + * unused entries hold ``fill_value``. + */ +MLX_API std::vector unique( + const array& a, + int size, + bool return_inverse = false, + bool return_counts = false, + const std::optional& fill_value = std::nullopt, + StreamOrDevice s = {}); + /** Cumulative logsumexp of an array. */ MLX_API array logcumsumexp( const array& a, diff --git a/python/src/ops.cpp b/python/src/ops.cpp index 678cd984cd..670c4aa6bf 100644 --- a/python/src/ops.cpp +++ b/python/src/ops.cpp @@ -3236,6 +3236,81 @@ void init_ops(nb::module_& m) { Returns: array: The top ``k`` elements from the input. )pbdoc"); + m.def( + "unique", + [](const mx::array& a, + int size, + bool return_inverse, + bool return_counts, + const std::optional& fill_value, + mx::StreamOrDevice s) -> nb::object { + std::optional fill_value_ = std::nullopt; + if (fill_value) { + fill_value_ = to_array(fill_value.value(), a.dtype()); + } + auto out = + mx::unique(a, size, return_inverse, return_counts, fill_value_, s); + if (out.size() == 1) { + return nb::cast(out.at(0)); + } + nb::list result; + for (auto& o : out) { + result.append(o); + } + return nb::tuple(result); + }, + nb::arg(), + "size"_a, + "return_inverse"_a = false, + "return_counts"_a = false, + "fill_value"_a = nb::none(), + nb::kw_only(), + "stream"_a = nb::none(), + nb::sig( + "def unique(a: array, /, size: int, return_inverse: bool = False, return_counts: bool = False, fill_value: scalar | array | None = None, *, stream: StreamOrDevice = None) -> array | tuple[array, ...]"), + R"pbdoc( + Returns the sorted unique elements of the flattened array. + + The output has the given ``size``, so ``a`` is not evaluated. Unlike + NumPy, the size is never inferred from the values and must be given. + Entries past the last unique element hold ``fill_value``. The count of + a padded entry is ``0``, so ``mx.sum(counts > 0)`` gives the number of + unique elements. + + Args: + a (array): Input array. + size (int): The size of the output. Use the size of the flattened + ``a`` to hold every unique element. A smaller size keeps only the + smallest ``size`` unique elements. + return_inverse (bool, optional): If ``True``, also return the + indices of the unique array that rebuild ``a``. The indices have + the same shape as ``a``. Default: ``False``. + return_counts (bool, optional): If ``True``, also return the number + of times each unique element occurs in ``a``. Default: ``False``. + fill_value (scalar or array, optional): The value of the entries + past the last unique element. If ``None``, defaults to the + smallest element of ``a``, which an empty ``a`` does not have. + Default: ``None``. + + Returns: + array or tuple(array, ...): The sorted unique elements. If + ``return_inverse`` or ``return_counts`` is ``True``, a tuple with + the requested arrays in the order values, inverse, counts. + + Example: + >>> a = mx.array([2, 1, 2, 3, 1]) + >>> mx.unique(a, 3) + array([1, 2, 3], dtype=int32) + >>> mx.unique(a, 5) + array([1, 2, 3, 1, 1], dtype=int32) + >>> values, inverse, counts = mx.unique(a, 4, True, True, fill_value=0) + >>> values + array([1, 2, 3, 0], dtype=int32) + >>> inverse + array([1, 0, 1, 2, 0], dtype=uint32) + >>> counts + array([2, 2, 1, 0], dtype=uint32) + )pbdoc"); m.def( "broadcast_to", [](const ScalarOrArray& a, const mx::Shape& shape, mx::StreamOrDevice s) { diff --git a/python/tests/test_export_import.py b/python/tests/test_export_import.py index 23f0645b3a..7e0be7249d 100644 --- a/python/tests/test_export_import.py +++ b/python/tests/test_export_import.py @@ -332,6 +332,24 @@ def fun(a, v): out = imported_fun(x, y)[0] self.assertTrue(mx.array_equal(expected, out)) + def test_export_unique(self): + path = os.path.join(self.test_dir, "fn.mlxfn") + + # the fixed output size is what makes this exportable + for size, fill_value in ((10, None), (4, -1)): + + def fun(a): + return mx.unique(a, size, True, True, fill_value=fill_value) + + x = mx.random.randint(0, 5, shape=(10,)) + mx.export_function(path, fun, x) + imported_fun = mx.import_function(path) + expected = fun(x) + out = imported_fun(x) + self.assertEqual(len(out), len(expected)) + for e, o in zip(expected, out): + self.assertTrue(mx.array_equal(e, o)) + def test_export_conv(self): path = os.path.join(self.test_dir, "fn.mlxfn") diff --git a/python/tests/test_ops.py b/python/tests/test_ops.py index d5568610e8..1b8b845dbb 100644 --- a/python/tests/test_ops.py +++ b/python/tests/test_ops.py @@ -3122,6 +3122,146 @@ def expect(out, want): with self.assertRaises(ValueError): mx.searchsorted(mx.array([1.0, 2.0]), mx.array([1.0]), side="middle") + def test_unique(self): + # size is required, and the exact size trims the padding away + a = mx.array([2, 1, 2, 3, 1]) + with self.assertRaises(TypeError): + mx.unique(a) + self.assertTrue(mx.array_equal(mx.unique(a, 3), mx.array([1, 2, 3]))) + self.assertEqual(mx.unique(a, 3).dtype, mx.int32) + + # the flattened input size always holds every unique element + self.assertTrue(mx.array_equal(mx.unique(a, a.size), mx.array([1, 2, 3, 1, 1]))) + + values, inverse, counts = mx.unique(a, 4, True, True, fill_value=0) + self.assertTrue(mx.array_equal(values, mx.array([1, 2, 3, 0]))) + self.assertTrue(mx.array_equal(inverse, mx.array([1, 0, 1, 2, 0]))) + self.assertTrue(mx.array_equal(counts, mx.array([2, 2, 1, 0]))) + self.assertEqual(inverse.dtype, mx.uint32) + self.assertEqual(counts.dtype, mx.uint32) + + # zero counts mark the padding, so they give the number of uniques + counts = mx.unique(a, a.size, False, True)[1] + self.assertEqual(mx.sum(counts > 0).item(), 3) + + # only the requested extras come back + self.assertEqual(len(mx.unique(a, 3, True)), 2) + self.assertEqual(len(mx.unique(a, 3, False, True)), 2) + + # a size larger than the input pads further + self.assertTrue( + mx.array_equal( + mx.unique(a, 7, fill_value=-1), + mx.array([1, 2, 3, -1, -1, -1, -1]), + ) + ) + + # a size smaller than the number of unique elements keeps the smallest + self.assertTrue(mx.array_equal(mx.unique(a, 2), mx.array([1, 2]))) + counts = mx.unique(a, 2, False, True)[1] + self.assertTrue(mx.array_equal(counts, mx.array([2, 2]))) + + # the input is flattened, but the inverse keeps the input shape + b = mx.array([[3, 1, 3], [2, 1, 4]]) + values, inverse = mx.unique(b, 4, True) + self.assertTrue(mx.array_equal(values, mx.array([1, 2, 3, 4]))) + self.assertEqual(inverse.shape, (2, 3)) + self.assertTrue(mx.array_equal(values[inverse], b)) + + # 0-d input + values, inverse = mx.unique(mx.array(5), 1, True) + self.assertTrue(mx.array_equal(values, mx.array([5]))) + self.assertEqual(inverse.shape, ()) + + # empty input with a zero size + values, inverse, counts = mx.unique(mx.array([], mx.float32), 0, True, True) + self.assertEqual(values.shape, (0,)) + self.assertEqual(values.dtype, mx.float32) + self.assertEqual(inverse.shape, (0,)) + self.assertEqual(counts.shape, (0,)) + + # an empty input has no smallest element to default the fill value to + with self.assertRaises(ValueError): + mx.unique(mx.array([], mx.float32), 3) + self.assertTrue( + mx.array_equal( + mx.unique(mx.array([], mx.float32), 3, fill_value=7), + mx.array([7.0, 7.0, 7.0]), + ) + ) + + # every element identical, and every element distinct + values, counts = mx.unique(mx.full((7,), 4, mx.int32), 7, False, True) + self.assertTrue(mx.array_equal(values, mx.full((7,), 4, mx.int32))) + self.assertTrue(mx.array_equal(counts, mx.array([7, 0, 0, 0, 0, 0, 0]))) + self.assertTrue(mx.array_equal(mx.unique(mx.arange(6), 6), mx.arange(6))) + + # floats, including negatives and a repeated zero + f = mx.array([0.0, -1.5, 2.25, -1.5, 0.0], mx.float32) + self.assertTrue(mx.array_equal(mx.unique(f, 3), mx.array([-1.5, 0.0, 2.25]))) + + # None selects the default fill + self.assertTrue( + mx.array_equal(mx.unique(a, 5, fill_value=None), mx.unique(a, 5)) + ) + + # a fill value is cast to the input dtype + self.assertEqual(mx.unique(f, 4, fill_value=1).dtype, mx.float32) + + # like torch, NaN never equals itself so each one is kept + nan = float("nan") + values = mx.unique(mx.array([1.0, nan, 2.0, nan], mx.float32), 4) + self.assertTrue(mx.array_equal(values[:2], mx.array([1.0, 2.0]))) + self.assertTrue(bool(mx.all(mx.isnan(values[2:])))) + + rng = np.random.RandomState(0) + for n in (1, 2, 17, 1000, 5000): + a_np = rng.randint(-20, 20, size=n).astype(np.int32) + a_mx = mx.array(a_np) + v_np, i_np, c_np = np.unique(a_np, return_inverse=True, return_counts=True) + values, inverse, counts = mx.unique(a_mx, len(v_np), True, True) + self.assertTrue(np.array_equal(np.array(values), v_np)) + self.assertTrue( + np.array_equal(np.array(inverse).reshape(-1), i_np.reshape(-1)) + ) + self.assertTrue(np.array_equal(np.array(counts), c_np)) + # the inverse rebuilds the input + self.assertTrue(mx.array_equal(values[inverse], a_mx)) + + # a padded size holds the same elements up front + padded, pad_counts = mx.unique(a_mx, n, False, True) + self.assertEqual(padded.shape, (n,)) + self.assertTrue(np.array_equal(np.array(padded)[: len(v_np)], v_np)) + self.assertEqual(mx.sum(pad_counts > 0).item(), len(v_np)) + + # non contiguous input + base_np = rng.randint(0, 5, size=(4, 6)).astype(np.int32) + base = mx.array(base_np) + for v_mx, v_np in [ + (base.T, base_np.T), + (base[::2], base_np[::2]), + (base[::-1], base_np[::-1]), + (base[:, ::-1], base_np[:, ::-1]), + (mx.broadcast_to(base[0], (3, 6)), np.broadcast_to(base_np[0], (3, 6))), + ]: + expected = np.unique(v_np) + self.assertTrue( + np.array_equal(np.array(mx.unique(v_mx, len(expected))), expected) + ) + + # the output size is static, so this composes with the transforms + c = mx.array([2, 1, 2, 3, 1]) + self.assertTrue( + mx.array_equal(mx.compile(lambda x: mx.unique(x, 3))(c), mx.unique(c, 3)) + ) + batched = mx.vmap(lambda x: mx.unique(x, 3))(mx.stack([c, c[::-1]])) + self.assertTrue(mx.array_equal(batched, mx.stack([mx.unique(c, 3)] * 2))) + + with self.assertRaises(ValueError): + mx.unique(a, -1) + with self.assertRaises(ValueError): + mx.unique(a, 3, fill_value=mx.array([1, 2])) + @unittest.skipIf( os.getenv("LOW_MEMORY", None) is not None, "This test requires a lot of memory", From f57c57151a8369667c3487c1b854364d77196368 Mon Sep 17 00:00:00 2001 From: Valentin Roussellet Date: Mon, 14 Sep 2026 00:30:15 -0700 Subject: [PATCH 2/8] add extra comments --- mlx/ops.cpp | 38 ++++++++++++++++++++++++++------------ mlx/ops.h | 6 ++++-- python/src/ops.cpp | 30 +++++++++++++++--------------- python/tests/test_ops.py | 26 +++++++++++++++++++++----- 4 files changed, 66 insertions(+), 34 deletions(-) diff --git a/mlx/ops.cpp b/mlx/ops.cpp index 76d98b5a3c..85c8706340 100644 --- a/mlx/ops.cpp +++ b/mlx/ops.cpp @@ -2951,6 +2951,7 @@ std::vector unique( bool return_counts /* = false */, const std::optional& fill_value /* = std::nullopt */, StreamOrDevice s /* = {} */) { + // Validate args if (size < 0) { std::ostringstream msg; msg << "[unique] Received negative size " << size << "."; @@ -2958,15 +2959,18 @@ std::vector unique( } if (fill_value && fill_value->size() != 1) { std::ostringstream msg; - msg << "[unique] Fill value must be a scalar but got shape " + msg << "[unique] Fill value must be a scalar, but got shape " << fill_value->shape() << "."; throw std::invalid_argument(msg.str()); } - auto flat = flatten(a, s); - int n = flat.size(); + const auto flat = flatten(a, s); + const int n = flat.size(); + // Handle the edge case of an empty array. + // The output is to be filled with `fill_value`. if (n == 0) { - // There is no smallest element to take the default fill value from. + // There is no smallest element to take the default fill value from + // so we throw (similar to jax) if (size > 0 && !fill_value) { throw std::invalid_argument( "[unique] A fill value is required for an empty input with a" @@ -2989,27 +2993,34 @@ std::vector unique( return out; } - auto order = argsort(flat, 0, s); - auto sorted = take(flat, order, 0, s); + // Sort the array + const auto order = argsort(flat, 0, s); + const auto sorted = take(flat, order, 0, s); - // True where a new unique value starts in the sorted array. - auto boundary = concatenate( + // Do edge detection on the sorted array to get a mask with + // true where a new unique value starts in the sorted array. + const auto boundary = concatenate( {array({true}), not_equal( slice(sorted, {1}, {n}, s), slice(sorted, {0}, {n - 1}, s), s)}, 0, s); - // Index of the unique value each sorted position belongs to. - auto group = subtract( + // Cumulative sum on boundary to get to the index of the unique + // value each sorted position belongs to. + const auto group = subtract( cumsum(astype(boundary, uint32, s), 0, false, true, s), array(1, uint32), s); // A buffer that fits every group index keeps the scatter in bounds. - int buffer_size = std::max(n, size); - auto fill = fill_value ? astype(*fill_value, flat.dtype(), s) : min(flat, s); + const int buffer_size = std::max(n, size); + // Use the minimum of the array to pad the result if size is bigger + // than the sorted array. + const auto fill = fill_value ? astype(*fill_value, flat.dtype(), s) + : slice(sorted, {0}, {1}, s); + // Build the output arrays std::vector out; out.push_back(slice( scatter( @@ -3026,6 +3037,9 @@ std::vector unique( scatter(zeros({n}, uint32, s), order, expand_dims(group, 1, s), 0, s); out.push_back(reshape(inverse, a.shape(), s)); } + // If output is padded with fill value, counts is padded with zeros + // so that ``count.sum() == a.size()`` holds if output was not + // truncated. if (return_counts) { out.push_back(slice( scatter_add( diff --git a/mlx/ops.h b/mlx/ops.h index 2f9e88df8d..f840338a17 100644 --- a/mlx/ops.h +++ b/mlx/ops.h @@ -880,8 +880,10 @@ MLX_API array topk(const array& a, int k, int axis, StreamOrDevice s = {}); /** * Returns the sorted unique elements of the flattened array, and optionally - * the inverse indices and the counts. The output has the given size and the - * unused entries hold ``fill_value``. + * the inverse indices and the counts. The output has the given size, and + * is truncated if ``size`` is smaller than the number of unique values. + * If ``size`` is larger, the result is padded with ``fill_value``, or + * ``min(a)`` if no fill value is provided. */ MLX_API std::vector unique( const array& a, diff --git a/python/src/ops.cpp b/python/src/ops.cpp index 670c4aa6bf..a654d82c09 100644 --- a/python/src/ops.cpp +++ b/python/src/ops.cpp @@ -1560,16 +1560,16 @@ void init_ops(nb::module_& m) { "stream"_a = nb::none(), R"pbdoc( Return the Bartlett window. - + The Bartlett window is a taper formed by using a weighted cosine. .. math:: w(n) = 1 - \frac{2|n - (M-1)/2|}{M-1} \qquad 0 \le n \le M-1 - + Args: M (int): Number of points in the output window. - + Returns: array: The window, with the maximum value normalized to one (the value one appears only if the number of samples is odd). @@ -1582,16 +1582,16 @@ void init_ops(nb::module_& m) { "stream"_a = nb::none(), R"pbdoc( Return the Hanning window. - + The Hanning window is a taper formed by using a weighted cosine. .. math:: w(n) = 0.5 - 0.5 \cos\left(\frac{2\pi n}{M-1}\right) \qquad 0 \le n \le M-1 - + Args: M (int): Number of points in the output window. - + Returns: array: The window, with the maximum value normalized to one (the value one appears only if the number of samples is odd). @@ -1629,16 +1629,16 @@ void init_ops(nb::module_& m) { "def blackman(M: int, *, stream: StreamOrDevice = None) -> array"), // <--- J'ai rajouté ça R"pbdoc( Return the Blackman window. - + The Blackman window is a taper formed by using the first three terms of a summation of cosines. .. math:: w(n) = 0.42 - 0.5 \cos\left(\frac{2\pi n}{M-1}\right) + 0.08 \cos\left(\frac{4\pi n}{M-1}\right) \qquad 0 \le n \le M-1 - + Args: M (int): Number of points in the output window. - + Returns: array: The window, with the maximum value normalized to one (the value one appears only if the number of samples is odd). @@ -3272,24 +3272,24 @@ void init_ops(nb::module_& m) { Returns the sorted unique elements of the flattened array. The output has the given ``size``, so ``a`` is not evaluated. Unlike - NumPy, the size is never inferred from the values and must be given. + NumPy, the size is never inferred from the values and must be provided. Entries past the last unique element hold ``fill_value``. The count of a padded entry is ``0``, so ``mx.sum(counts > 0)`` gives the number of unique elements. Args: a (array): Input array. - size (int): The size of the output. Use the size of the flattened - ``a`` to hold every unique element. A smaller size keeps only the - smallest ``size`` unique elements. + size (int): The size of the output. If the size is smaller than the + number of unique elements of ``a``, the output will be truncated. + If the size is larger, it will be padded with ``fill_value``. return_inverse (bool, optional): If ``True``, also return the indices of the unique array that rebuild ``a``. The indices have the same shape as ``a``. Default: ``False``. return_counts (bool, optional): If ``True``, also return the number of times each unique element occurs in ``a``. Default: ``False``. fill_value (scalar or array, optional): The value of the entries - past the last unique element. If ``None``, defaults to the - smallest element of ``a``, which an empty ``a`` does not have. + past the last unique element. If ``None``, this defaults to the + minimum unique value. Default: ``None``. Returns: diff --git a/python/tests/test_ops.py b/python/tests/test_ops.py index 1b8b845dbb..0aa691a44f 100644 --- a/python/tests/test_ops.py +++ b/python/tests/test_ops.py @@ -3123,10 +3123,12 @@ def expect(out, want): mx.searchsorted(mx.array([1.0, 2.0]), mx.array([1.0]), side="middle") def test_unique(self): - # size is required, and the exact size trims the padding away + # size is required a = mx.array([2, 1, 2, 3, 1]) with self.assertRaises(TypeError): mx.unique(a) + + # passing the correct size (number of unique values) produces a sorted array self.assertTrue(mx.array_equal(mx.unique(a, 3), mx.array([1, 2, 3]))) self.assertEqual(mx.unique(a, 3).dtype, mx.int32) @@ -3180,9 +3182,7 @@ def test_unique(self): self.assertEqual(inverse.shape, (0,)) self.assertEqual(counts.shape, (0,)) - # an empty input has no smallest element to default the fill value to - with self.assertRaises(ValueError): - mx.unique(mx.array([], mx.float32), 3) + # empty input is handled correctly when given a fill_value self.assertTrue( mx.array_equal( mx.unique(mx.array([], mx.float32), 3, fill_value=7), @@ -3190,6 +3190,11 @@ def test_unique(self): ) ) + # empty input has no smallest element to default without a fill_value + # and produces an error + with self.assertRaises(ValueError): + mx.unique(mx.array([], mx.float32), 3) + # every element identical, and every element distinct values, counts = mx.unique(mx.full((7,), 4, mx.int32), 7, False, True) self.assertTrue(mx.array_equal(values, mx.full((7,), 4, mx.int32))) @@ -3214,6 +3219,16 @@ def test_unique(self): self.assertTrue(mx.array_equal(values[:2], mx.array([1.0, 2.0]))) self.assertTrue(bool(mx.all(mx.isnan(values[2:])))) + # NaN sorts last, so it is never the default fill + # note that mx.min(a) would return `nan` but we keep + # parity with jax which pads with the min of the sorted + # array in this case. + values = mx.unique(mx.array([3.0, 1.0, nan, 2.0], mx.float32), 6) + self.assertTrue(mx.array_equal(values[:3], mx.array([1.0, 2.0, 3.0]))) + self.assertTrue(bool(mx.isnan(values[3]))) + self.assertTrue(mx.array_equal(values[4:], mx.array([1.0, 1.0]))) + + # Compare the output with np.unique rng = np.random.RandomState(0) for n in (1, 2, 17, 1000, 5000): a_np = rng.randint(-20, 20, size=n).astype(np.int32) @@ -3234,7 +3249,7 @@ def test_unique(self): self.assertTrue(np.array_equal(np.array(padded)[: len(v_np)], v_np)) self.assertEqual(mx.sum(pad_counts > 0).item(), len(v_np)) - # non contiguous input + # test various non-contiguous slices base_np = rng.randint(0, 5, size=(4, 6)).astype(np.int32) base = mx.array(base_np) for v_mx, v_np in [ @@ -3257,6 +3272,7 @@ def test_unique(self): batched = mx.vmap(lambda x: mx.unique(x, 3))(mx.stack([c, c[::-1]])) self.assertTrue(mx.array_equal(batched, mx.stack([mx.unique(c, 3)] * 2))) + # Test input validation with self.assertRaises(ValueError): mx.unique(a, -1) with self.assertRaises(ValueError): From dd87aaf73bcb0456a48bb036b1749fea8132ba72 Mon Sep 17 00:00:00 2001 From: Valentin Roussellet Date: Mon, 14 Sep 2026 01:31:31 -0700 Subject: [PATCH 3/8] jax parity, differentiability --- mlx/ops.cpp | 81 +++++++++++++++++++++++++++------------- mlx/ops.h | 7 +++- python/src/ops.cpp | 43 ++++++++++++--------- python/tests/test_ops.py | 43 +++++++++++++++++++++ 4 files changed, 130 insertions(+), 44 deletions(-) diff --git a/mlx/ops.cpp b/mlx/ops.cpp index 85c8706340..d7b69e736a 100644 --- a/mlx/ops.cpp +++ b/mlx/ops.cpp @@ -2959,7 +2959,7 @@ std::vector unique( } if (fill_value && fill_value->size() != 1) { std::ostringstream msg; - msg << "[unique] Fill value must be a scalar, but got shape " + msg << "[unique] Fill value must have one element, but got shape " << fill_value->shape() << "."; throw std::invalid_argument(msg.str()); } @@ -2981,7 +2981,7 @@ std::vector unique( size == 0 ? flat : full( {size}, - astype(*fill_value, flat.dtype(), s), + astype(reshape(*fill_value, {}, s), flat.dtype(), s), flat.dtype(), s)); if (return_inverse) { @@ -2993,9 +2993,13 @@ std::vector unique( return out; } - // Sort the array - const auto order = argsort(flat, 0, s); - const auto sorted = take(flat, order, 0, s); + // Sort the array. The argsort is only needed to build the inverse, and the + // indices it feeds are not differentiable. + std::optional order; + if (return_inverse) { + order = stop_gradient(argsort(flat, 0, s), s); + } + const auto sorted = order ? take(flat, *order, 0, s) : sort(flat, 0, s); // Do edge detection on the sorted array to get a mask with // true where a new unique value starts in the sorted array. @@ -3007,34 +3011,61 @@ std::vector unique( s); // Cumulative sum on boundary to get to the index of the unique - // value each sorted position belongs to. - const auto group = subtract( - cumsum(astype(boundary, uint32, s), 0, false, true, s), - array(1, uint32), + // value each sorted position belongs to. It indexes the output, so it is + // not differentiable. + const auto group = stop_gradient( + subtract( + cumsum(astype(boundary, uint32, s), 0, false, true, s), + array(1, uint32), + s), s); // A buffer that fits every group index keeps the scatter in bounds. const int buffer_size = std::max(n, size); - // Use the minimum of the array to pad the result if size is bigger - // than the sorted array. - const auto fill = fill_value ? astype(*fill_value, flat.dtype(), s) - : slice(sorted, {0}, {1}, s); + // Use the smallest element of the sorted array to pad the result if size is + // bigger than the number of unique values. + const auto fill = fill_value + ? astype(reshape(*fill_value, {}, s), flat.dtype(), s) + : slice(sorted, {0}, {1}, s); + + // Scatter positions in the sorted array rather than the values themselves: + // a GPU scatter rejects 8 byte types such as int64. A slot holds its + // position plus one, so a zero marks a slot no unique value landed in. + // Only the first position of a group is non-zero, so the max over a group + // picks it without depending on the order duplicate indices are written in. + const auto slots = stop_gradient( + slice( + scatter_max( + zeros({buffer_size}, uint32, s), + group, + expand_dims( + where( + boundary, + add(arange(n, uint32, s), array(1, uint32), s), + array(0, uint32), + s), + 1, + s), + 0, + s), + {0}, + {size}, + s), + s); + const auto used = greater(slots, array(0, uint32), s); + const auto positions = + subtract(maximum(slots, array(1, uint32), s), array(1, uint32), s); // Build the output arrays std::vector out; - out.push_back(slice( - scatter( - full({buffer_size}, fill, flat.dtype(), s), - group, - expand_dims(sorted, 1, s), - 0, - s), - {0}, - {size}, - s)); + out.push_back(where(used, take(sorted, positions, 0, s), fill, s)); if (return_inverse) { - auto inverse = - scatter(zeros({n}, uint32, s), order, expand_dims(group, 1, s), 0, s); + // Clamp so that the indices stay inside a truncated output. Without a + // truncation every group index is already smaller than size. + const auto clamped = + minimum(group, array(std::max(size - 1, 0), uint32), s); + auto inverse = scatter( + zeros({n}, uint32, s), *order, expand_dims(clamped, 1, s), 0, s); out.push_back(reshape(inverse, a.shape(), s)); } // If output is padded with fill value, counts is padded with zeros diff --git a/mlx/ops.h b/mlx/ops.h index f840338a17..de4773222b 100644 --- a/mlx/ops.h +++ b/mlx/ops.h @@ -882,8 +882,11 @@ MLX_API array topk(const array& a, int k, int axis, StreamOrDevice s = {}); * Returns the sorted unique elements of the flattened array, and optionally * the inverse indices and the counts. The output has the given size, and * is truncated if ``size`` is smaller than the number of unique values. - * If ``size`` is larger, the result is padded with ``fill_value``, or - * ``min(a)`` if no fill value is provided. + * If ``size`` is larger, the result is padded with ``fill_value``, or the + * smallest unique value if no fill value is provided. An empty input has no + * such value, so it throws unless ``size`` is zero or a fill value is given. + * A truncated output loses the counts of the values it drops, and its inverse + * indices are clamped to the last entry. */ MLX_API std::vector unique( const array& a, diff --git a/python/src/ops.cpp b/python/src/ops.cpp index a654d82c09..c7f9688268 100644 --- a/python/src/ops.cpp +++ b/python/src/ops.cpp @@ -1560,16 +1560,16 @@ void init_ops(nb::module_& m) { "stream"_a = nb::none(), R"pbdoc( Return the Bartlett window. - + The Bartlett window is a taper formed by using a weighted cosine. .. math:: w(n) = 1 - \frac{2|n - (M-1)/2|}{M-1} \qquad 0 \le n \le M-1 - + Args: M (int): Number of points in the output window. - + Returns: array: The window, with the maximum value normalized to one (the value one appears only if the number of samples is odd). @@ -1582,16 +1582,16 @@ void init_ops(nb::module_& m) { "stream"_a = nb::none(), R"pbdoc( Return the Hanning window. - + The Hanning window is a taper formed by using a weighted cosine. .. math:: w(n) = 0.5 - 0.5 \cos\left(\frac{2\pi n}{M-1}\right) \qquad 0 \le n \le M-1 - + Args: M (int): Number of points in the output window. - + Returns: array: The window, with the maximum value normalized to one (the value one appears only if the number of samples is odd). @@ -1629,16 +1629,16 @@ void init_ops(nb::module_& m) { "def blackman(M: int, *, stream: StreamOrDevice = None) -> array"), // <--- J'ai rajouté ça R"pbdoc( Return the Blackman window. - + The Blackman window is a taper formed by using the first three terms of a summation of cosines. .. math:: w(n) = 0.42 - 0.5 \cos\left(\frac{2\pi n}{M-1}\right) + 0.08 \cos\left(\frac{4\pi n}{M-1}\right) \qquad 0 \le n \le M-1 - + Args: M (int): Number of points in the output window. - + Returns: array: The window, with the maximum value normalized to one (the value one appears only if the number of samples is odd). @@ -3271,17 +3271,26 @@ void init_ops(nb::module_& m) { R"pbdoc( Returns the sorted unique elements of the flattened array. - The output has the given ``size``, so ``a`` is not evaluated. Unlike - NumPy, the size is never inferred from the values and must be provided. - Entries past the last unique element hold ``fill_value``. The count of - a padded entry is ``0``, so ``mx.sum(counts > 0)`` gives the number of - unique elements. + The shape of the output does not depend on the values in ``a``, so + ``a`` is not evaluated. Unlike NumPy, the size is never inferred from + the values and must be provided. + + Entries past the last unique element hold ``fill_value`` and their + count is ``0``, so ``mx.sum(counts > 0)`` gives the number of unique + elements as long as the output is not truncated. + + A truncated output keeps the smallest ``size`` unique elements, and + ``inverse`` and ``counts`` then stop describing all of ``a`` and stop + agreeing with each other: the counts of the dropped elements are gone, + while their indices in ``inverse`` are clamped to the last entry. + ``inverse`` is meaningless for a ``size`` of ``0``, since there is no + element left for it to point at. Args: a (array): Input array. - size (int): The size of the output. If the size is smaller than the - number of unique elements of ``a``, the output will be truncated. - If the size is larger, it will be padded with ``fill_value``. + size (int): The size of the output. If the size is smaller than + the number of unique elements of ``a``, the output is truncated. + If it is larger, the output is padded with ``fill_value``. return_inverse (bool, optional): If ``True``, also return the indices of the unique array that rebuild ``a``. The indices have the same shape as ``a``. Default: ``False``. diff --git a/python/tests/test_ops.py b/python/tests/test_ops.py index 0aa691a44f..5655058c03 100644 --- a/python/tests/test_ops.py +++ b/python/tests/test_ops.py @@ -3277,6 +3277,49 @@ def test_unique(self): mx.unique(a, -1) with self.assertRaises(ValueError): mx.unique(a, 3, fill_value=mx.array([1, 2])) + # any one element array fill value works, whatever its rank + for fv in (7, mx.array(7), mx.array([7]), mx.array([[7]])): + self.assertTrue( + mx.array_equal(mx.unique(a, 4, fill_value=fv), mx.array([1, 2, 3, 7])) + ) + + # inverse is clamped, so it points nowhere when size is zero + values, inverse = mx.unique(a, 0, True) + self.assertEqual(values.size, 0) + self.assertTrue(mx.array_equal(inverse, mx.zeros(a.shape, mx.uint32))) + + # 8 byte types have to avoid a scatter on the data, since the GPU + # scatter does not support them + for dtype in (mx.int64, mx.uint64, mx.complex64): + values, inverse, counts = mx.unique( + mx.array([3, 1, 2, 1], dtype), 4, True, True + ) + self.assertEqual(values.dtype, dtype) + self.assertTrue(mx.array_equal(values, mx.array([1, 2, 3, 1], dtype))) + self.assertTrue(mx.array_equal(inverse, mx.array([2, 0, 1, 0]))) + self.assertTrue(mx.array_equal(counts, mx.array([2, 1, 1, 0]))) + # the dtype numpy hands over by default + self.assertTrue( + mx.array_equal(mx.unique(mx.array(np.arange(5)), 5), mx.arange(5)) + ) + + # truncation clamps the inverse so that it stays in range, otherwise + # values[inverse] reads past the end of the output + values, inverse = mx.unique(a, 2, True) + self.assertTrue(mx.array_equal(values, mx.array([1, 2]))) + self.assertTrue(bool(mx.all(inverse < values.size))) + self.assertTrue(mx.array_equal(values[inverse], mx.array([2, 1, 2, 2, 1]))) + + # the gradient credits the first occurrence of each unique value + grad = mx.grad(lambda x: mx.sum(mx.unique(x, 3)))( + mx.array([3.0, 1.0, 2.0, 1.0]) + ) + self.assertTrue(mx.array_equal(grad, mx.array([1.0, 1.0, 1.0, 0.0]))) + # and a padded output routes the padding to that same element + grad = mx.grad(lambda x: mx.sum(mx.unique(x, 4)))( + mx.array([3.0, 1.0, 2.0, 1.0]) + ) + self.assertTrue(mx.array_equal(grad, mx.array([1.0, 2.0, 1.0, 0.0]))) @unittest.skipIf( os.getenv("LOW_MEMORY", None) is not None, From 658910118ddc7f4a26c6ca0624965588101e8c9c Mon Sep 17 00:00:00 2001 From: Valentin Roussellet Date: Mon, 14 Sep 2026 01:51:23 -0700 Subject: [PATCH 4/8] count type to int32 (jax parity) --- mlx/ops.cpp | 6 +++--- python/src/ops.cpp | 2 +- python/tests/test_ops.py | 2 +- 3 files changed, 5 insertions(+), 5 deletions(-) diff --git a/mlx/ops.cpp b/mlx/ops.cpp index d7b69e736a..fc139d6104 100644 --- a/mlx/ops.cpp +++ b/mlx/ops.cpp @@ -2988,7 +2988,7 @@ std::vector unique( out.push_back(zeros(a.shape(), uint32, s)); } if (return_counts) { - out.push_back(zeros({size}, uint32, s)); + out.push_back(zeros({size}, int32, s)); } return out; } @@ -3074,9 +3074,9 @@ std::vector unique( if (return_counts) { out.push_back(slice( scatter_add( - zeros({buffer_size}, uint32, s), + zeros({buffer_size}, int32, s), group, - ones({n, 1}, uint32, s), + ones({n, 1}, int32, s), 0, s), {0}, diff --git a/python/src/ops.cpp b/python/src/ops.cpp index c7f9688268..7ed3fc5238 100644 --- a/python/src/ops.cpp +++ b/python/src/ops.cpp @@ -3318,7 +3318,7 @@ void init_ops(nb::module_& m) { >>> inverse array([1, 0, 1, 2, 0], dtype=uint32) >>> counts - array([2, 2, 1, 0], dtype=uint32) + array([2, 2, 1, 0], dtype=int32) )pbdoc"); m.def( "broadcast_to", diff --git a/python/tests/test_ops.py b/python/tests/test_ops.py index 5655058c03..1619b5301f 100644 --- a/python/tests/test_ops.py +++ b/python/tests/test_ops.py @@ -3140,7 +3140,7 @@ def test_unique(self): self.assertTrue(mx.array_equal(inverse, mx.array([1, 0, 1, 2, 0]))) self.assertTrue(mx.array_equal(counts, mx.array([2, 2, 1, 0]))) self.assertEqual(inverse.dtype, mx.uint32) - self.assertEqual(counts.dtype, mx.uint32) + self.assertEqual(counts.dtype, mx.int32) # zero counts mark the padding, so they give the number of uniques counts = mx.unique(a, a.size, False, True)[1] From b2ee0eee375c98cbc10220f10e4a84aca73bc9d4 Mon Sep 17 00:00:00 2001 From: Valentin Roussellet Date: Mon, 14 Sep 2026 09:13:31 -0700 Subject: [PATCH 5/8] comment simplification --- mlx/ops.cpp | 34 +++++++++++------------------- python/tests/test_ops.py | 45 ++++++++++++++++++++++++++++++++++++---- 2 files changed, 53 insertions(+), 26 deletions(-) diff --git a/mlx/ops.cpp b/mlx/ops.cpp index fc139d6104..2969bb0bd2 100644 --- a/mlx/ops.cpp +++ b/mlx/ops.cpp @@ -2969,8 +2969,7 @@ std::vector unique( // Handle the edge case of an empty array. // The output is to be filled with `fill_value`. if (n == 0) { - // There is no smallest element to take the default fill value from - // so we throw (similar to jax) + // Throw as is no smallest element to take the default fill value from. if (size > 0 && !fill_value) { throw std::invalid_argument( "[unique] A fill value is required for an empty input with a" @@ -2993,16 +2992,15 @@ std::vector unique( return out; } - // Sort the array. The argsort is only needed to build the inverse, and the - // indices it feeds are not differentiable. + // Sort the array. Not differentiable. std::optional order; if (return_inverse) { order = stop_gradient(argsort(flat, 0, s), s); } const auto sorted = order ? take(flat, *order, 0, s) : sort(flat, 0, s); - // Do edge detection on the sorted array to get a mask with - // true where a new unique value starts in the sorted array. + // Edge detection on the sorted array, true where a new unique + // value starts in the sorted array. const auto boundary = concatenate( {array({true}), not_equal( @@ -3010,9 +3008,7 @@ std::vector unique( 0, s); - // Cumulative sum on boundary to get to the index of the unique - // value each sorted position belongs to. It indexes the output, so it is - // not differentiable. + // Cumsum on boundary gets the index of unique elements. Not differentiable. const auto group = stop_gradient( subtract( cumsum(astype(boundary, uint32, s), 0, false, true, s), @@ -3020,19 +3016,14 @@ std::vector unique( s), s); - // A buffer that fits every group index keeps the scatter in bounds. - const int buffer_size = std::max(n, size); - // Use the smallest element of the sorted array to pad the result if size is - // bigger than the number of unique values. + // Use smallest element of the sorted array (index 0) as padding by default. const auto fill = fill_value ? astype(reshape(*fill_value, {}, s), flat.dtype(), s) : slice(sorted, {0}, {1}, s); - // Scatter positions in the sorted array rather than the values themselves: - // a GPU scatter rejects 8 byte types such as int64. A slot holds its - // position plus one, so a zero marks a slot no unique value landed in. - // Only the first position of a group is non-zero, so the max over a group - // picks it without depending on the order duplicate indices are written in. + const int buffer_size = std::max(n, size); + // Scatter positions rather than values in the sorted array (64 bit values + // would fail to scatter on the GPU). const auto slots = stop_gradient( slice( scatter_max( @@ -3056,7 +3047,7 @@ std::vector unique( const auto positions = subtract(maximum(slots, array(1, uint32), s), array(1, uint32), s); - // Build the output arrays + // Build the output arrays. std::vector out; out.push_back(where(used, take(sorted, positions, 0, s), fill, s)); if (return_inverse) { @@ -3064,13 +3055,12 @@ std::vector unique( // truncation every group index is already smaller than size. const auto clamped = minimum(group, array(std::max(size - 1, 0), uint32), s); - auto inverse = scatter( + const auto inverse = scatter( zeros({n}, uint32, s), *order, expand_dims(clamped, 1, s), 0, s); out.push_back(reshape(inverse, a.shape(), s)); } // If output is padded with fill value, counts is padded with zeros - // so that ``count.sum() == a.size()`` holds if output was not - // truncated. + // so that ``count.sum() == a.size()`` holds when output is not truncated. if (return_counts) { out.push_back(slice( scatter_add( diff --git a/python/tests/test_ops.py b/python/tests/test_ops.py index 1619b5301f..99fe037f99 100644 --- a/python/tests/test_ops.py +++ b/python/tests/test_ops.py @@ -3146,6 +3146,15 @@ def test_unique(self): counts = mx.unique(a, a.size, False, True)[1] self.assertEqual(mx.sum(counts > 0).item(), 3) + # counts.sum() == a.size if and only if the output was not truncated, + # and is smaller when it was + for size in (3, 4, 5, 7): + counts = mx.unique(a, size, False, True)[1] + self.assertEqual(mx.sum(counts).item(), a.size) + for size in (1, 2): + counts = mx.unique(a, size, False, True)[1] + self.assertLess(mx.sum(counts).item(), a.size) + # only the requested extras come back self.assertEqual(len(mx.unique(a, 3, True)), 2) self.assertEqual(len(mx.unique(a, 3, False, True)), 2) @@ -3181,6 +3190,14 @@ def test_unique(self): self.assertEqual(values.dtype, mx.float32) self.assertEqual(inverse.shape, (0,)) self.assertEqual(counts.shape, (0,)) + # the inverse is empty too, so the reconstruction stays valid + self.assertEqual(values[inverse].shape, (0,)) + empty_values, empty_inverse = mx.unique( + mx.array([], mx.float32), 3, True, fill_value=7 + ) + self.assertEqual(empty_values.shape, (3,)) + self.assertEqual(empty_inverse.shape, (0,)) + self.assertEqual(empty_values[empty_inverse].shape, (0,)) # empty input is handled correctly when given a fill_value self.assertTrue( @@ -3283,13 +3300,20 @@ def test_unique(self): mx.array_equal(mx.unique(a, 4, fill_value=fv), mx.array([1, 2, 3, 7])) ) - # inverse is clamped, so it points nowhere when size is zero + # inverse indices are clamped to size. If size is 0, this means an all zero array. values, inverse = mx.unique(a, 0, True) self.assertEqual(values.size, 0) self.assertTrue(mx.array_equal(inverse, mx.zeros(a.shape, mx.uint32))) + # the clamp has no index to clamp to, so this is the one case where + # values[inverse] is not a valid array + with self.assertRaises(ValueError): + mx.eval(values[inverse]) + for shape in ((2, 3), ()): + values, inverse = mx.unique(mx.zeros(shape, mx.int32), 0, True) + with self.assertRaises(ValueError): + mx.eval(values[inverse]) - # 8 byte types have to avoid a scatter on the data, since the GPU - # scatter does not support them + # check unique support 64 bit types for dtype in (mx.int64, mx.uint64, mx.complex64): values, inverse, counts = mx.unique( mx.array([3, 1, 2, 1], dtype), 4, True, True @@ -3298,11 +3322,21 @@ def test_unique(self): self.assertTrue(mx.array_equal(values, mx.array([1, 2, 3, 1], dtype))) self.assertTrue(mx.array_equal(inverse, mx.array([2, 0, 1, 0]))) self.assertTrue(mx.array_equal(counts, mx.array([2, 1, 1, 0]))) - # the dtype numpy hands over by default + + # in particular, the int64 dtype numpy hands over by default self.assertTrue( mx.array_equal(mx.unique(mx.array(np.arange(5)), 5), mx.arange(5)) ) + # bool only works on the cpu stream, since the metal sort has no bool + # kernel to build on + values, inverse, counts = mx.unique( + mx.array([True, False, True]), 2, True, True, stream=mx.cpu + ) + self.assertTrue(mx.array_equal(values, mx.array([False, True]))) + self.assertTrue(mx.array_equal(inverse, mx.array([1, 0, 1]))) + self.assertTrue(mx.array_equal(counts, mx.array([1, 2]))) + # truncation clamps the inverse so that it stays in range, otherwise # values[inverse] reads past the end of the output values, inverse = mx.unique(a, 2, True) @@ -3315,6 +3349,9 @@ def test_unique(self): mx.array([3.0, 1.0, 2.0, 1.0]) ) self.assertTrue(mx.array_equal(grad, mx.array([1.0, 1.0, 1.0, 0.0]))) + # the leading duplicate is the one credited, not a later one + grad = mx.grad(lambda x: mx.sum(mx.unique(x, 2)))(mx.array([1.0, 1.0, 2.0])) + self.assertTrue(mx.array_equal(grad, mx.array([1.0, 0.0, 1.0]))) # and a padded output routes the padding to that same element grad = mx.grad(lambda x: mx.sum(mx.unique(x, 4)))( mx.array([3.0, 1.0, 2.0, 1.0]) From cc17d3cc32347d5048c9f1f2f1be854807326789 Mon Sep 17 00:00:00 2001 From: Valentin Roussellet Date: Mon, 14 Sep 2026 09:36:25 -0700 Subject: [PATCH 6/8] test consolidation --- python/tests/test_ops.py | 335 +++++++++++++++++---------------------- 1 file changed, 142 insertions(+), 193 deletions(-) diff --git a/python/tests/test_ops.py b/python/tests/test_ops.py index 99fe037f99..c80539416a 100644 --- a/python/tests/test_ops.py +++ b/python/tests/test_ops.py @@ -3123,240 +3123,189 @@ def expect(out, want): mx.searchsorted(mx.array([1.0, 2.0]), mx.array([1.0]), side="middle") def test_unique(self): - # size is required - a = mx.array([2, 1, 2, 3, 1]) + def expect(out, want): + self.assertTrue(mx.array_equal(out, mx.array(want)), f"got {out}") + + a = mx.array([2, 1, 2, 3, 1]) # three unique values + + # size is required and cannot be negative with self.assertRaises(TypeError): mx.unique(a) + with self.assertRaises(ValueError): + mx.unique(a, -1) - # passing the correct size (number of unique values) produces a sorted array - self.assertTrue(mx.array_equal(mx.unique(a, 3), mx.array([1, 2, 3]))) - self.assertEqual(mx.unique(a, 3).dtype, mx.int32) - - # the flattened input size always holds every unique element - self.assertTrue(mx.array_equal(mx.unique(a, a.size), mx.array([1, 2, 3, 1, 1]))) + # size decides whether the output is exact, padded or truncated, and a + # truncated output keeps the smallest unique values + expect(mx.unique(a, 3), [1, 2, 3]) + expect(mx.unique(a, 5), [1, 2, 3, 1, 1]) + expect(mx.unique(a, 2), [1, 2]) + self.assertEqual(mx.unique(a, 3).dtype, a.dtype) + # only the requested extras come back, ordered values, inverse, counts + self.assertEqual(len(mx.unique(a, 3, True)), 2) + self.assertEqual(len(mx.unique(a, 3, False, True)), 2) values, inverse, counts = mx.unique(a, 4, True, True, fill_value=0) - self.assertTrue(mx.array_equal(values, mx.array([1, 2, 3, 0]))) - self.assertTrue(mx.array_equal(inverse, mx.array([1, 0, 1, 2, 0]))) - self.assertTrue(mx.array_equal(counts, mx.array([2, 2, 1, 0]))) + expect(values, [1, 2, 3, 0]) + expect(inverse, [1, 0, 1, 2, 0]) + expect(counts, [2, 2, 1, 0]) self.assertEqual(inverse.dtype, mx.uint32) self.assertEqual(counts.dtype, mx.int32) - # zero counts mark the padding, so they give the number of uniques - counts = mx.unique(a, a.size, False, True)[1] - self.assertEqual(mx.sum(counts > 0).item(), 3) + # the fill value defaults to the smallest unique element, and any one + # element array works whatever its rank + expect(mx.unique(a, 5, fill_value=None), [1, 2, 3, 1, 1]) + expect(mx.unique(a, 7, fill_value=-1), [1, 2, 3, -1, -1, -1, -1]) + for fv in (7, mx.array(7), mx.array([7]), mx.array([[7]])): + expect(mx.unique(a, 4, fill_value=fv), [1, 2, 3, 7]) + with self.assertRaises(ValueError): + mx.unique(a, 3, fill_value=mx.array([1, 2])) - # counts.sum() == a.size if and only if the output was not truncated, - # and is smaller when it was + # floats, including negatives and a repeated zero, and a fill value + # cast to the input dtype + f = mx.array([0.0, -1.5, 2.25, -1.5, 0.0], mx.float32) + expect(mx.unique(f, 3), [-1.5, 0.0, 2.25]) + self.assertEqual(mx.unique(f, 4, fill_value=1).dtype, mx.float32) + + # padding counts are zero, so counts sum to a.size unless the output + # was truncated, and their non zero entries count the unique values for size in (3, 4, 5, 7): - counts = mx.unique(a, size, False, True)[1] - self.assertEqual(mx.sum(counts).item(), a.size) + self.assertEqual(mx.sum(mx.unique(a, size, False, True)[1]).item(), a.size) for size in (1, 2): - counts = mx.unique(a, size, False, True)[1] - self.assertLess(mx.sum(counts).item(), a.size) - - # only the requested extras come back - self.assertEqual(len(mx.unique(a, 3, True)), 2) - self.assertEqual(len(mx.unique(a, 3, False, True)), 2) - - # a size larger than the input pads further - self.assertTrue( - mx.array_equal( - mx.unique(a, 7, fill_value=-1), - mx.array([1, 2, 3, -1, -1, -1, -1]), - ) - ) - - # a size smaller than the number of unique elements keeps the smallest - self.assertTrue(mx.array_equal(mx.unique(a, 2), mx.array([1, 2]))) - counts = mx.unique(a, 2, False, True)[1] - self.assertTrue(mx.array_equal(counts, mx.array([2, 2]))) + self.assertLess(mx.sum(mx.unique(a, size, False, True)[1]).item(), a.size) + self.assertEqual(mx.sum(mx.unique(a, a.size, False, True)[1] > 0).item(), 3) + expect(mx.unique(a, 2, False, True)[1], [2, 2]) - # the input is flattened, but the inverse keeps the input shape + # the inverse has the shape of the input and indexes into the output, + # so values[inverse] rebuilds the input when nothing was truncated b = mx.array([[3, 1, 3], [2, 1, 4]]) values, inverse = mx.unique(b, 4, True) - self.assertTrue(mx.array_equal(values, mx.array([1, 2, 3, 4]))) - self.assertEqual(inverse.shape, (2, 3)) - self.assertTrue(mx.array_equal(values[inverse], b)) + expect(values, [1, 2, 3, 4]) + self.assertEqual(inverse.shape, b.shape) + expect(values[inverse], [[3, 1, 3], [2, 1, 4]]) - # 0-d input + # truncation clamps the inverse to keep it in range, so a dropped value + # comes back as the largest kept one instead of reading past the output + values, inverse = mx.unique(a, 2, True) + self.assertTrue(bool(mx.all(inverse < values.size))) + expect(values[inverse], [2, 1, 2, 2, 1]) + + # a size of zero leaves no index to clamp to, the only case where the + # reconstruction is not a valid array + for x in (a, mx.zeros((2, 3), mx.int32), mx.array(5)): + values, inverse = mx.unique(x, 0, True) + self.assertEqual(values.size, 0) + self.assertTrue(mx.array_equal(inverse, mx.zeros(x.shape, mx.uint32))) + with self.assertRaises(ValueError): + mx.eval(values[inverse]) + + # a 0-d input is flattened like any other values, inverse = mx.unique(mx.array(5), 1, True) - self.assertTrue(mx.array_equal(values, mx.array([5]))) + expect(values, [5]) self.assertEqual(inverse.shape, ()) - # empty input with a zero size - values, inverse, counts = mx.unique(mx.array([], mx.float32), 0, True, True) - self.assertEqual(values.shape, (0,)) + # an empty input has no smallest element, so it needs a fill value + e = mx.array([], mx.float32) + values, inverse, counts = mx.unique(e, 0, True, True) self.assertEqual(values.dtype, mx.float32) - self.assertEqual(inverse.shape, (0,)) - self.assertEqual(counts.shape, (0,)) - # the inverse is empty too, so the reconstruction stays valid - self.assertEqual(values[inverse].shape, (0,)) - empty_values, empty_inverse = mx.unique( - mx.array([], mx.float32), 3, True, fill_value=7 - ) - self.assertEqual(empty_values.shape, (3,)) - self.assertEqual(empty_inverse.shape, (0,)) - self.assertEqual(empty_values[empty_inverse].shape, (0,)) - - # empty input is handled correctly when given a fill_value - self.assertTrue( - mx.array_equal( - mx.unique(mx.array([], mx.float32), 3, fill_value=7), - mx.array([7.0, 7.0, 7.0]), - ) - ) - - # empty input has no smallest element to default without a fill_value - # and produces an error + self.assertEqual((values.shape, inverse.shape, counts.shape), ((0,),) * 3) with self.assertRaises(ValueError): - mx.unique(mx.array([], mx.float32), 3) + mx.unique(e, 3) + values, inverse = mx.unique(e, 3, True, fill_value=7) + expect(values, [7.0, 7.0, 7.0]) + # its inverse is empty as well, so the reconstruction stays valid + self.assertEqual(values[inverse].shape, (0,)) # every element identical, and every element distinct values, counts = mx.unique(mx.full((7,), 4, mx.int32), 7, False, True) - self.assertTrue(mx.array_equal(values, mx.full((7,), 4, mx.int32))) - self.assertTrue(mx.array_equal(counts, mx.array([7, 0, 0, 0, 0, 0, 0]))) - self.assertTrue(mx.array_equal(mx.unique(mx.arange(6), 6), mx.arange(6))) - - # floats, including negatives and a repeated zero - f = mx.array([0.0, -1.5, 2.25, -1.5, 0.0], mx.float32) - self.assertTrue(mx.array_equal(mx.unique(f, 3), mx.array([-1.5, 0.0, 2.25]))) + expect(values, [4] * 7) + expect(counts, [7] + [0] * 6) + expect(mx.unique(mx.arange(6), 6), [0, 1, 2, 3, 4, 5]) - # None selects the default fill - self.assertTrue( - mx.array_equal(mx.unique(a, 5, fill_value=None), mx.unique(a, 5)) + # check 64 bit types work + for dtype in (mx.int64, mx.uint64, mx.complex64): + with self.subTest(dtype=dtype): + values, inverse, counts = mx.unique( + mx.array([3, 1, 2, 1], dtype), 4, True, True + ) + self.assertEqual(values.dtype, dtype) + self.assertTrue(mx.array_equal(values, mx.array([1, 2, 3, 1], dtype))) + expect(inverse, [2, 0, 1, 0]) + expect(counts, [2, 1, 1, 0]) + # in particular the int64 numpy hands over by default + expect(mx.unique(mx.array(np.arange(5)), 5), [0, 1, 2, 3, 4]) + + # bool only works on the cpu stream, the metal sort has no bool kernel + values, inverse, counts = mx.unique( + mx.array([True, False, True]), 2, True, True, stream=mx.cpu ) + expect(values, [False, True]) + expect(inverse, [1, 0, 1]) + expect(counts, [1, 2]) - # a fill value is cast to the input dtype - self.assertEqual(mx.unique(f, 4, fill_value=1).dtype, mx.float32) - - # like torch, NaN never equals itself so each one is kept + # NaN sorts last and never equals itself, so each one is kept, and none + # is ever picked as the default fill even though mx.min would return it nan = float("nan") values = mx.unique(mx.array([1.0, nan, 2.0, nan], mx.float32), 4) - self.assertTrue(mx.array_equal(values[:2], mx.array([1.0, 2.0]))) + expect(values[:2], [1.0, 2.0]) self.assertTrue(bool(mx.all(mx.isnan(values[2:])))) - - # NaN sorts last, so it is never the default fill - # note that mx.min(a) would return `nan` but we keep - # parity with jax which pads with the min of the sorted - # array in this case. values = mx.unique(mx.array([3.0, 1.0, nan, 2.0], mx.float32), 6) - self.assertTrue(mx.array_equal(values[:3], mx.array([1.0, 2.0, 3.0]))) + expect(values[:3], [1.0, 2.0, 3.0]) self.assertTrue(bool(mx.isnan(values[3]))) - self.assertTrue(mx.array_equal(values[4:], mx.array([1.0, 1.0]))) + expect(values[4:], [1.0, 1.0]) + + # the gradient credits the first occurrence of each unique value, and + # padding routes to that same element + for size, data, want in ( + (3, [3.0, 1.0, 2.0, 1.0], [1.0, 1.0, 1.0, 0.0]), + (4, [3.0, 1.0, 2.0, 1.0], [1.0, 2.0, 1.0, 0.0]), + (2, [1.0, 1.0, 2.0], [1.0, 0.0, 1.0]), + ): + with self.subTest(size=size, data=data): + grad = mx.grad(lambda x: mx.sum(mx.unique(x, size)))(mx.array(data)) + expect(grad, want) + + # the output shape does not depend on the values, so this composes with + # the function transforms + self.assertTrue( + mx.array_equal(mx.compile(lambda x: mx.unique(x, 3))(a), mx.unique(a, 3)) + ) + batched = mx.vmap(lambda x: mx.unique(x, 3))(mx.stack([a, a[::-1]])) + self.assertTrue(mx.array_equal(batched, mx.stack([mx.unique(a, 3)] * 2))) - # Compare the output with np.unique + # compare against numpy on random inputs rng = np.random.RandomState(0) for n in (1, 2, 17, 1000, 5000): - a_np = rng.randint(-20, 20, size=n).astype(np.int32) - a_mx = mx.array(a_np) - v_np, i_np, c_np = np.unique(a_np, return_inverse=True, return_counts=True) - values, inverse, counts = mx.unique(a_mx, len(v_np), True, True) - self.assertTrue(np.array_equal(np.array(values), v_np)) - self.assertTrue( - np.array_equal(np.array(inverse).reshape(-1), i_np.reshape(-1)) - ) - self.assertTrue(np.array_equal(np.array(counts), c_np)) - # the inverse rebuilds the input - self.assertTrue(mx.array_equal(values[inverse], a_mx)) - - # a padded size holds the same elements up front - padded, pad_counts = mx.unique(a_mx, n, False, True) - self.assertEqual(padded.shape, (n,)) - self.assertTrue(np.array_equal(np.array(padded)[: len(v_np)], v_np)) - self.assertEqual(mx.sum(pad_counts > 0).item(), len(v_np)) - - # test various non-contiguous slices + with self.subTest(n=n): + a_np = rng.randint(-20, 20, size=n).astype(np.int32) + a_mx = mx.array(a_np) + v_np, i_np, c_np = np.unique( + a_np, return_inverse=True, return_counts=True + ) + values, inverse, counts = mx.unique(a_mx, len(v_np), True, True) + self.assertTrue(np.array_equal(np.array(values), v_np)) + self.assertTrue( + np.array_equal(np.array(inverse).reshape(-1), i_np.reshape(-1)) + ) + self.assertTrue(np.array_equal(np.array(counts), c_np)) + self.assertTrue(mx.array_equal(values[inverse], a_mx)) + # a padded output holds the same values up front + padded, pad_counts = mx.unique(a_mx, n, False, True) + self.assertEqual(padded.shape, (n,)) + self.assertTrue(np.array_equal(np.array(padded)[: len(v_np)], v_np)) + self.assertEqual(mx.sum(pad_counts > 0).item(), len(v_np)) + + # non contiguous inputs are read in the output order base_np = rng.randint(0, 5, size=(4, 6)).astype(np.int32) base = mx.array(base_np) - for v_mx, v_np in [ + for v_mx, v_np in ( (base.T, base_np.T), (base[::2], base_np[::2]), (base[::-1], base_np[::-1]), (base[:, ::-1], base_np[:, ::-1]), (mx.broadcast_to(base[0], (3, 6)), np.broadcast_to(base_np[0], (3, 6))), - ]: - expected = np.unique(v_np) - self.assertTrue( - np.array_equal(np.array(mx.unique(v_mx, len(expected))), expected) - ) - - # the output size is static, so this composes with the transforms - c = mx.array([2, 1, 2, 3, 1]) - self.assertTrue( - mx.array_equal(mx.compile(lambda x: mx.unique(x, 3))(c), mx.unique(c, 3)) - ) - batched = mx.vmap(lambda x: mx.unique(x, 3))(mx.stack([c, c[::-1]])) - self.assertTrue(mx.array_equal(batched, mx.stack([mx.unique(c, 3)] * 2))) - - # Test input validation - with self.assertRaises(ValueError): - mx.unique(a, -1) - with self.assertRaises(ValueError): - mx.unique(a, 3, fill_value=mx.array([1, 2])) - # any one element array fill value works, whatever its rank - for fv in (7, mx.array(7), mx.array([7]), mx.array([[7]])): - self.assertTrue( - mx.array_equal(mx.unique(a, 4, fill_value=fv), mx.array([1, 2, 3, 7])) - ) - - # inverse indices are clamped to size. If size is 0, this means an all zero array. - values, inverse = mx.unique(a, 0, True) - self.assertEqual(values.size, 0) - self.assertTrue(mx.array_equal(inverse, mx.zeros(a.shape, mx.uint32))) - # the clamp has no index to clamp to, so this is the one case where - # values[inverse] is not a valid array - with self.assertRaises(ValueError): - mx.eval(values[inverse]) - for shape in ((2, 3), ()): - values, inverse = mx.unique(mx.zeros(shape, mx.int32), 0, True) - with self.assertRaises(ValueError): - mx.eval(values[inverse]) - - # check unique support 64 bit types - for dtype in (mx.int64, mx.uint64, mx.complex64): - values, inverse, counts = mx.unique( - mx.array([3, 1, 2, 1], dtype), 4, True, True - ) - self.assertEqual(values.dtype, dtype) - self.assertTrue(mx.array_equal(values, mx.array([1, 2, 3, 1], dtype))) - self.assertTrue(mx.array_equal(inverse, mx.array([2, 0, 1, 0]))) - self.assertTrue(mx.array_equal(counts, mx.array([2, 1, 1, 0]))) - - # in particular, the int64 dtype numpy hands over by default - self.assertTrue( - mx.array_equal(mx.unique(mx.array(np.arange(5)), 5), mx.arange(5)) - ) - - # bool only works on the cpu stream, since the metal sort has no bool - # kernel to build on - values, inverse, counts = mx.unique( - mx.array([True, False, True]), 2, True, True, stream=mx.cpu - ) - self.assertTrue(mx.array_equal(values, mx.array([False, True]))) - self.assertTrue(mx.array_equal(inverse, mx.array([1, 0, 1]))) - self.assertTrue(mx.array_equal(counts, mx.array([1, 2]))) - - # truncation clamps the inverse so that it stays in range, otherwise - # values[inverse] reads past the end of the output - values, inverse = mx.unique(a, 2, True) - self.assertTrue(mx.array_equal(values, mx.array([1, 2]))) - self.assertTrue(bool(mx.all(inverse < values.size))) - self.assertTrue(mx.array_equal(values[inverse], mx.array([2, 1, 2, 2, 1]))) - - # the gradient credits the first occurrence of each unique value - grad = mx.grad(lambda x: mx.sum(mx.unique(x, 3)))( - mx.array([3.0, 1.0, 2.0, 1.0]) - ) - self.assertTrue(mx.array_equal(grad, mx.array([1.0, 1.0, 1.0, 0.0]))) - # the leading duplicate is the one credited, not a later one - grad = mx.grad(lambda x: mx.sum(mx.unique(x, 2)))(mx.array([1.0, 1.0, 2.0])) - self.assertTrue(mx.array_equal(grad, mx.array([1.0, 0.0, 1.0]))) - # and a padded output routes the padding to that same element - grad = mx.grad(lambda x: mx.sum(mx.unique(x, 4)))( - mx.array([3.0, 1.0, 2.0, 1.0]) - ) - self.assertTrue(mx.array_equal(grad, mx.array([1.0, 2.0, 1.0, 0.0]))) + ): + want = np.unique(v_np) + self.assertTrue(np.array_equal(np.array(mx.unique(v_mx, len(want))), want)) @unittest.skipIf( os.getenv("LOW_MEMORY", None) is not None, From 2c70d1b11a222e83cba7c4ce12cd9e0edd3c7ecd Mon Sep 17 00:00:00 2001 From: Valentin Roussellet Date: Mon, 14 Sep 2026 10:38:46 -0700 Subject: [PATCH 7/8] some fix to tests --- mlx/ops.cpp | 12 ++++++------ mlx/ops.h | 5 +++-- python/src/ops.cpp | 6 ++++-- python/tests/test_export_import.py | 8 +++++--- python/tests/test_ops.py | 29 +++++++++++++++++++++-------- 5 files changed, 39 insertions(+), 21 deletions(-) diff --git a/mlx/ops.cpp b/mlx/ops.cpp index 2969bb0bd2..c15dbe7cf5 100644 --- a/mlx/ops.cpp +++ b/mlx/ops.cpp @@ -2967,9 +2967,8 @@ std::vector unique( const int n = flat.size(); // Handle the edge case of an empty array. - // The output is to be filled with `fill_value`. if (n == 0) { - // Throw as is no smallest element to take the default fill value from. + // Without an element there is no default fill value. if (size > 0 && !fill_value) { throw std::invalid_argument( "[unique] A fill value is required for an empty input with a" @@ -2992,7 +2991,7 @@ std::vector unique( return out; } - // Sort the array. Not differentiable. + // Sort the array. The values stay differentiable, the permutation does not. std::optional order; if (return_inverse) { order = stop_gradient(argsort(flat, 0, s), s); @@ -3008,7 +3007,8 @@ std::vector unique( 0, s); - // Cumsum on boundary gets the index of unique elements. Not differentiable. + // Cumsum on boundary gives each sorted element the index of the unique + // value it belongs to. It indexes the output, so it is not differentiable. const auto group = stop_gradient( subtract( cumsum(astype(boundary, uint32, s), 0, false, true, s), @@ -3059,8 +3059,8 @@ std::vector unique( zeros({n}, uint32, s), *order, expand_dims(clamped, 1, s), 0, s); out.push_back(reshape(inverse, a.shape(), s)); } - // If output is padded with fill value, counts is padded with zeros - // so that ``count.sum() == a.size()`` holds when output is not truncated. + // Padding entries count zero, so counts sum to the input size unless the + // output was truncated. if (return_counts) { out.push_back(slice( scatter_add( diff --git a/mlx/ops.h b/mlx/ops.h index de4773222b..6a08dd4d21 100644 --- a/mlx/ops.h +++ b/mlx/ops.h @@ -883,8 +883,9 @@ MLX_API array topk(const array& a, int k, int axis, StreamOrDevice s = {}); * the inverse indices and the counts. The output has the given size, and * is truncated if ``size`` is smaller than the number of unique values. * If ``size`` is larger, the result is padded with ``fill_value``, or the - * smallest unique value if no fill value is provided. An empty input has no - * such value, so it throws unless ``size`` is zero or a fill value is given. + * first of the sorted unique values if no fill value is provided. An empty + * input has none, so it throws unless ``size`` is zero or a fill value is + * given. * A truncated output loses the counts of the values it drops, and its inverse * indices are clamped to the last entry. */ diff --git a/python/src/ops.cpp b/python/src/ops.cpp index 7ed3fc5238..cb2c6384c0 100644 --- a/python/src/ops.cpp +++ b/python/src/ops.cpp @@ -3286,6 +3286,9 @@ void init_ops(nb::module_& m) { ``inverse`` is meaningless for a ``size`` of ``0``, since there is no element left for it to point at. + This op builds on ``mx.sort``, so a ``bool`` input needs the CPU + stream, which is where ``mx.sort`` supports it. + Args: a (array): Input array. size (int): The size of the output. If the size is smaller than @@ -3298,8 +3301,7 @@ void init_ops(nb::module_& m) { of times each unique element occurs in ``a``. Default: ``False``. fill_value (scalar or array, optional): The value of the entries past the last unique element. If ``None``, this defaults to the - minimum unique value. - Default: ``None``. + first of the sorted unique elements. Default: ``None``. Returns: array or tuple(array, ...): The sorted unique elements. If diff --git a/python/tests/test_export_import.py b/python/tests/test_export_import.py index 7e0be7249d..21304095bd 100644 --- a/python/tests/test_export_import.py +++ b/python/tests/test_export_import.py @@ -335,13 +335,15 @@ def fun(a, v): def test_export_unique(self): path = os.path.join(self.test_dir, "fn.mlxfn") - # the fixed output size is what makes this exportable - for size, fill_value in ((10, None), (4, -1)): + # the fixed output size is what makes this exportable. The input is + # fixed so that the padded, exact and truncated cases are all covered + # rather than left to chance. + x = mx.array([3, 1, 2, 1, 3, 2, 1, 0, 2, 1]) # four unique values + for size, fill_value in ((10, None), (6, -1), (4, None), (2, None)): def fun(a): return mx.unique(a, size, True, True, fill_value=fill_value) - x = mx.random.randint(0, 5, shape=(10,)) mx.export_function(path, fun, x) imported_fun = mx.import_function(path) expected = fun(x) diff --git a/python/tests/test_ops.py b/python/tests/test_ops.py index c80539416a..3866b8f21a 100644 --- a/python/tests/test_ops.py +++ b/python/tests/test_ops.py @@ -3203,17 +3203,20 @@ def expect(out, want): expect(values, [5]) self.assertEqual(inverse.shape, ()) - # an empty input has no smallest element, so it needs a fill value + # an empty input needs a fill value if size > 0 e = mx.array([], mx.float32) values, inverse, counts = mx.unique(e, 0, True, True) self.assertEqual(values.dtype, mx.float32) self.assertEqual((values.shape, inverse.shape, counts.shape), ((0,),) * 3) with self.assertRaises(ValueError): mx.unique(e, 3) - values, inverse = mx.unique(e, 3, True, fill_value=7) - expect(values, [7.0, 7.0, 7.0]) - # its inverse is empty as well, so the reconstruction stays valid - self.assertEqual(values[inverse].shape, (0,)) + # the reconstruction stays valid even with size 0 + for size, want in ((0, []), (3, [7.0, 7.0, 7.0])): + values, inverse = mx.unique(e, size, True, fill_value=7) + expect(values, want) + rebuilt = values[inverse] + mx.eval(rebuilt) + self.assertEqual(rebuilt.shape, (0,)) # every element identical, and every element distinct values, counts = mx.unique(mx.full((7,), 4, mx.int32), 7, False, True) @@ -3252,6 +3255,10 @@ def expect(out, want): expect(values[:3], [1.0, 2.0, 3.0]) self.assertTrue(bool(mx.isnan(values[3]))) expect(values[4:], [1.0, 1.0]) + # with nothing but NaN there is no other element to fall back on, so + # the default fill is NaN as well + values = mx.unique(mx.array([nan, nan], mx.float32), 5) + self.assertTrue(bool(mx.all(mx.isnan(values)))) # the gradient credits the first occurrence of each unique value, and # padding routes to that same element @@ -3269,8 +3276,13 @@ def expect(out, want): self.assertTrue( mx.array_equal(mx.compile(lambda x: mx.unique(x, 3))(a), mx.unique(a, 3)) ) - batched = mx.vmap(lambda x: mx.unique(x, 3))(mx.stack([a, a[::-1]])) - self.assertTrue(mx.array_equal(batched, mx.stack([mx.unique(a, 3)] * 2))) + # the batch rows must have different unique sets, otherwise broadcasting + # one row over the batch would pass + rows = mx.array([[2, 1, 2, 3, 1], [7, 7, 5, 5, 9]]) + values, inverse, counts = mx.vmap(lambda x: mx.unique(x, 3, True, True))(rows) + expect(values, [[1, 2, 3], [5, 7, 9]]) + expect(inverse, [[1, 0, 1, 2, 0], [1, 1, 0, 0, 2]]) + expect(counts, [[2, 2, 1], [2, 2, 1]]) # compare against numpy on random inputs rng = np.random.RandomState(0) @@ -3294,7 +3306,8 @@ def expect(out, want): self.assertTrue(np.array_equal(np.array(padded)[: len(v_np)], v_np)) self.assertEqual(mx.sum(pad_counts > 0).item(), len(v_np)) - # non contiguous inputs are read in the output order + # non contiguous inputs are flattened over their logical elements, + # broadcast multiplicity included, not over their backing memory base_np = rng.randint(0, 5, size=(4, 6)).astype(np.int32) base = mx.array(base_np) for v_mx, v_np in ( From 0b46283f6e09c95ca556b2076b4ec532adbc0062 Mon Sep 17 00:00:00 2001 From: Valentin Roussellet Date: Fri, 25 Sep 2026 12:48:06 -0700 Subject: [PATCH 8/8] add indices --- mlx/ops.cpp | 12 +++- mlx/ops.h | 4 +- python/src/ops.cpp | 36 +++++++--- python/tests/test_export_import.py | 9 ++- python/tests/test_ops.py | 105 +++++++++++++++++++++++------ 5 files changed, 132 insertions(+), 34 deletions(-) diff --git a/mlx/ops.cpp b/mlx/ops.cpp index d38dcf1e3f..b31de2b7c0 100644 --- a/mlx/ops.cpp +++ b/mlx/ops.cpp @@ -2947,6 +2947,7 @@ array topk(const array& a, int k, int axis, StreamOrDevice s /* = {}*/) { std::vector unique( const array& a, int size, + bool return_index /* = false */, bool return_inverse /* = false */, bool return_counts /* = false */, const std::optional& fill_value /* = std::nullopt */, @@ -2982,6 +2983,9 @@ std::vector unique( astype(reshape(*fill_value, {}, s), flat.dtype(), s), flat.dtype(), s)); + if (return_index) { + out.push_back(zeros({size}, uint32, s)); + } if (return_inverse) { out.push_back(zeros(a.shape(), uint32, s)); } @@ -2993,7 +2997,7 @@ std::vector unique( // Sort the array. The values stay differentiable, the permutation does not. std::optional order; - if (return_inverse) { + if (return_index || return_inverse) { order = stop_gradient(argsort(flat, 0, s), s); } const auto sorted = order ? take(flat, *order, 0, s) : sort(flat, 0, s); @@ -3050,6 +3054,12 @@ std::vector unique( // Build the output arrays. std::vector out; out.push_back(where(used, take(sorted, positions, 0, s), fill, s)); + if (return_index) { + // The sort is stable, so the first sorted position of a group holds the + // first occurrence in the input. Unused slots point at position zero, + // which pads with the index of the smallest unique value. + out.push_back(take(*order, positions, 0, s)); + } if (return_inverse) { // Clamp so that the indices stay inside a truncated output. Without a // truncation every group index is already smaller than size. diff --git a/mlx/ops.h b/mlx/ops.h index 6a08dd4d21..7d38416ae1 100644 --- a/mlx/ops.h +++ b/mlx/ops.h @@ -880,7 +880,8 @@ MLX_API array topk(const array& a, int k, int axis, StreamOrDevice s = {}); /** * Returns the sorted unique elements of the flattened array, and optionally - * the inverse indices and the counts. The output has the given size, and + * the index of the first occurrence of each, the inverse indices and the + * counts. The output has the given size, and * is truncated if ``size`` is smaller than the number of unique values. * If ``size`` is larger, the result is padded with ``fill_value``, or the * first of the sorted unique values if no fill value is provided. An empty @@ -892,6 +893,7 @@ MLX_API array topk(const array& a, int k, int axis, StreamOrDevice s = {}); MLX_API std::vector unique( const array& a, int size, + bool return_index = false, bool return_inverse = false, bool return_counts = false, const std::optional& fill_value = std::nullopt, diff --git a/python/src/ops.cpp b/python/src/ops.cpp index 315ac858fc..b0561fa698 100644 --- a/python/src/ops.cpp +++ b/python/src/ops.cpp @@ -3240,6 +3240,7 @@ void init_ops(nb::module_& m) { "unique", [](const mx::array& a, int size, + bool return_index, bool return_inverse, bool return_counts, const std::optional& fill_value, @@ -3248,8 +3249,14 @@ void init_ops(nb::module_& m) { if (fill_value) { fill_value_ = to_array(fill_value.value(), a.dtype()); } - auto out = - mx::unique(a, size, return_inverse, return_counts, fill_value_, s); + auto out = mx::unique( + a, + size, + return_index, + return_inverse, + return_counts, + fill_value_, + s); if (out.size() == 1) { return nb::cast(out.at(0)); } @@ -3261,13 +3268,14 @@ void init_ops(nb::module_& m) { }, nb::arg(), "size"_a, + "return_index"_a = false, "return_inverse"_a = false, "return_counts"_a = false, "fill_value"_a = nb::none(), nb::kw_only(), "stream"_a = nb::none(), nb::sig( - "def unique(a: array, /, size: int, return_inverse: bool = False, return_counts: bool = False, fill_value: scalar | array | None = None, *, stream: StreamOrDevice = None) -> array | tuple[array, ...]"), + "def unique(a: array, /, size: int, return_index: bool = False, return_inverse: bool = False, return_counts: bool = False, fill_value: scalar | array | None = None, *, stream: StreamOrDevice = None) -> array | tuple[array, ...]"), R"pbdoc( Returns the sorted unique elements of the flattened array. @@ -3280,9 +3288,10 @@ void init_ops(nb::module_& m) { elements as long as the output is not truncated. A truncated output keeps the smallest ``size`` unique elements, and - ``inverse`` and ``counts`` then stop describing all of ``a`` and stop - agreeing with each other: the counts of the dropped elements are gone, - while their indices in ``inverse`` are clamped to the last entry. + ``index``, ``inverse`` and ``counts`` then stop describing all of ``a``, + and ``inverse`` and ``counts`` stop agreeing with each other: the counts + of the dropped elements are gone, while their indices in ``inverse`` are + clamped to the last entry. ``inverse`` is meaningless for a ``size`` of ``0``, since there is no element left for it to point at. @@ -3294,6 +3303,9 @@ void init_ops(nb::module_& m) { size (int): The size of the output. If the size is smaller than the number of unique elements of ``a``, the output is truncated. If it is larger, the output is padded with ``fill_value``. + return_index (bool, optional): If ``True``, also return the index + in the flattened ``a`` of the first occurrence of each unique + element. Default: ``False``. return_inverse (bool, optional): If ``True``, also return the indices of the unique array that rebuild ``a``. The indices have the same shape as ``a``. Default: ``False``. @@ -3304,9 +3316,10 @@ void init_ops(nb::module_& m) { first of the sorted unique elements. Default: ``None``. Returns: - array or tuple(array, ...): The sorted unique elements. If - ``return_inverse`` or ``return_counts`` is ``True``, a tuple with - the requested arrays in the order values, inverse, counts. + array or tuple(array, ...): The sorted unique elements. If any of + ``return_index``, ``return_inverse`` or ``return_counts`` is + ``True``, a tuple with the requested arrays in the order values, + index, inverse, counts. Example: >>> a = mx.array([2, 1, 2, 3, 1]) @@ -3314,9 +3327,12 @@ void init_ops(nb::module_& m) { array([1, 2, 3], dtype=int32) >>> mx.unique(a, 5) array([1, 2, 3, 1, 1], dtype=int32) - >>> values, inverse, counts = mx.unique(a, 4, True, True, fill_value=0) + >>> values, index, inverse, counts = mx.unique( + ... a, 4, True, True, True, fill_value=0) >>> values array([1, 2, 3, 0], dtype=int32) + >>> index + array([1, 0, 3, 1], dtype=uint32) >>> inverse array([1, 0, 1, 2, 0], dtype=uint32) >>> counts diff --git a/python/tests/test_export_import.py b/python/tests/test_export_import.py index 21304095bd..a718dc9d79 100644 --- a/python/tests/test_export_import.py +++ b/python/tests/test_export_import.py @@ -342,7 +342,14 @@ def test_export_unique(self): for size, fill_value in ((10, None), (6, -1), (4, None), (2, None)): def fun(a): - return mx.unique(a, size, True, True, fill_value=fill_value) + return mx.unique( + a, + size, + return_index=True, + return_inverse=True, + return_counts=True, + fill_value=fill_value, + ) mx.export_function(path, fun, x) imported_fun = mx.import_function(path) diff --git a/python/tests/test_ops.py b/python/tests/test_ops.py index 3866b8f21a..5b4c0ed83c 100644 --- a/python/tests/test_ops.py +++ b/python/tests/test_ops.py @@ -3127,6 +3127,7 @@ def expect(out, want): self.assertTrue(mx.array_equal(out, mx.array(want)), f"got {out}") a = mx.array([2, 1, 2, 3, 1]) # three unique values + b = mx.array([[3, 1, 3], [2, 1, 4]]) # four unique values when flat # size is required and cannot be negative with self.assertRaises(TypeError): @@ -3142,15 +3143,40 @@ def expect(out, want): self.assertEqual(mx.unique(a, 3).dtype, a.dtype) # only the requested extras come back, ordered values, inverse, counts - self.assertEqual(len(mx.unique(a, 3, True)), 2) - self.assertEqual(len(mx.unique(a, 3, False, True)), 2) - values, inverse, counts = mx.unique(a, 4, True, True, fill_value=0) + self.assertEqual(len(mx.unique(a, 3, return_inverse=True)), 2) + self.assertEqual(len(mx.unique(a, 3, return_counts=True)), 2) + values, inverse, counts = mx.unique( + a, 4, return_inverse=True, return_counts=True, fill_value=0 + ) expect(values, [1, 2, 3, 0]) expect(inverse, [1, 0, 1, 2, 0]) expect(counts, [2, 2, 1, 0]) self.assertEqual(inverse.dtype, mx.uint32) self.assertEqual(counts.dtype, mx.int32) + # return_index gives the first occurrence of each unique element in the + # flattened input, and pads with the index of the smallest one + expect(mx.unique(a, 3, return_index=True)[1], [1, 0, 3]) + expect(mx.unique(a, 5, return_index=True)[1], [1, 0, 3, 1, 1]) + expect(mx.unique(a, 2, return_index=True)[1], [1, 0]) + # the fill value applies to the values, never to the indices + expect(mx.unique(a, 5, return_index=True, fill_value=99)[1], [1, 0, 3, 1, 1]) + # the sort is stable, so a leading duplicate wins + expect(mx.unique(mx.array([5, 1, 5, 1]), 4, return_index=True)[1], [1, 0, 1, 1]) + # the index refers to the flattened input + expect(mx.unique(b, 4, return_index=True)[1], [1, 3, 0, 5]) + + # all four outputs at once, in the order values, index, inverse, counts + values, index, inverse, counts = mx.unique( + a, 4, return_index=True, return_inverse=True, return_counts=True + ) + expect(values, [1, 2, 3, 1]) + expect(index, [1, 0, 3, 1]) + expect(inverse, [1, 0, 1, 2, 0]) + expect(counts, [2, 2, 1, 0]) + self.assertEqual(index.dtype, mx.uint32) + self.assertEqual(len(mx.unique(a, 3, return_index=True)), 2) + # the fill value defaults to the smallest unique element, and any one # element array works whatever its rank expect(mx.unique(a, 5, fill_value=None), [1, 2, 3, 1, 1]) @@ -3169,57 +3195,64 @@ def expect(out, want): # padding counts are zero, so counts sum to a.size unless the output # was truncated, and their non zero entries count the unique values for size in (3, 4, 5, 7): - self.assertEqual(mx.sum(mx.unique(a, size, False, True)[1]).item(), a.size) + self.assertEqual( + mx.sum(mx.unique(a, size, return_counts=True)[1]).item(), a.size + ) for size in (1, 2): - self.assertLess(mx.sum(mx.unique(a, size, False, True)[1]).item(), a.size) - self.assertEqual(mx.sum(mx.unique(a, a.size, False, True)[1] > 0).item(), 3) - expect(mx.unique(a, 2, False, True)[1], [2, 2]) + self.assertLess( + mx.sum(mx.unique(a, size, return_counts=True)[1]).item(), a.size + ) + self.assertEqual( + mx.sum(mx.unique(a, a.size, return_counts=True)[1] > 0).item(), 3 + ) + expect(mx.unique(a, 2, return_counts=True)[1], [2, 2]) # the inverse has the shape of the input and indexes into the output, # so values[inverse] rebuilds the input when nothing was truncated - b = mx.array([[3, 1, 3], [2, 1, 4]]) - values, inverse = mx.unique(b, 4, True) + values, inverse = mx.unique(b, 4, return_inverse=True) expect(values, [1, 2, 3, 4]) self.assertEqual(inverse.shape, b.shape) expect(values[inverse], [[3, 1, 3], [2, 1, 4]]) # truncation clamps the inverse to keep it in range, so a dropped value # comes back as the largest kept one instead of reading past the output - values, inverse = mx.unique(a, 2, True) + values, inverse = mx.unique(a, 2, return_inverse=True) self.assertTrue(bool(mx.all(inverse < values.size))) expect(values[inverse], [2, 1, 2, 2, 1]) # a size of zero leaves no index to clamp to, the only case where the # reconstruction is not a valid array for x in (a, mx.zeros((2, 3), mx.int32), mx.array(5)): - values, inverse = mx.unique(x, 0, True) + values, inverse = mx.unique(x, 0, return_inverse=True) self.assertEqual(values.size, 0) self.assertTrue(mx.array_equal(inverse, mx.zeros(x.shape, mx.uint32))) with self.assertRaises(ValueError): mx.eval(values[inverse]) # a 0-d input is flattened like any other - values, inverse = mx.unique(mx.array(5), 1, True) + values, inverse = mx.unique(mx.array(5), 1, return_inverse=True) expect(values, [5]) self.assertEqual(inverse.shape, ()) # an empty input needs a fill value if size > 0 e = mx.array([], mx.float32) - values, inverse, counts = mx.unique(e, 0, True, True) + values, inverse, counts = mx.unique( + e, 0, return_inverse=True, return_counts=True + ) self.assertEqual(values.dtype, mx.float32) self.assertEqual((values.shape, inverse.shape, counts.shape), ((0,),) * 3) with self.assertRaises(ValueError): mx.unique(e, 3) # the reconstruction stays valid even with size 0 for size, want in ((0, []), (3, [7.0, 7.0, 7.0])): - values, inverse = mx.unique(e, size, True, fill_value=7) + values, inverse = mx.unique(e, size, return_inverse=True, fill_value=7) expect(values, want) rebuilt = values[inverse] mx.eval(rebuilt) self.assertEqual(rebuilt.shape, (0,)) # every element identical, and every element distinct - values, counts = mx.unique(mx.full((7,), 4, mx.int32), 7, False, True) + values, counts = mx.unique(mx.full((7,), 4, mx.int32), 7, return_counts=True) expect(values, [4] * 7) expect(counts, [7] + [0] * 6) expect(mx.unique(mx.arange(6), 6), [0, 1, 2, 3, 4, 5]) @@ -3228,7 +3261,10 @@ def expect(out, want): for dtype in (mx.int64, mx.uint64, mx.complex64): with self.subTest(dtype=dtype): values, inverse, counts = mx.unique( - mx.array([3, 1, 2, 1], dtype), 4, True, True + mx.array([3, 1, 2, 1], dtype), + 4, + return_inverse=True, + return_counts=True, ) self.assertEqual(values.dtype, dtype) self.assertTrue(mx.array_equal(values, mx.array([1, 2, 3, 1], dtype))) @@ -3239,7 +3275,11 @@ def expect(out, want): # bool only works on the cpu stream, the metal sort has no bool kernel values, inverse, counts = mx.unique( - mx.array([True, False, True]), 2, True, True, stream=mx.cpu + mx.array([True, False, True]), + 2, + return_inverse=True, + return_counts=True, + stream=mx.cpu, ) expect(values, [False, True]) expect(inverse, [1, 0, 1]) @@ -3276,10 +3316,27 @@ def expect(out, want): self.assertTrue( mx.array_equal(mx.compile(lambda x: mx.unique(x, 3))(a), mx.unique(a, 3)) ) + + # every optional output and a fill value have to survive compilation too + def all_outputs(x): + return mx.unique( + x, + 5, + return_index=True, + return_inverse=True, + return_counts=True, + fill_value=-1, + ) + + for got, want in zip(mx.compile(all_outputs)(a), all_outputs(a)): + self.assertTrue(mx.array_equal(got, want)) + self.assertEqual(got.dtype, want.dtype) # the batch rows must have different unique sets, otherwise broadcasting # one row over the batch would pass rows = mx.array([[2, 1, 2, 3, 1], [7, 7, 5, 5, 9]]) - values, inverse, counts = mx.vmap(lambda x: mx.unique(x, 3, True, True))(rows) + values, inverse, counts = mx.vmap( + lambda x: mx.unique(x, 3, return_inverse=True, return_counts=True) + )(rows) expect(values, [[1, 2, 3], [5, 7, 9]]) expect(inverse, [[1, 0, 1, 2, 0], [1, 1, 0, 0, 2]]) expect(counts, [[2, 2, 1], [2, 2, 1]]) @@ -3293,7 +3350,9 @@ def expect(out, want): v_np, i_np, c_np = np.unique( a_np, return_inverse=True, return_counts=True ) - values, inverse, counts = mx.unique(a_mx, len(v_np), True, True) + values, inverse, counts = mx.unique( + a_mx, len(v_np), return_inverse=True, return_counts=True + ) self.assertTrue(np.array_equal(np.array(values), v_np)) self.assertTrue( np.array_equal(np.array(inverse).reshape(-1), i_np.reshape(-1)) @@ -3301,7 +3360,7 @@ def expect(out, want): self.assertTrue(np.array_equal(np.array(counts), c_np)) self.assertTrue(mx.array_equal(values[inverse], a_mx)) # a padded output holds the same values up front - padded, pad_counts = mx.unique(a_mx, n, False, True) + padded, pad_counts = mx.unique(a_mx, n, return_counts=True) self.assertEqual(padded.shape, (n,)) self.assertTrue(np.array_equal(np.array(padded)[: len(v_np)], v_np)) self.assertEqual(mx.sum(pad_counts > 0).item(), len(v_np)) @@ -3318,7 +3377,11 @@ def expect(out, want): (mx.broadcast_to(base[0], (3, 6)), np.broadcast_to(base_np[0], (3, 6))), ): want = np.unique(v_np) - self.assertTrue(np.array_equal(np.array(mx.unique(v_mx, len(want))), want)) + values, inverse = mx.unique(v_mx, len(want), return_inverse=True) + self.assertTrue(np.array_equal(np.array(values), want)) + # the inverse follows the shape of the view, not of its backing array + self.assertEqual(inverse.shape, v_mx.shape) + self.assertTrue(np.array_equal(np.array(values[inverse]), v_np)) @unittest.skipIf( os.getenv("LOW_MEMORY", None) is not None,