diff --git a/include/boost/multiprecision/detail/default_ops.hpp b/include/boost/multiprecision/detail/default_ops.hpp index e278777b..5dd099c4 100644 --- a/include/boost/multiprecision/detail/default_ops.hpp +++ b/include/boost/multiprecision/detail/default_ops.hpp @@ -45,6 +45,9 @@ void generic_interconvert(To& to, const From& from, const std::integral_constant template void generic_interconvert(To& to, const From& from, const std::integral_constant& /*to_type*/, const std::integral_constant& /*from_type*/); +template +BOOST_MP_CXX14_CONSTEXPR Integer karatsuba_sqrt(const Integer& x, Integer& r, Integer& t, size_t bits); + } // namespace detail namespace default_ops { @@ -1588,62 +1591,100 @@ inline BOOST_MP_CXX14_CONSTEXPR void eval_bit_unset(T& val, unsigned index) eval_bitwise_xor(val, mask); } -template -void BOOST_MP_CXX14_CONSTEXPR eval_integer_sqrt(B& s, B& r, const B& x) +template +BOOST_MP_CXX14_CONSTEXPR void eval_qr(const Backend& x, const Backend& y, Backend& q, Backend& r); + +template +BOOST_MP_CXX14_CONSTEXPR void eval_karatsuba_sqrt(Backend& result, const Backend& x, Backend& r, Backend& t, size_t bits) { - // - // This is slow bit-by-bit integer square root, see for example - // http://en.wikipedia.org/wiki/Methods_of_computing_square_roots#Binary_numeral_system_.28base_2.29 - // There are better methods such as http://hal.inria.fr/docs/00/07/28/54/PDF/RR-3805.pdf - // and http://hal.inria.fr/docs/00/07/21/13/PDF/RR-4475.pdf which should be implemented - // at some point. - // - using ui_type = typename boost::multiprecision::detail::canonical::type; + using default_ops::eval_is_zero; + using default_ops::eval_subtract; + using default_ops::eval_right_shift; + using default_ops::eval_left_shift; + using default_ops::eval_bit_set; + using default_ops::eval_decrement; + using default_ops::eval_bitwise_and; + using default_ops::eval_add; + using default_ops::eval_qr; - s = ui_type(0u); - if (eval_get_sign(x) == 0) + using small_uint = typename std::tuple_element<0, typename Backend::unsigned_types>::type; + + constexpr small_uint zero = 0u; + + // we can calculate it faster with std::sqrt +#ifdef BOOST_HAS_INT128 + if (bits <= 128) { - r = ui_type(0u); + unsigned __int128 a, b, c; + eval_convert_to(&a, x); + c = boost::multiprecision::detail::karatsuba_sqrt(a, b, c, bits); + r = number::canonical_value(b); + result = number::canonical_value(c); return; } - int g = eval_msb(x); - if (g <= 1) +#else + if (bits <= std::numeric_limits::digits) { - s = ui_type(1); - eval_subtract(r, x, s); + std::uintmax_t a, b, c; + eval_convert_to(&a, x); + c = boost::multiprecision::detail::karatsuba_sqrt(a, b, c, bits); + r = number::canonical_value(b); + result = number::canonical_value(c); return; } - - B t; - r = x; - g /= 2; - int org_g = g; - eval_bit_set(s, g); - eval_bit_set(t, 2 * g); - eval_subtract(r, x, t); - --g; - if (eval_get_sign(r) == 0) - return; - int msbr = eval_msb(r); - do +#endif + // https://hal.inria.fr/file/index/docid/72854/filename/RR-3805.pdf + std::size_t b = bits / 4; + Backend q(x); + eval_right_shift(q, b * 2); + Backend s; + eval_karatsuba_sqrt(s, q, r, t, bits - b * 2); + t = zero; + eval_bit_set(t, b * 2); + eval_left_shift(r, b); + eval_decrement(t); + eval_bitwise_and(t, x); + eval_right_shift(t, b); + eval_add(t, r); + eval_left_shift(s, 1u); + eval_qr(t, s, q, r); + eval_left_shift(r, b); + t = zero; + eval_bit_set(t, b); + eval_decrement(t); + eval_bitwise_and(t, x); + eval_add(r, t); + eval_left_shift(s, b - 1); + eval_add(s, q); + eval_multiply(q, q); + // we substract after, so it works for unsigned integers too + if (r.compare(q) < 0) { - if (msbr >= org_g + g + 1) - { - t = s; - eval_left_shift(t, g + 1); - eval_bit_set(t, 2 * g); - if (t.compare(r) <= 0) - { - BOOST_ASSERT(g >= 0); - eval_bit_set(s, g); - eval_subtract(r, t); - if (eval_get_sign(r) == 0) - return; - msbr = eval_msb(r); - } - } - --g; - } while (g >= 0); + t = s; + eval_left_shift(t, 1u); + eval_decrement(t); + eval_add(r, t); + eval_decrement(s); + } + eval_subtract(r, q); + result = s; +} + +template +BOOST_MP_CXX14_CONSTEXPR void eval_integer_sqrt(Backend& result, Backend& r, const Backend& x) +{ + using small_uint = typename std::tuple_element<0, typename Backend::unsigned_types>::type; + + constexpr small_uint zero = 0u; + + if (eval_is_zero(x)) + { + r = zero; + result = zero; + return; + } + Backend t; + eval_karatsuba_sqrt(result, x, r, t, eval_msb(x) + 1); } template diff --git a/include/boost/multiprecision/integer.hpp b/include/boost/multiprecision/integer.hpp index eeb6a82b..1db6f613 100644 --- a/include/boost/multiprecision/integer.hpp +++ b/include/boost/multiprecision/integer.hpp @@ -189,52 +189,63 @@ BOOST_MP_CXX14_CONSTEXPR typename std::enable_if -BOOST_MP_CXX14_CONSTEXPR typename std::enable_if::value, Integer>::type karatsuba_sqrt(const Integer& x, Integer& r, Integer& t, size_t bits) +BOOST_MP_CXX14_CONSTEXPR Integer karatsuba_sqrt(const Integer& x, Integer& r, Integer& t, size_t bits) { #ifndef BOOST_MP_NO_CONSTEXPR_DETECTION - // std::sqrt is not constexpr by standard, so use this - if (BOOST_MP_IS_CONST_EVALUATED(bits)) { - if (bits <= 4) { - if (x == 0) { + // std::sqrt is not constexpr by standard, so use this + if (BOOST_MP_IS_CONST_EVALUATED(bits)) + { + if (bits <= 4) + { + if (x == 0) + { r = 0u; - return 0u; + return 0u; } - else if (x < 4) { + else if (x < 4) + { r = x - 1; - return 1u; + return 1u; } - else if (x < 9) { + else if (x < 9) + { r = x - 4; - return 2u; + return 2u; } - else { + else + { r = x - 9; - return 3u; + return 3u; } } } else #endif // we can calculate it faster with std::sqrt - if (bits <= 64) { - const uint64_t int32max = uint64_t((std::numeric_limits::max)()); - uint64_t val = static_cast(x); - uint64_t s64 = static_cast(std::sqrt(static_cast(val))); + if (bits <= 64) + { + const std::uint64_t int32max = std::uint64_t((std::numeric_limits::max)()); + std::uint64_t val = static_cast(x); + std::uint64_t s64 = static_cast(std::sqrt(static_cast(val))); // converting to long double can loose some precision, and `sqrt` can give eps error, so we'll fix this // this is needed - while (s64 > int32max || s64 * s64 > val) s64--; + while (s64 > int32max || s64 * s64 > val) + s64--; // in my tests this never fired, but theoretically this might be needed - while (s64 < int32max && (s64 + 1) * (s64 + 1) <= val) s64++; + while (s64 < int32max && (s64 + 1) * (s64 + 1) <= val) + s64++; r = val - s64 * s64; return s64; } // https://hal.inria.fr/file/index/docid/72854/filename/RR-3805.pdf - size_t b = bits / 4; + std::size_t b = bits / 4; Integer q = x; q >>= b * 2; Integer s = karatsuba_sqrt(q, r, t, bits - b * 2); - t = 0u; + t = 0u; bit_set(t, b * 2); r <<= b; t--; @@ -253,7 +264,8 @@ BOOST_MP_CXX14_CONSTEXPR typename std::enable_if BOOST_MP_CXX14_CONSTEXPR typename std::enable_if::value, Integer>::type sqrt(const Integer& x, Integer& r) { @@ -272,7 +286,7 @@ BOOST_MP_CXX14_CONSTEXPR typename std::enable_if diff --git a/performance/sqrt_bench.cpp b/performance/sqrt_bench.cpp new file mode 100644 index 00000000..6dd6413c --- /dev/null +++ b/performance/sqrt_bench.cpp @@ -0,0 +1,144 @@ +// Copyright 2020 John Maddock. Distributed under the Boost +// Software License, Version 1.0. (See accompanying file +// LICENSE_1_0.txt or copy at https://www.boost.org/LICENSE_1_0.txt + +#include +#include +#include +#include +#include +#include +#include + +#include + +using namespace boost::multiprecision; +using namespace boost::random; + +template +BOOST_MP_CXX14_CONSTEXPR Integer sqrt_old(const Integer& x, Integer& r) +{ + // + // This is slow bit-by-bit integer square root, see for example + // http://en.wikipedia.org/wiki/Methods_of_computing_square_roots#Binary_numeral_system_.28base_2.29 + // There are better methods such as http://hal.inria.fr/docs/00/07/28/54/PDF/RR-3805.pdf + // and http://hal.inria.fr/docs/00/07/21/13/PDF/RR-4475.pdf which should be implemented + // at some point. + // + Integer s = 0; + if (x == 0) + { + r = 0; + return s; + } + int g = msb(x); + if (g == 0) + { + r = 1; + return s; + } + + Integer t = 0; + r = x; + g /= 2; + bit_set(s, g); + bit_set(t, 2 * g); + r = x - t; + --g; + do + { + t = s; + t <<= g + 1; + bit_set(t, 2 * g); + if (t <= r) + { + bit_set(s, g); + r -= t; + } + --g; + } while (g >= 0); + return s; +} + +template +BOOST_MP_CXX14_CONSTEXPR Integer sqrt_old(const Integer& x) +{ + Integer r(0); + return sqrt_old(x, r); +} + +template +std::tuple, std::vector >& get_test_vector(unsigned bits) +{ + static std::map, std::vector > > data; + + std::tuple, std::vector >& result = data[bits]; + + if (std::get<0>(result).size() == 0) + { + mt19937 mt; + uniform_int_distribution ui(T(1) << (bits - 1), T(1) << bits); + + std::vector& a = std::get<0>(result); + std::vector& b = std::get<1>(result); + + for (unsigned i = 0; i < 1000; ++i) + { + a.push_back(ui(mt)); + b.push_back(0); + } + } + return result; +} + +template +std::vector& get_test_vector_a(unsigned bits) +{ + return std::get<0>(get_test_vector(bits)); +} +template +std::vector& get_test_vector_b(unsigned bits) +{ + return std::get<1>(get_test_vector(bits)); +} + +template +static void BM_sqrt_old(benchmark::State& state) +{ + int bits = state.range(0); + + std::vector& a = get_test_vector_a(bits); + std::vector& b = get_test_vector_b(bits); + + for (auto _ : state) + { + for (unsigned i = 0; i < a.size(); ++i) + b[i] = sqrt_old(a[i]); + } + state.SetComplexityN(bits); +} + +template +static void BM_sqrt_current(benchmark::State& state) +{ + int bits = state.range(0); + + std::vector& a = get_test_vector_a(bits); + std::vector& b = get_test_vector_b(bits); + + for (auto _ : state) + { + for (unsigned i = 0; i < a.size(); ++i) + b[i] = sqrt(a[i]); + } + state.SetComplexityN(bits); +} + +constexpr unsigned lower_range = 512; +constexpr unsigned upper_range = 1 << 15; + +BENCHMARK_TEMPLATE(BM_sqrt_old, cpp_int)->RangeMultiplier(2)->Range(lower_range, upper_range)->Unit(benchmark::kMillisecond)->Complexity(); +BENCHMARK_TEMPLATE(BM_sqrt_current, cpp_int)->RangeMultiplier(2)->Range(lower_range, upper_range)->Unit(benchmark::kMillisecond)->Complexity(); +BENCHMARK_TEMPLATE(BM_sqrt_current, mpz_int)->RangeMultiplier(2)->Range(lower_range, upper_range)->Unit(benchmark::kMillisecond)->Complexity(); + +BENCHMARK_MAIN();