--- a/mlx/primitives.cpp +++ b/mlx/primitives.cpp @@ -1803,6 +1803,35 @@ return vjps; } +namespace { + +array real_divide_denominator_partial( + const array& weight, const array& numerator, const array& denominator, Stream s) { + auto one = array(1., denominator.dtype()); + auto aw = abs(stop_gradient(weight, s), s); + auto an = abs(stop_gradient(numerator, s), s); + auto ad = abs(stop_gradient(denominator, s), s); + auto weight_is_large = greater_equal(aw, an, s); + auto large = where(weight_is_large, weight, numerator, s); + auto small = where(weight_is_large, numerator, weight, s); + auto denominator_is_large = greater_equal(ad, one, s); + auto split = where( + denominator_is_large, + greater_equal(minimum(aw, an, s), ad, s), + isfinite(divide(stop_gradient(large, s), stop_gradient(denominator, s), s), s), + s); + auto divide_first = logical_or(denominator_is_large, split, s); + auto first_divisor = where(divide_first, denominator, one, s); + auto second_divisor = where(split, denominator, one, s); + auto next_divisor = where(split, one, denominator, s); + auto last_divisor = where(divide_first, one, denominator, s); + auto product = multiply( + divide(large, first_divisor, s), divide(small, second_divisor, s), s); + return negative(divide(divide(product, next_divisor, s), last_divisor, s), s); +} + +} // namespace + std::vector Divide::vjp( const std::vector& primals, const std::vector& cotangents, @@ -1813,6 +1842,9 @@ for (auto arg : argnums) { if (arg == 0) { vjps.push_back(divide(cotangents[0], denominator_bar, stream())); + } else if (!issubdtype(primals[0].dtype(), complexfloating)) { + vjps.push_back(real_divide_denominator_partial( + cotangents[0], primals[0], primals[1], stream())); } else { vjps.push_back(negative( divide( @@ -1862,6 +1894,9 @@ int arg = argnums[i]; if (arg == 0) { return divide(tangents[i], primals[1], stream()); + } else if (!issubdtype(primals[0].dtype(), complexfloating)) { + return real_divide_denominator_partial( + tangents[i], primals[0], primals[1], stream()); } else { return negative( divide(