| 23 | |
| 24 | class TestDevice(mlx_tests.MLXTestCase): |
| 25 | def test_device(self): |
| 26 | device = mx.default_device() |
| 27 | |
| 28 | cpu = mx.Device(mx.cpu) |
| 29 | mx.set_default_device(cpu) |
| 30 | self.assertEqual(mx.default_device(), cpu) |
| 31 | self.assertEqual(str(cpu), "Device(cpu, 0)") |
| 32 | |
| 33 | mx.set_default_device(mx.cpu) |
| 34 | self.assertEqual(mx.default_device(), mx.cpu) |
| 35 | self.assertEqual(cpu, mx.cpu) |
| 36 | self.assertEqual(mx.cpu, cpu) |
| 37 | |
| 38 | # Restore device |
| 39 | mx.set_default_device(device) |
| 40 | |
| 41 | @unittest.skipIf(not mx.is_available(mx.gpu), "GPU is not available") |
| 42 | def test_device_context(self): |