Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions docs/src/python/ops.rst
Original file line number Diff line number Diff line change
Expand Up @@ -203,6 +203,7 @@ Operations
triu
trunc
unflatten
unique
unstack
vecdot
var
Expand Down
142 changes: 142 additions & 0 deletions mlx/ops.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -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<array> unique(
const array& a,
int size,
bool return_index /* = false */,
bool return_inverse /* = false */,
bool return_counts /* = false */,
Comment thread
louen marked this conversation as resolved.
const std::optional<array>& 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<array> 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<array> 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<array> 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<int> axes(a.ndim());
std::iota(axes.begin(), axes.end(), 0);
Expand Down
21 changes: 21 additions & 0 deletions mlx/ops.h
Original file line number Diff line number Diff line change
Expand Up @@ -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<array> unique(
const array& a,
int size,
bool return_index = false,
bool return_inverse = false,
bool return_counts = false,
const std::optional<array>& fill_value = std::nullopt,
StreamOrDevice s = {});

/** Cumulative logsumexp of an array. */
MLX_API array logcumsumexp(
const array& a,
Expand Down
102 changes: 102 additions & 0 deletions python/src/ops.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -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<ScalarOrArray>& fill_value,
mx::StreamOrDevice s) -> nb::object {
std::optional<mx::array> 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) {
Expand Down
27 changes: 27 additions & 0 deletions python/tests/test_export_import.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")

Expand Down
Loading
Loading