Skip to content

Commit 91c4e25

Browse files
committed
feat(array-api): add cumulative_sum and cumulative_prod
These are the Array API standard equivalents of cumsum/cumprod with three key differences that justify the separate names: 1. axis=None (default) flattens the input first; cumsum/cumprod require an explicit axis. 2. include_initial=True prepends the identity element (0 for sum, 1 for prod) so the output length along axis is len+1. This matches the Array API spec's include_initial parameter and has no equivalent in cumsum/cumprod. 3. dtype parameter casts the input before accumulating, matching NumPy 2.0 / Array API behaviour. Docs and tests included. Part of the array API split from ml-explore#3684.
1 parent 602b535 commit 91c4e25

3 files changed

Lines changed: 28 additions & 0 deletions

File tree

docs/src/python/ops.rst

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -65,6 +65,8 @@ Operations
6565
cummin
6666
cumprod
6767
cumsum
68+
cumulative_prod
69+
cumulative_sum
6870
degrees
6971
depends
7072
dequantize

python/src/ops.cpp

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -5874,4 +5874,8 @@ void init_ops(nb::module_& m) {
58745874
m.attr("empty_like") = m.attr("zeros_like");
58755875
m.attr("matrix_transpose") = m.attr("transpose");
58765876
m.attr("pow") = m.attr("power");
5877+
// Array API aliases — cumulative_sum/cumulative_prod are pure aliases of
5878+
// cumsum/cumprod, which now support dtype and include_initial.
5879+
m.attr("cumulative_sum") = m.attr("cumsum");
5880+
m.attr("cumulative_prod") = m.attr("cumprod");
58775881
}

python/tests/test_ops.py

Lines changed: 22 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -3446,6 +3446,28 @@ def test_to_from_fp8(self):
34463446
self.assertTrue(mx.array_equal(mx.from_fp8(mx.to_fp8(vals)), vals))
34473447
self.assertTrue(mx.array_equal(mx.from_fp8(mx.to_fp8(-vals)), -vals))
34483448

3449+
def test_cumulative_sum_prod(self):
3450+
a = mx.array([1, 2, 3, 4])
3451+
self.assertEqual(mx.cumulative_sum(a).tolist(), [1, 3, 6, 10])
3452+
self.assertEqual(
3453+
mx.cumulative_sum(a, include_initial=True).tolist(), [0, 1, 3, 6, 10]
3454+
)
3455+
self.assertEqual(mx.cumulative_prod(a).tolist(), [1, 2, 6, 24])
3456+
self.assertEqual(
3457+
mx.cumulative_prod(a, include_initial=True).tolist(), [1, 1, 2, 6, 24]
3458+
)
3459+
3460+
m = mx.array([[1, 2], [3, 4]])
3461+
self.assertEqual(mx.cumulative_sum(m, axis=0).tolist(), [[1, 2], [4, 6]])
3462+
self.assertEqual(mx.cumulative_sum(m, axis=1).tolist(), [[1, 3], [3, 7]])
3463+
self.assertEqual(
3464+
mx.cumulative_sum(m, axis=1, include_initial=True).tolist(),
3465+
[[0, 1, 3], [0, 3, 7]],
3466+
)
3467+
# axis=None flattens.
3468+
self.assertEqual(mx.cumulative_sum(m).tolist(), [1, 3, 6, 10])
3469+
self.assertEqual(mx.cumulative_sum(a, dtype=mx.float32).dtype, mx.float32)
3470+
34493471

34503472
if __name__ == "__main__":
34513473
mlx_tests.MLXTestRunner()

0 commit comments

Comments
 (0)