--- a/mlx/primitives.cpp +++ b/mlx/primitives.cpp @@ -521,6 +521,43 @@ return {{arctan(inputs[0], stream())}, axes}; } +namespace { + +std::pair arctan2_scale_norm( + const array& y, const array& x, Stream s) { + auto one = array(1., y.dtype()); + auto magnitude = maximum(abs(y, s), abs(x, s), s); + auto valid = logical_and( + isfinite(magnitude, s), greater(magnitude, array(0., y.dtype()), s), s); + auto scale = stop_gradient(where(valid, magnitude, one, s), s); + auto u = divide(y, scale, s); + auto v = divide(x, scale, s); + auto norm = add(square(u, s), square(v, s), s); + return {scale, norm}; +} + +array arctan2_weighted_partial( + const array& weight, + const array& coordinate, + const array& scale, + const array& norm, + Stream s) { + auto one = array(1., weight.dtype()); + auto divide_first = isfinite(divide(stop_gradient(weight, s), scale, s), s); + auto divide_coordinate = logical_and(divide_first, less(scale, one, s), s); + auto first_divisor = where(divide_first, scale, one, s); + auto coordinate_divisor = where(divide_coordinate, scale, one, s); + auto next_divisor = where(divide_coordinate, one, scale, s); + auto last_divisor = where(divide_first, one, scale, s); + auto numerator = multiply( + divide(weight, first_divisor, s), + divide(coordinate, coordinate_divisor, s), s); + auto term = divide(numerator, norm, s); + return divide(divide(term, next_divisor, s), last_divisor, s); +} + +} // namespace + std::vector ArcTan2::vjp( const std::vector& primals, const std::vector& cotangents, @@ -528,24 +565,15 @@ const std::vector&) { assert(primals.size() == 2); assert(cotangents.size() == 1); - const auto& s = stream(); - const array& x1 = primals[0]; - const array& x2 = primals[1]; - const array& dy = cotangents[0]; - + auto [scale, norm] = arctan2_scale_norm(primals[0], primals[1], s); std::vector grads; - array dy_over_x1_x2_squared = - divide(dy, add(square(x1, s), square(x2, s)), s); - for (auto arg : argnums) { - if (arg == 0) { - grads.emplace_back(multiply(x2, dy_over_x1_x2_squared, s)); - } else { - grads.emplace_back(multiply(negative(x1, s), dy_over_x1_x2_squared, s)); - } - } - + grads.push_back(arctan2_weighted_partial( + cotangents[0], + arg == 0 ? primals[1] : negative(primals[0], s), + scale, norm, s)); + } return grads; } @@ -555,22 +583,20 @@ const std::vector& argnums) { assert(primals.size() == 2); assert(tangents.size() == argnums.size()); - + assert(!argnums.empty()); const auto& s = stream(); - const array& x1 = primals[0]; - const array& x2 = primals[1]; - - auto numerator = [&]() { - if (argnums.size() == 2) { - return subtract( - multiply(x2, tangents[0], s), multiply(x1, tangents[1], s), s); - } else if (argnums[0] == 0) { - return multiply(x2, tangents[0], s); - } else { - return negative(multiply(x1, tangents[0], s), s); - } - }(); - return {divide(numerator, add(square(x1, s), square(x2, s), s), s)}; + auto [scale, norm] = arctan2_scale_norm(primals[0], primals[1], s); + auto term = [&](int i) { + return arctan2_weighted_partial( + tangents[i], + argnums[i] == 0 ? primals[1] : negative(primals[0], s), + scale, norm, s); + }; + auto out = term(0); + for (int i = 1; i < argnums.size(); ++i) { + out = add(out, term(i), s); + } + return {out}; } std::pair, std::vector> ArcTan2::vmap(