MCPcopy Create free account
hub / github.com/ml-explore/mlx / test_logsumexp

Method test_logsumexp

python/tests/test_ops.py:733–761  ·  view source on GitHub ↗
(self)

Source from the content-addressed store, hash-verified

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(

Callers

nothing calls this directly

Calls 4

itemMethod · 0.80
logsumexpMethod · 0.80
arrayMethod · 0.60
logsumexpFunction · 0.50

Tested by

no test coverage detected