--- a/mlx/primitives.cpp +++ b/mlx/primitives.cpp @@ -413,6 +413,26 @@ return {{arccos(inputs[0], stream())}, axes}; } +namespace { + +array real_inverse_hyperbolic_slope( + const array& x, bool acosh, Stream s) { + auto one = array(1., x.dtype()); + auto scale = stop_gradient( + where(isfinite(x, s), maximum(abs(x, s), one, s), one, s), s); + if (acosh) { + auto lower = divide(subtract(x, one, s), scale, s); + auto upper = divide(add(x, one, s), scale, s); + return divide(rsqrt(multiply(lower, upper, s), s), scale, s); + } + auto value = divide(x, scale, s); + auto unit = divide(one, scale, s); + auto norm = add(square(value, s), square(unit, s), s); + return divide(rsqrt(norm, s), scale, s); +} + +} // namespace + std::vector ArcCosh::vjp( const std::vector& primals, const std::vector& cotangents, @@ -427,6 +447,12 @@ const std::vector& argnums) { assert(primals.size() == 1); assert(argnums.size() == 1); + if (!issubdtype(primals[0].dtype(), complexfloating)) { + return {multiply( + tangents[0], + real_inverse_hyperbolic_slope(primals[0], true, stream()), + stream())}; + } array one = array(1., primals[0].dtype()); array t = subtract(square(primals[0], stream()), one, stream()); return {multiply(tangents[0], rsqrt(t, stream()), stream())}; @@ -481,6 +507,12 @@ const std::vector& argnums) { assert(primals.size() == 1); assert(argnums.size() == 1); + if (!issubdtype(primals[0].dtype(), complexfloating)) { + return {multiply( + tangents[0], + real_inverse_hyperbolic_slope(primals[0], false, stream()), + stream())}; + } array one = array(1., primals[0].dtype()); array t = add(square(primals[0], stream()), one, stream()); return {multiply(tangents[0], rsqrt(t, stream()), stream())};