From 7b13d41a29f2b0a219265cf65b70ed2f646107b7 Mon Sep 17 00:00:00 2001 From: Ronald Mannak Date: Wed, 23 Sep 2026 20:48:25 -0700 Subject: [PATCH 1/2] Commit GPU streams before the wait --- mlx/transforms.cpp | 9 +++++++++ python/tests/test_eval.py | 13 +++++++++++++ 2 files changed, 22 insertions(+) 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..f3af9b48ee 100644 --- a/python/tests/test_eval.py +++ b/python/tests/test_eval.py @@ -236,6 +236,19 @@ 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.full((2, 64), 3.0) * 2.0 + wq, scales = mx.quantize(mx.ones((32, 64)), mode="nvfp4")[: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() From c6c00c722814762aafcd71a4051552476e524257 Mon Sep 17 00:00:00 2001 From: Ronald Mannak Date: Wed, 23 Sep 2026 21:12:07 -0700 Subject: [PATCH 2/2] set streams to .gpu explicitly --- python/tests/test_eval.py | 12 ++++++++++-- 1 file changed, 10 insertions(+), 2 deletions(-) diff --git a/python/tests/test_eval.py b/python/tests/test_eval.py index f3af9b48ee..9c9cdfff03 100644 --- a/python/tests/test_eval.py +++ b/python/tests/test_eval.py @@ -240,8 +240,16 @@ def test_async_eval_error_in_synchronize(self): 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.full((2, 64), 3.0) * 2.0 - wq, scales = mx.quantize(mx.ones((32, 64)), mode="nvfp4")[:2] + 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)