LogCumSumExp (#2069)
This commit is contained in:
@@ -1202,6 +1202,28 @@ void init_array(nb::module_& m) {
|
||||
nb::kw_only(),
|
||||
"stream"_a = nb::none(),
|
||||
"See :func:`max`.")
|
||||
.def(
|
||||
"logcumsumexp",
|
||||
[](const mx::array& a,
|
||||
std::optional<int> axis,
|
||||
bool reverse,
|
||||
bool inclusive,
|
||||
mx::StreamOrDevice s) {
|
||||
if (axis) {
|
||||
return mx::logcumsumexp(a, *axis, reverse, inclusive, s);
|
||||
} else {
|
||||
// TODO: Implement that in the C++ API as well. See concatenate
|
||||
// above.
|
||||
return mx::logcumsumexp(
|
||||
mx::reshape(a, {-1}, s), 0, reverse, inclusive, s);
|
||||
}
|
||||
},
|
||||
"axis"_a = nb::none(),
|
||||
nb::kw_only(),
|
||||
"reverse"_a = false,
|
||||
"inclusive"_a = true,
|
||||
"stream"_a = nb::none(),
|
||||
"See :func:`logcumsumexp`.")
|
||||
.def(
|
||||
"logsumexp",
|
||||
[](const mx::array& a,
|
||||
|
||||
@@ -2382,6 +2382,43 @@ void init_ops(nb::module_& m) {
|
||||
Returns:
|
||||
array: The output array with the corresponding axes reduced.
|
||||
)pbdoc");
|
||||
m.def(
|
||||
"logcumsumexp",
|
||||
[](const mx::array& a,
|
||||
std::optional<int> axis,
|
||||
bool reverse,
|
||||
bool inclusive,
|
||||
mx::StreamOrDevice s) {
|
||||
if (axis) {
|
||||
return mx::logcumsumexp(a, *axis, reverse, inclusive, s);
|
||||
} else {
|
||||
return mx::logcumsumexp(
|
||||
mx::reshape(a, {-1}, s), 0, reverse, inclusive, s);
|
||||
}
|
||||
},
|
||||
nb::arg(),
|
||||
"axis"_a = nb::none(),
|
||||
nb::kw_only(),
|
||||
"reverse"_a = false,
|
||||
"inclusive"_a = true,
|
||||
"stream"_a = nb::none(),
|
||||
nb::sig(
|
||||
"def logcumsumexp(a: array, /, axis: Optional[int] = None, *, reverse: bool = False, inclusive: bool = True, stream: Union[None, Stream, Device] = None) -> array"),
|
||||
R"pbdoc(
|
||||
Return the cumulative logsumexp of the elements along the given axis.
|
||||
|
||||
Args:
|
||||
a (array): Input array
|
||||
axis (int, optional): Optional axis to compute the cumulative logsumexp
|
||||
over. If unspecified the cumulative logsumexp of the flattened array is
|
||||
returned.
|
||||
reverse (bool): Perform the cumulative logsumexp in reverse.
|
||||
inclusive (bool): The i-th element of the output includes the i-th
|
||||
element of the input.
|
||||
|
||||
Returns:
|
||||
array: The output array.
|
||||
)pbdoc");
|
||||
m.def(
|
||||
"logsumexp",
|
||||
[](const mx::array& a,
|
||||
|
||||
@@ -1508,6 +1508,7 @@ class TestArray(mlx_tests.MLXTestCase):
|
||||
("prod", 1),
|
||||
("min", 1),
|
||||
("max", 1),
|
||||
("logcumsumexp", 1),
|
||||
("logsumexp", 1),
|
||||
("mean", 1),
|
||||
("var", 1),
|
||||
|
||||
@@ -1857,6 +1857,30 @@ class TestOps(mlx_tests.MLXTestCase):
|
||||
y = mx.as_strided(x, (x.size,), (-1,), x.size - 1)
|
||||
self.assertTrue(mx.array_equal(y, x[::-1]))
|
||||
|
||||
def test_logcumsumexp(self):
|
||||
npop = np.logaddexp.accumulate
|
||||
mxop = mx.logcumsumexp
|
||||
|
||||
a_npy = np.random.randn(32, 32, 32).astype(np.float32)
|
||||
a_mlx = mx.array(a_npy)
|
||||
|
||||
for axis in (0, 1, 2):
|
||||
c_npy = npop(a_npy, axis=axis)
|
||||
c_mlx = mxop(a_mlx, axis=axis)
|
||||
self.assertTrue(np.allclose(c_npy, c_mlx, rtol=1e-3, atol=1e-3))
|
||||
|
||||
edge_cases_npy = [
|
||||
np.float32([-float("inf")] * 8),
|
||||
np.float32([-float("inf"), 0, -float("inf")]),
|
||||
np.float32([-float("inf"), float("inf"), -float("inf")]),
|
||||
]
|
||||
edge_cases_mlx = [mx.array(a) for a in edge_cases_npy]
|
||||
|
||||
for a_npy, a_mlx in zip(edge_cases_npy, edge_cases_mlx):
|
||||
c_npy = npop(a_npy, axis=0)
|
||||
c_mlx = mxop(a_mlx, axis=0)
|
||||
self.assertTrue(np.allclose(c_npy, c_mlx, rtol=1e-3, atol=1e-3))
|
||||
|
||||
def test_scans(self):
|
||||
a_npy = np.random.randn(32, 32, 32).astype(np.float32)
|
||||
a_mlx = mx.array(a_npy)
|
||||
|
||||
Reference in New Issue
Block a user