| 1119 | self.assertEqual(mx.all(a, axis=1).tolist(), [False, True]) |
| 1120 | |
| 1121 | def test_any(self): |
| 1122 | a = mx.array([[True, False], [False, False]]) |
| 1123 | |
| 1124 | self.assertTrue(mx.any(a).item()) |
| 1125 | self.assertEqual(mx.any(a, keepdims=True).shape, (1, 1)) |
| 1126 | self.assertTrue(mx.any(a, axis=[0, 1]).item()) |
| 1127 | self.assertEqual(mx.any(a, axis=[0]).tolist(), [True, False]) |
| 1128 | self.assertEqual(mx.any(a, axis=[1]).tolist(), [True, False]) |
| 1129 | self.assertEqual(mx.any(a, axis=0).tolist(), [True, False]) |
| 1130 | self.assertEqual(mx.any(a, axis=1).tolist(), [True, False]) |
| 1131 | |
| 1132 | def test_stop_gradient(self): |
| 1133 | def func(x): |