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 83094d25b8..1e60d7a2ad 100644 --- a/mlx/ops.cpp +++ b/mlx/ops.cpp @@ -2955,6 +2955,148 @@ 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_index /* = false */, + bool return_inverse /* = false */, + 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 << "."; + throw std::invalid_argument(msg.str()); + } + if (fill_value && fill_value->size() != 1) { + std::ostringstream msg; + msg << "[unique] Fill value must have one element, but got shape " + << fill_value->shape() << "."; + throw std::invalid_argument(msg.str()); + } + const auto flat = flatten(a, s); + const int n = flat.size(); + + // Handle the edge case of an empty array. + if (n == 0) { + // 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" + " non-zero size."); + } + std::vector out; + out.push_back( + size == 0 ? flat + : full( + {size}, + 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)); + } + if (return_counts) { + out.push_back(zeros({size}, int32, s)); + } + return out; + } + + // Sort the array. The values stay differentiable, the permutation does not. + std::optional order; + 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); + + // 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( + slice(sorted, {1}, {n}, s), slice(sorted, {0}, {n - 1}, s), s)}, + 0, + s); + + // 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), + array(1, uint32), + s), + s); + + // 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); + + 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( + 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(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. + const auto clamped = + minimum(group, array(std::max(size - 1, 0), uint32), s); + const auto inverse = scatter( + zeros({n}, uint32, s), *order, expand_dims(clamped, 1, s), 0, s); + out.push_back(reshape(inverse, a.shape(), s)); + } + // 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( + zeros({buffer_size}, int32, s), + group, + ones({n, 1}, int32, 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 ab676ca6c3..fb04e6e074 100644 --- a/mlx/ops.h +++ b/mlx/ops.h @@ -881,6 +881,27 @@ 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 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 + * 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. + */ +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, + 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 fe63bbf71f..a4d4d7f67e 100644 --- a/python/src/ops.cpp +++ b/python/src/ops.cpp @@ -3246,6 +3246,108 @@ 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_index, + 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_index, + 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_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_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. + + 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 + ``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. + + 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 + 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``. + 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``, this defaults to the + first of the sorted unique elements. Default: ``None``. + + Returns: + 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]) + >>> mx.unique(a, 3) + array([1, 2, 3], dtype=int32) + >>> mx.unique(a, 5) + array([1, 2, 3, 1, 1], dtype=int32) + >>> 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 + array([2, 2, 1, 0], dtype=int32) + )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..a718dc9d79 100644 --- a/python/tests/test_export_import.py +++ b/python/tests/test_export_import.py @@ -332,6 +332,33 @@ 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. 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, + 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) + 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 151eddea8e..e78870c533 100644 --- a/python/tests/test_ops.py +++ b/python/tests/test_ops.py @@ -3202,6 +3202,267 @@ 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): + 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): + mx.unique(a) + with self.assertRaises(ValueError): + mx.unique(a, -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, 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]) + 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])) + + # 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): + 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, 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 + 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, 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, 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, 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, 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, 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, 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]) + + # 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, + return_inverse=True, + return_counts=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, + return_inverse=True, + return_counts=True, + stream=mx.cpu, + ) + expect(values, [False, True]) + expect(inverse, [1, 0, 1]) + expect(counts, [1, 2]) + + # 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) + expect(values[:2], [1.0, 2.0]) + self.assertTrue(bool(mx.all(mx.isnan(values[2:])))) + values = mx.unique(mx.array([3.0, 1.0, nan, 2.0], mx.float32), 6) + 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 + 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)) + ) + + # 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, 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]]) + + # compare against numpy on random inputs + rng = np.random.RandomState(0) + for n in (1, 2, 17, 1000, 5000): + 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), 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)) + ) + 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, 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)) + + # 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 ( + (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))), + ): + want = np.unique(v_np) + 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, "This test requires a lot of memory",