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:
jzmaddock
2021-06-18 20:05:02 +01:00
parent 019a169414
commit 8caae32acf
3 changed files with 268 additions and 69 deletions
@@ -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>
+36 -22
View File
@@ -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>
+144
View File
@@ -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();