| 731 | self.assertTrue(np.array_equal(b_npy, b_mlx)) |
| 732 | |
| 733 | def test_logsumexp(self): |
| 734 | def logsumexp(x, axes=None): |
| 735 | maxs = mx.max(x, axis=axes, keepdims=True) |
| 736 | return mx.log(mx.sum(mx.exp(x - maxs), axis=axes, keepdims=True)) + maxs |
| 737 | |
| 738 | x = mx.array( |
| 739 | [ |
| 740 | [1.0, 2.0], |
| 741 | [3.0, 4.0], |
| 742 | ] |
| 743 | ) |
| 744 | self.assertTrue(math.isclose(mx.logsumexp(x).item(), logsumexp(x).item())) |
| 745 | |
| 746 | x = mx.random.uniform(shape=(1025,)) |
| 747 | self.assertTrue(mx.allclose(mx.logsumexp(x), logsumexp(x))) |
| 748 | |
| 749 | # Transposed |
| 750 | x = mx.random.uniform(shape=(2, 2, 8)) |
| 751 | x = x.swapaxes(0, 1) |
| 752 | self.assertTrue(mx.allclose(mx.logsumexp(x), logsumexp(x))) |
| 753 | |
| 754 | # Broadcast |
| 755 | x = mx.broadcast_to(mx.random.uniform(shape=(2, 1, 8)), (2, 2, 8)) |
| 756 | self.assertTrue(mx.allclose(mx.logsumexp(x), logsumexp(x))) |
| 757 | |
| 758 | # Large |
| 759 | x = mx.random.uniform(shape=(1025,)) |
| 760 | x = mx.broadcast_to(mx.random.uniform(shape=(2, 1, 8)), (2, 2, 8)) |
| 761 | self.assertTrue(mx.allclose(mx.logsumexp(x), logsumexp(x))) |
| 762 | |
| 763 | def test_mean(self): |
| 764 | x = mx.array( |