(self)
| 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 |
nothing calls this directly
no test coverage detected