| 354 | self.assertEqual(mx.not_equal(a, b).tolist(), [False, True, False, True]) |
| 355 | |
| 356 | def test_array_equal(self): |
| 357 | x = mx.array([1, 2, 3, 4]) |
| 358 | y = mx.array([1, 2, 3, 4]) |
| 359 | self.assertTrue(mx.array_equal(x, y)) |
| 360 | |
| 361 | y = mx.array([1, 2, 4, 5]) |
| 362 | self.assertFalse(mx.array_equal(x, y)) |
| 363 | |
| 364 | y = mx.array([1, 2, 3]) |
| 365 | self.assertFalse(mx.array_equal(x, y)) |
| 366 | |
| 367 | # Can still be equal with different types |
| 368 | y = mx.array([1.0, 2.0, 3.0, 4.0]) |
| 369 | self.assertTrue(mx.array_equal(x, y)) |
| 370 | |
| 371 | x = mx.array([0.0, float("nan")]) |
| 372 | y = mx.array([0.0, float("nan")]) |
| 373 | self.assertFalse(mx.array_equal(x, y)) |
| 374 | self.assertTrue(mx.array_equal(x, y, equal_nan=True)) |
| 375 | |
| 376 | for t in [mx.float32, mx.float16, mx.bfloat16, mx.complex64]: |
| 377 | with self.subTest(type=t): |
| 378 | x = mx.array([0.0, float("nan")]).astype(t) |
| 379 | y = mx.array([0.0, float("nan")]).astype(t) |
| 380 | self.assertFalse(mx.array_equal(x, y)) |
| 381 | self.assertTrue(mx.array_equal(x, y, equal_nan=True)) |
| 382 | |
| 383 | def test_isnan(self): |
| 384 | x = mx.array([0.0, float("nan")]) |