mirror of
https://github.com/boostorg/multiprecision.git
synced 2026-07-21 13:23:49 +00:00
Minor tweaks to Karatsuba sqrt.
Hooks up backend-based version for cpp_int. Also adds simple google-benchmark. From https://github.com/boostorg/multiprecision/pull/328.
This commit is contained in:
@@ -45,6 +45,9 @@ void generic_interconvert(To& to, const From& from, const std::integral_constant
|
||||
template <class To, class From>
|
||||
void generic_interconvert(To& to, const From& from, const std::integral_constant<int, number_kind_rational>& /*to_type*/, const std::integral_constant<int, number_kind_integer>& /*from_type*/);
|
||||
|
||||
template <class Integer>
|
||||
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 <class B>
|
||||
void BOOST_MP_CXX14_CONSTEXPR eval_integer_sqrt(B& s, B& r, const B& x)
|
||||
template <class Backend>
|
||||
BOOST_MP_CXX14_CONSTEXPR void eval_qr(const Backend& x, const Backend& y, Backend& q, Backend& r);
|
||||
|
||||
template <class Backend>
|
||||
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<unsigned char, B>::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<Backend>::canonical_value(b);
|
||||
result = number<Backend>::canonical_value(c);
|
||||
return;
|
||||
}
|
||||
int g = eval_msb(x);
|
||||
if (g <= 1)
|
||||
#else
|
||||
if (bits <= std::numeric_limits<std::uintmax_t>::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<Backend>::canonical_value(b);
|
||||
result = number<Backend>::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 <class Backend>
|
||||
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 <class B>
|
||||
|
||||
@@ -189,52 +189,63 @@ BOOST_MP_CXX14_CONSTEXPR typename std::enable_if<boost::multiprecision::detail::
|
||||
return val;
|
||||
}
|
||||
|
||||
namespace detail {
|
||||
|
||||
template <class Integer>
|
||||
BOOST_MP_CXX14_CONSTEXPR typename std::enable_if<boost::multiprecision::detail::is_integral<Integer>::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<uint32_t>::max)());
|
||||
uint64_t val = static_cast<uint64_t>(x);
|
||||
uint64_t s64 = static_cast<uint64_t>(std::sqrt(static_cast<long double>(val)));
|
||||
if (bits <= 64)
|
||||
{
|
||||
const std::uint64_t int32max = std::uint64_t((std::numeric_limits<std::uint32_t>::max)());
|
||||
std::uint64_t val = static_cast<std::uint64_t>(x);
|
||||
std::uint64_t s64 = static_cast<std::uint64_t>(std::sqrt(static_cast<long double>(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::multiprecision::detail::
|
||||
s += q;
|
||||
q *= q;
|
||||
// we substract after, so it works for unsigned integers too
|
||||
if (r < q) {
|
||||
if (r < q)
|
||||
{
|
||||
t = s;
|
||||
t <<= 1;
|
||||
t--;
|
||||
@@ -264,6 +276,8 @@ BOOST_MP_CXX14_CONSTEXPR typename std::enable_if<boost::multiprecision::detail::
|
||||
return s;
|
||||
}
|
||||
|
||||
} // namespace detail
|
||||
|
||||
template <class Integer>
|
||||
BOOST_MP_CXX14_CONSTEXPR typename std::enable_if<boost::multiprecision::detail::is_integral<Integer>::value, Integer>::type sqrt(const Integer& x, Integer& r)
|
||||
{
|
||||
@@ -272,7 +286,7 @@ BOOST_MP_CXX14_CONSTEXPR typename std::enable_if<boost::multiprecision::detail::
|
||||
return 0u;
|
||||
}
|
||||
Integer t{};
|
||||
return karatsuba_sqrt(x, r, t, msb(x) + 1);
|
||||
return detail::karatsuba_sqrt(x, r, t, msb(x) + 1);
|
||||
}
|
||||
|
||||
template <class Integer>
|
||||
|
||||
@@ -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 <iostream>
|
||||
#include <benchmark/benchmark.h>
|
||||
#include <boost/multiprecision/cpp_int.hpp>
|
||||
#include <boost/multiprecision/gmp.hpp>
|
||||
#include <boost/multiprecision/integer.hpp>
|
||||
#include <boost/random.hpp>
|
||||
#include <cmath>
|
||||
|
||||
#include <immintrin.h>
|
||||
|
||||
using namespace boost::multiprecision;
|
||||
using namespace boost::random;
|
||||
|
||||
template <class Integer>
|
||||
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 <class Integer>
|
||||
BOOST_MP_CXX14_CONSTEXPR Integer sqrt_old(const Integer& x)
|
||||
{
|
||||
Integer r(0);
|
||||
return sqrt_old(x, r);
|
||||
}
|
||||
|
||||
template <class T>
|
||||
std::tuple<std::vector<T>, std::vector<T> >& get_test_vector(unsigned bits)
|
||||
{
|
||||
static std::map<unsigned, std::tuple<std::vector<T>, std::vector<T> > > data;
|
||||
|
||||
std::tuple<std::vector<T>, std::vector<T> >& result = data[bits];
|
||||
|
||||
if (std::get<0>(result).size() == 0)
|
||||
{
|
||||
mt19937 mt;
|
||||
uniform_int_distribution<T> ui(T(1) << (bits - 1), T(1) << bits);
|
||||
|
||||
std::vector<T>& a = std::get<0>(result);
|
||||
std::vector<T>& b = std::get<1>(result);
|
||||
|
||||
for (unsigned i = 0; i < 1000; ++i)
|
||||
{
|
||||
a.push_back(ui(mt));
|
||||
b.push_back(0);
|
||||
}
|
||||
}
|
||||
return result;
|
||||
}
|
||||
|
||||
template <class T>
|
||||
std::vector<T>& get_test_vector_a(unsigned bits)
|
||||
{
|
||||
return std::get<0>(get_test_vector<T>(bits));
|
||||
}
|
||||
template <class T>
|
||||
std::vector<T>& get_test_vector_b(unsigned bits)
|
||||
{
|
||||
return std::get<1>(get_test_vector<T>(bits));
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
static void BM_sqrt_old(benchmark::State& state)
|
||||
{
|
||||
int bits = state.range(0);
|
||||
|
||||
std::vector<T>& a = get_test_vector_a<T>(bits);
|
||||
std::vector<T>& b = get_test_vector_b<T>(bits);
|
||||
|
||||
for (auto _ : state)
|
||||
{
|
||||
for (unsigned i = 0; i < a.size(); ++i)
|
||||
b[i] = sqrt_old(a[i]);
|
||||
}
|
||||
state.SetComplexityN(bits);
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
static void BM_sqrt_current(benchmark::State& state)
|
||||
{
|
||||
int bits = state.range(0);
|
||||
|
||||
std::vector<T>& a = get_test_vector_a<T>(bits);
|
||||
std::vector<T>& b = get_test_vector_b<T>(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();
|
||||
Reference in New Issue
Block a user