diff --git a/mlx/primitives.cpp b/mlx/primitives.cpp index 92e54f99..5ac2548f 100644 --- a/mlx/primitives.cpp +++ b/mlx/primitives.cpp @@ -1863,7 +1863,7 @@ std::vector Equal::jvp( const std::vector& tangents, const std::vector& argnums) { auto shape = broadcast_shapes(primals[0].shape(), primals[1].shape()); - return {zeros(shape, bool_, stream())}; + return {zeros(shape, tangents[0].dtype(), stream())}; } std::vector Erf::vjp( @@ -2530,7 +2530,7 @@ std::vector Greater::jvp( const std::vector& tangents, const std::vector& argnums) { auto shape = broadcast_shapes(primals[0].shape(), primals[1].shape()); - return {zeros(shape, bool_, stream())}; + return {zeros(shape, tangents[0].dtype(), stream())}; } std::pair, std::vector> GreaterEqual::vmap( @@ -2557,7 +2557,7 @@ std::vector GreaterEqual::jvp( const std::vector& tangents, const std::vector& argnums) { auto shape = broadcast_shapes(primals[0].shape(), primals[1].shape()); - return {zeros(shape, bool_, stream())}; + return {zeros(shape, tangents[0].dtype(), stream())}; } std::vector Imag::vjp( @@ -2614,7 +2614,7 @@ std::vector Less::jvp( const std::vector& tangents, const std::vector& argnums) { auto shape = broadcast_shapes(primals[0].shape(), primals[1].shape()); - return {zeros(shape, bool_, stream())}; + return {zeros(shape, tangents[0].dtype(), stream())}; } std::pair, std::vector> LessEqual::vmap( @@ -2641,7 +2641,7 @@ std::vector LessEqual::jvp( const std::vector& tangents, const std::vector& argnums) { auto shape = broadcast_shapes(primals[0].shape(), primals[1].shape()); - return {zeros(shape, bool_, stream())}; + return {zeros(shape, tangents[0].dtype(), stream())}; } std::vector Log::vjp( @@ -3188,7 +3188,7 @@ std::vector NotEqual::jvp( const std::vector& tangents, const std::vector& argnums) { auto shape = broadcast_shapes(primals[0].shape(), primals[1].shape()); - return {zeros(shape, bool_, stream())}; + return {zeros(shape, tangents[0].dtype(), stream())}; } std::vector Pad::vjp( diff --git a/python/tests/test_autograd.py b/python/tests/test_autograd.py index c37161a4..c0609db9 100644 --- a/python/tests/test_autograd.py +++ b/python/tests/test_autograd.py @@ -30,6 +30,24 @@ class TestAutograd(mlx_tests.MLXTestCase): self.assertEqual(out[0].item(), 4.0 * 1.0 + 2.0 * 3.0) self.assertEqual(out[1].item(), 4.0 * 1.0 + 6.0 * 3.0) + def test_jvp_comparison_tangent_dtype(self): + # Comparison op JVP tangents should preserve the input tangent's + # dtype (e.g. float32), not return bool. Using bool tangents causes + # downstream ops like negative to crash. (issue #3081) + x = mx.array([1.0, -2.0, 3.0]) + t = mx.ones_like(x) + + for op in [ + mx.greater, + mx.less, + mx.equal, + mx.greater_equal, + mx.less_equal, + mx.not_equal, + ]: + _, tangents = mx.jvp(lambda x, _op=op: _op(x, 0.0), [x], [t]) + self.assertEqual(tangents[0].dtype, mx.float32) + def test_vjp(self): fun = lambda x: 2 * x out, dout = mx.vjp(fun, [mx.array(1.0)], [mx.array(2.0)])