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

Method test_trace

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

Source from the content-addressed store, hash-verified

2678 self.assertTrue(mx.array_equal(result, expected))
2679
2680 def test_trace(self):
2681 a_mx = mx.arange(9, dtype=mx.int64).reshape((3, 3))
2682 a_np = np.arange(9, dtype=np.int64).reshape((3, 3))
2683
2684 # Test 2D array
2685 result = mx.trace(a_mx)
2686 expected = np.trace(a_np)
2687 self.assertEqualArray(result, mx.array(expected))
2688
2689 # Test dtype
2690 result = mx.trace(a_mx, dtype=mx.float16)
2691 expected = np.trace(a_np, dtype=np.float16)
2692 self.assertEqualArray(result, mx.array(expected))
2693
2694 # Test offset
2695 result = mx.trace(a_mx, offset=1)
2696 expected = np.trace(a_np, offset=1)
2697 self.assertEqualArray(result, mx.array(expected))
2698
2699 # Test axis1 and axis2
2700 b_mx = mx.arange(27, dtype=mx.int64).reshape(3, 3, 3)
2701 b_np = np.arange(27, dtype=np.int64).reshape(3, 3, 3)
2702
2703 result = mx.trace(b_mx, axis1=1, axis2=2)
2704 expected = np.trace(b_np, axis1=1, axis2=2)
2705 self.assertEqualArray(result, mx.array(expected))
2706
2707 # Test offset, axis1, axis2, and dtype
2708 result = mx.trace(b_mx, offset=1, axis1=1, axis2=2, dtype=mx.float32)
2709 expected = np.trace(b_np, offset=1, axis1=1, axis2=2, dtype=np.float32)
2710 self.assertEqualArray(result, mx.array(expected))
2711
2712 def test_atleast_1d(self):
2713 # Test 1D input

Callers

nothing calls this directly

Calls 2

assertEqualArrayMethod · 0.80
arrayMethod · 0.60

Tested by

no test coverage detected