Skip to content
Open
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
63 changes: 21 additions & 42 deletions mlx/ops.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,13 @@ namespace mlx::core {

namespace {

// Returns {0, 1, ..., ndim - 1}, i.e. all axes of an ndim-dimensional array.
std::vector<int> all_axes(int ndim) {
std::vector<int> axes(ndim);
std::iota(axes.begin(), axes.end(), 0);
return axes;
}

std::tuple<Shape, std::vector<int>, bool> compute_reduce_shape(
const std::vector<int>& axes,
const Shape& shape) {
Expand Down Expand Up @@ -699,9 +706,7 @@ array flip(const array& a, int axis, StreamOrDevice s /* = {} */) {
}

array flip(const array& a, StreamOrDevice s /* = {} */) {
std::vector<int> axes(a.ndim());
std::iota(axes.begin(), axes.end(), 0);
return flip(a, axes, s);
return flip(a, all_axes(a.ndim()), s);
}

// Slice helper
Expand Down Expand Up @@ -2073,9 +2078,7 @@ array isclose(
}

array all(const array& a, bool keepdims, StreamOrDevice s /* = {}*/) {
std::vector<int> axes(a.ndim());
std::iota(axes.begin(), axes.end(), 0);
return all(a, axes, keepdims, s);
return all(a, all_axes(a.ndim()), keepdims, s);
}

array all(
Expand Down Expand Up @@ -2107,9 +2110,7 @@ array all(
}

array any(const array& a, bool keepdims, StreamOrDevice s /* = {}*/) {
std::vector<int> axes(a.ndim());
std::iota(axes.begin(), axes.end(), 0);
return any(a, axes, keepdims, s);
return any(a, all_axes(a.ndim()), keepdims, s);
}

array any(
Expand Down Expand Up @@ -2141,9 +2142,7 @@ array any(
}

array sum(const array& a, bool keepdims, StreamOrDevice s /* = {}*/) {
std::vector<int> axes(a.ndim());
std::iota(axes.begin(), axes.end(), 0);
return sum(a, axes, keepdims, s);
return sum(a, all_axes(a.ndim()), keepdims, s);
}

array sum(
Expand Down Expand Up @@ -2189,9 +2188,7 @@ array count_nonzero(
const array& a,
bool keepdims /* = false */,
StreamOrDevice s /* = {} */) {
std::vector<int> axes(a.ndim());
std::iota(axes.begin(), axes.end(), 0);
return count_nonzero(a, axes, keepdims, s);
return count_nonzero(a, all_axes(a.ndim()), keepdims, s);
}

array count_nonzero(
Expand All @@ -2212,9 +2209,7 @@ array count_nonzero(
}

array mean(const array& a, bool keepdims, StreamOrDevice s /* = {}*/) {
std::vector<int> axes(a.ndim());
std::iota(axes.begin(), axes.end(), 0);
return mean(a, axes, keepdims, to_stream(s));
return mean(a, all_axes(a.ndim()), keepdims, to_stream(s));
}

array mean(
Expand Down Expand Up @@ -2245,9 +2240,7 @@ array mean(
}

array median(const array& a, bool keepdims, StreamOrDevice s /* = {}*/) {
std::vector<int> axes(a.ndim());
std::iota(axes.begin(), axes.end(), 0);
return median(a, axes, keepdims, to_stream(s));
return median(a, all_axes(a.ndim()), keepdims, to_stream(s));
}

array median(
Expand Down Expand Up @@ -2330,9 +2323,7 @@ array var(
bool keepdims,
int ddof /* = 0*/,
StreamOrDevice s /* = {}*/) {
std::vector<int> axes(a.ndim());
std::iota(axes.begin(), axes.end(), 0);
return var(a, axes, keepdims, ddof, to_stream(s));
return var(a, all_axes(a.ndim()), keepdims, ddof, to_stream(s));
}

array var(
Expand Down Expand Up @@ -2376,9 +2367,7 @@ array std(
bool keepdims,
int ddof /* = 0*/,
StreamOrDevice s /* = {}*/) {
std::vector<int> axes(a.ndim());
std::iota(axes.begin(), axes.end(), 0);
return std(a, axes, keepdims, ddof, to_stream(s));
return std(a, all_axes(a.ndim()), keepdims, ddof, to_stream(s));
}

array std(
Expand All @@ -2400,9 +2389,7 @@ array std(
}

array prod(const array& a, bool keepdims, StreamOrDevice s /* = {}*/) {
std::vector<int> axes(a.ndim());
std::iota(axes.begin(), axes.end(), 0);
return prod(a, axes, keepdims, s);
return prod(a, all_axes(a.ndim()), keepdims, s);
}

array prod(
Expand Down Expand Up @@ -2445,9 +2432,7 @@ array prod(
}

array max(const array& a, bool keepdims, StreamOrDevice s /* = {}*/) {
std::vector<int> axes(a.ndim());
std::iota(axes.begin(), axes.end(), 0);
return max(a, axes, keepdims, s);
return max(a, all_axes(a.ndim()), keepdims, s);
}

array max(
Expand Down Expand Up @@ -2482,9 +2467,7 @@ array max(
}

array min(const array& a, bool keepdims, StreamOrDevice s /* = {}*/) {
std::vector<int> axes(a.ndim());
std::iota(axes.begin(), axes.end(), 0);
return min(a, axes, keepdims, s);
return min(a, all_axes(a.ndim()), keepdims, s);
}

array min(
Expand Down Expand Up @@ -2831,9 +2814,7 @@ array topk(const array& a, int k, int axis, StreamOrDevice s /* = {}*/) {
}

array logsumexp(const array& a, bool keepdims, StreamOrDevice s /* = {}*/) {
std::vector<int> axes(a.ndim());
std::iota(axes.begin(), axes.end(), 0);
return logsumexp(a, axes, keepdims, s);
return logsumexp(a, all_axes(a.ndim()), keepdims, s);
}

array logsumexp(
Expand Down Expand Up @@ -4006,9 +3987,7 @@ array softmax(
const array& a,
bool precise /* = false */,
StreamOrDevice s /* = {}*/) {
std::vector<int> axes(a.ndim());
std::iota(axes.begin(), axes.end(), 0);
return softmax(a, axes, precise, s);
return softmax(a, all_axes(a.ndim()), precise, s);
}

array power(const array& a, const array& b, StreamOrDevice s /* = {} */) {
Expand Down