From d167dffd677aa5b888d30b2f5defe134a8e70686 Mon Sep 17 00:00:00 2001 From: Ryan11c Date: Mon, 14 Sep 2026 20:26:41 -0700 Subject: [PATCH 1/3] Fix precision of constants in compiled kernels --- mlx/backend/common/compiled.h | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/mlx/backend/common/compiled.h b/mlx/backend/common/compiled.h index 8c6466da03..94ba38b8f6 100644 --- a/mlx/backend/common/compiled.h +++ b/mlx/backend/common/compiled.h @@ -37,9 +37,9 @@ void print_float_constant(std::ostream& os, const array& x) { auto old_precision = os.precision(); if constexpr (std::is_same_v) { - os << std::setprecision(std::numeric_limits::digits10 + 1); + os << std::setprecision(std::numeric_limits::max_digits10); } else { - os << std::setprecision(std::numeric_limits::digits10 + 1); + os << std::setprecision(std::numeric_limits::max_digits10); } os << value << std::setprecision(old_precision); } From 3364618d2f512356974592524823d89c7efa77c5 Mon Sep 17 00:00:00 2001 From: Ryan11c Date: Mon, 14 Sep 2026 20:28:51 -0700 Subject: [PATCH 2/3] Add tests for compiled constant precision --- python/tests/test_compile.py | 15 ++++++++++----- 1 file changed, 10 insertions(+), 5 deletions(-) diff --git a/python/tests/test_compile.py b/python/tests/test_compile.py index ed8a7369b7..2e1192fb28 100644 --- a/python/tests/test_compile.py +++ b/python/tests/test_compile.py @@ -88,6 +88,12 @@ def test_compile_nonfinite_constants(self): self.assertEqual(out[0].item(), 1.0) self.assertEqual(out[1].item(), float("-inf")) + def test_compile_float_constant_precision(self): + x = mx.ones((4,), dtype=mx.float32) + for constant in (1 / 3, 128**-0.5, 2**-0.5): + fun = lambda x, constant=constant: (x * x) * constant + self.assertTrue(mx.array_equal(mx.compile(fun)(x), fun(x))) + def test_compile_tuple_output_in_thread(self): @mx.compile def fun(x): @@ -1337,14 +1343,13 @@ def f(x): def test_double_constant(self): with mx.stream(mx.cpu): - x = mx.array(1.0, dtype=mx.float64) + x = mx.array([1.0], dtype=mx.float64) + constant = math.nextafter(1.0, 2.0) def fun(x): - return (x + math.pi) * 2.0 + return (x * x) * constant - y = fun(x).item() - y_compiled = mx.compile(fun)(x).item() - self.assertEqual(y, y_compiled) + self.assertTrue(mx.array_equal(fun(x), mx.compile(fun)(x))) def test_shared_broadcast(self): def fun(x, y, z): From ce352a2c89020eb86576eef302683688b2e65388 Mon Sep 17 00:00:00 2001 From: Ryan11c Date: Mon, 14 Sep 2026 21:07:36 -0700 Subject: [PATCH 3/3] Align float precision test with issue --- python/tests/test_compile.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/python/tests/test_compile.py b/python/tests/test_compile.py index 2e1192fb28..185019342a 100644 --- a/python/tests/test_compile.py +++ b/python/tests/test_compile.py @@ -90,7 +90,7 @@ def test_compile_nonfinite_constants(self): def test_compile_float_constant_precision(self): x = mx.ones((4,), dtype=mx.float32) - for constant in (1 / 3, 128**-0.5, 2**-0.5): + for constant in (1 / 3, 128**-0.5, 0.7071067811865476): fun = lambda x, constant=constant: (x * x) * constant self.assertTrue(mx.array_equal(mx.compile(fun)(x), fun(x)))