| 857 | self.assertEqual(y.dtype, dtype_out) |
| 858 | |
| 859 | def test_array_comparison(self): |
| 860 | a = mx.array([0.0, 1.0, 5.0]) |
| 861 | b = mx.array([-1.0, 2.0, 5.0]) |
| 862 | |
| 863 | self.assertEqual((a < b).tolist(), [False, True, False]) |
| 864 | self.assertEqual((a <= b).tolist(), [False, True, True]) |
| 865 | self.assertEqual((a > b).tolist(), [True, False, False]) |
| 866 | self.assertEqual((a >= b).tolist(), [True, False, True]) |
| 867 | |
| 868 | self.assertEqual((a < 5).tolist(), [True, True, False]) |
| 869 | self.assertEqual((5 < a).tolist(), [False, False, False]) |
| 870 | self.assertEqual((5 <= a).tolist(), [False, False, True]) |
| 871 | self.assertEqual((a > 1).tolist(), [False, False, True]) |
| 872 | self.assertEqual((a >= 1).tolist(), [False, True, True]) |
| 873 | |
| 874 | def test_array_neg(self): |
| 875 | a = mx.array([-1.0, 4.0, 0.0]) |