diff --git a/mlx/transforms.cpp b/mlx/transforms.cpp index 694bca53d8..55e75013bf 100644 --- a/mlx/transforms.cpp +++ b/mlx/transforms.cpp @@ -324,6 +324,15 @@ array eval_impl(std::vector outputs, bool async) { } catch (...) { } } + // Commit GPU streams before the wait. + for (auto& s : open_streams) { + if (s.device == Device::gpu) { + try { + gpu::finalize(s); + } catch (...) { + } + } + } for (auto& s : open_streams) { try { synchronize(s); diff --git a/python/tests/test_eval.py b/python/tests/test_eval.py index 265a090aa2..9c9cdfff03 100644 --- a/python/tests/test_eval.py +++ b/python/tests/test_eval.py @@ -236,6 +236,27 @@ def test_async_eval_error_in_synchronize(self): with self.assertRaises(RuntimeError): mx.synchronize(mx.cpu) + @unittest.skipIf(not mx.metal.is_available(), "Metal is not available") + def test_eval_exception_after_cross_stream_wait(self): + # gather_qqmm has no CPU kernel. It fails in eval after the CPU + # stream waits for x from the GPU. + x = mx.multiply( + mx.full((2, 64), 3.0, stream=mx.gpu), + 2.0, + stream=mx.gpu, + ) + wq, scales = mx.quantize( + mx.ones((32, 64), stream=mx.gpu), + mode="nvfp4", + stream=mx.gpu, + )[:2] + y = mx.gather_qqmm(x, wq, scales, mode="nvfp4", stream=mx.cpu) + with self.assertRaises(RuntimeError): + mx.eval(y) + + self.assertTrue(mx.all(x == 6.0).item()) + self.assertEqual((mx.ones((4,), stream=mx.cpu) + 1).sum().item(), 8.0) + if __name__ == "__main__": mlx_tests.MLXTestRunner()