382 lines
12 KiB
C++
382 lines
12 KiB
C++
/// @file grotto/principal_lut.hpp
|
|
/// @brief Principal-domain cubics at fractional precisions 8, 12, ..., 32.
|
|
/// @details Eleven maps reuse the elementary2 tournament knots (the
|
|
/// arbitrary-breakpoint cubics chosen among Mathematica, Maple, and
|
|
/// Chebfun). A requested precision only rounds those coefficients
|
|
/// down to `k + 16` fractional bits. `coth` is different: its
|
|
/// principal function depends on `k` through
|
|
/// `beta = ln(2^{k+1}+1)/2`, so each precision has its own
|
|
/// partition. `inv` (1/x), `rsqrt` (1/sqrt(x)), and `invsq`
|
|
/// (1/x^2) are also per precision: each is a longest-feasible
|
|
/// cubic march on the closed principal interval [1/2, 1], with
|
|
/// absolute error at most half an ulp at that precision. They do
|
|
/// not apply an exponent lift. The returned raw value is
|
|
/// `round_half_away(p(x) * 2^k)`. On the closed principal domain,
|
|
/// `p` stays within one unit in the last place of that precision.
|
|
|
|
#ifndef LIBDPF_INCLUDE_GROTTO_PRINCIPAL_LUT_HPP__
|
|
#define LIBDPF_INCLUDE_GROTTO_PRINCIPAL_LUT_HPP__
|
|
|
|
#include "hedley/hedley.h"
|
|
|
|
#include <cstdint>
|
|
#include <stdexcept>
|
|
|
|
namespace grotto
|
|
{
|
|
|
|
enum class principal : unsigned
|
|
{
|
|
ln = 0,
|
|
exp,
|
|
sin,
|
|
tanf,
|
|
tang,
|
|
sinh,
|
|
cosh,
|
|
sqrt,
|
|
coth,
|
|
sec,
|
|
gsec,
|
|
csch,
|
|
inv,
|
|
rsqrt,
|
|
invsq,
|
|
};
|
|
|
|
inline constexpr unsigned principal_precisions[] = {8u, 12u, 16u, 20u, 24u, 28u, 32u};
|
|
|
|
HEDLEY_NO_THROW
|
|
inline constexpr bool principal_precision(unsigned fractional_bits) noexcept
|
|
{
|
|
for (unsigned k : principal_precisions)
|
|
{
|
|
if (k == fractional_bits)
|
|
return true;
|
|
}
|
|
return false;
|
|
}
|
|
|
|
namespace principal_detail
|
|
{
|
|
|
|
struct cubic_bits
|
|
{
|
|
std::int64_t hi[4];
|
|
std::uint64_t lo[4];
|
|
};
|
|
|
|
struct knot
|
|
{
|
|
std::int64_t num;
|
|
std::uint8_t sh;
|
|
};
|
|
|
|
struct table_ref
|
|
{
|
|
const knot * knots;
|
|
const cubic_bits * pieces;
|
|
std::uint16_t nparts;
|
|
std::uint16_t q;
|
|
};
|
|
|
|
#include "grotto/principal_tables.inc"
|
|
#include "grotto/principal_recip_tables.inc"
|
|
|
|
struct w256
|
|
{
|
|
unsigned __int128 lo;
|
|
__int128 hi;
|
|
};
|
|
|
|
inline w256 w_from_i128(__int128 value)
|
|
{
|
|
w256 out;
|
|
out.lo = static_cast<unsigned __int128>(value);
|
|
out.hi = value < 0 ? __int128(-1) : __int128(0);
|
|
return out;
|
|
}
|
|
|
|
inline w256 w_add(w256 lhs, w256 rhs)
|
|
{
|
|
w256 out;
|
|
out.lo = lhs.lo + rhs.lo;
|
|
const unsigned __int128 carry = out.lo < lhs.lo ? 1u : 0u;
|
|
const unsigned __int128 hi = static_cast<unsigned __int128>(lhs.hi)
|
|
+ static_cast<unsigned __int128>(rhs.hi) + carry;
|
|
out.hi = static_cast<__int128>(hi);
|
|
return out;
|
|
}
|
|
|
|
inline w256 w_shl(w256 value, unsigned shift)
|
|
{
|
|
if (shift == 0)
|
|
return value;
|
|
unsigned __int128 hi = static_cast<unsigned __int128>(value.hi);
|
|
w256 out;
|
|
if (shift >= 128)
|
|
{
|
|
const unsigned off = shift - 128;
|
|
out.lo = 0;
|
|
out.hi = static_cast<__int128>(off >= 128 ? 0 : value.lo << off);
|
|
return out;
|
|
}
|
|
out.lo = value.lo << shift;
|
|
const unsigned __int128 spill = value.lo >> (128u - shift);
|
|
out.hi = static_cast<__int128>((hi << shift) | spill);
|
|
return out;
|
|
}
|
|
|
|
inline w256 w_mul_u64(w256 value, std::uint64_t factor)
|
|
{
|
|
const unsigned __int128 hi = static_cast<unsigned __int128>(value.hi);
|
|
const std::uint64_t limbs[4] = {
|
|
static_cast<std::uint64_t>(value.lo),
|
|
static_cast<std::uint64_t>(value.lo >> 64),
|
|
static_cast<std::uint64_t>(hi),
|
|
static_cast<std::uint64_t>(hi >> 64),
|
|
};
|
|
std::uint64_t out_limbs[4];
|
|
unsigned __int128 carry = 0;
|
|
for (int i = 0; i < 4; ++i)
|
|
{
|
|
const unsigned __int128 prod = static_cast<unsigned __int128>(limbs[i]) * factor + carry;
|
|
out_limbs[i] = static_cast<std::uint64_t>(prod);
|
|
carry = prod >> 64;
|
|
}
|
|
w256 out;
|
|
out.lo = static_cast<unsigned __int128>(out_limbs[0])
|
|
| (static_cast<unsigned __int128>(out_limbs[1]) << 64);
|
|
const unsigned __int128 new_hi = static_cast<unsigned __int128>(out_limbs[2])
|
|
| (static_cast<unsigned __int128>(out_limbs[3]) << 64);
|
|
out.hi = static_cast<__int128>(new_hi);
|
|
return out;
|
|
}
|
|
|
|
HEDLEY_NO_THROW
|
|
inline bool w_negative(w256 value) noexcept
|
|
{
|
|
return value.hi < 0;
|
|
}
|
|
|
|
inline w256 w_neg(w256 value)
|
|
{
|
|
unsigned __int128 lo = ~value.lo + 1;
|
|
unsigned __int128 hi = ~static_cast<unsigned __int128>(value.hi);
|
|
if (lo == 0)
|
|
hi += 1;
|
|
w256 out;
|
|
out.lo = lo;
|
|
out.hi = static_cast<__int128>(hi);
|
|
return out;
|
|
}
|
|
|
|
inline w256 w_add_pow2(w256 value, unsigned shift)
|
|
{
|
|
w256 add;
|
|
add.lo = 0;
|
|
add.hi = 0;
|
|
if (shift < 128)
|
|
add.lo = static_cast<unsigned __int128>(1) << shift;
|
|
else if (shift < 256)
|
|
add.hi = static_cast<__int128>(static_cast<unsigned __int128>(1) << (shift - 128));
|
|
return w_add(value, add);
|
|
}
|
|
|
|
inline w256 w_shr(w256 value, unsigned shift)
|
|
{
|
|
if (shift == 0)
|
|
return value;
|
|
const unsigned __int128 hi = static_cast<unsigned __int128>(value.hi);
|
|
w256 out;
|
|
if (shift >= 256)
|
|
{
|
|
out.lo = 0;
|
|
out.hi = 0;
|
|
return out;
|
|
}
|
|
if (shift >= 128)
|
|
{
|
|
out.hi = 0;
|
|
out.lo = hi >> (shift - 128);
|
|
return out;
|
|
}
|
|
out.lo = (value.lo >> shift) | (hi << (128u - shift));
|
|
out.hi = static_cast<__int128>(hi >> shift);
|
|
return out;
|
|
}
|
|
|
|
inline std::int64_t w_to_i64(w256 value)
|
|
{
|
|
const unsigned __int128 hi = static_cast<unsigned __int128>(value.hi);
|
|
if (hi == 0)
|
|
{
|
|
if (value.lo > static_cast<unsigned __int128>(INT64_MAX))
|
|
throw std::overflow_error("principal lut: value does not fit int64");
|
|
return static_cast<std::int64_t>(value.lo);
|
|
}
|
|
if (value.hi == __int128(-1))
|
|
{
|
|
const auto hi_lo = static_cast<std::uint64_t>(value.lo >> 64);
|
|
if (hi_lo != ~std::uint64_t{0})
|
|
throw std::overflow_error("principal lut: value does not fit int64");
|
|
return static_cast<std::int64_t>(static_cast<std::uint64_t>(value.lo));
|
|
}
|
|
throw std::overflow_error("principal lut: value does not fit int64");
|
|
}
|
|
|
|
inline std::int64_t round_half_away_pow2(w256 numerator, unsigned shift)
|
|
{
|
|
if (shift == 0)
|
|
return w_to_i64(numerator);
|
|
const bool neg = w_negative(numerator);
|
|
w256 mag = neg ? w_neg(numerator) : numerator;
|
|
mag = w_add_pow2(mag, shift - 1);
|
|
mag = w_shr(mag, shift);
|
|
if (neg)
|
|
mag = w_neg(mag);
|
|
return w_to_i64(mag);
|
|
}
|
|
|
|
inline __int128 unpack_coeff(std::int64_t hi, std::uint64_t lo)
|
|
{
|
|
__int128 value = hi;
|
|
value <<= 64;
|
|
value |= static_cast<__int128>(lo);
|
|
return value;
|
|
}
|
|
|
|
inline __int128 rshift_ties_even(__int128 value, int shift)
|
|
{
|
|
if (shift <= 0)
|
|
return value;
|
|
const bool neg = value < 0;
|
|
const auto mag = static_cast<unsigned __int128>(neg ? -value : value);
|
|
const auto half = static_cast<unsigned __int128>(1) << (shift - 1);
|
|
const auto mask = (static_cast<unsigned __int128>(1) << shift) - 1;
|
|
unsigned __int128 quot = mag >> shift;
|
|
const unsigned __int128 rem = mag & mask;
|
|
if (rem > half || (rem == half && (quot & 1u)))
|
|
++quot;
|
|
const auto out = static_cast<__int128>(quot);
|
|
return neg ? -out : out;
|
|
}
|
|
|
|
inline bool knot_le_raw(knot endpoint, std::int64_t raw, unsigned fractional_bits)
|
|
{
|
|
const __int128 left = static_cast<__int128>(endpoint.num) << fractional_bits;
|
|
const __int128 right = static_cast<__int128>(raw) << endpoint.sh;
|
|
return left <= right;
|
|
}
|
|
|
|
inline int piece_index(const table_ref & table, std::int64_t raw, unsigned fractional_bits)
|
|
{
|
|
if (!knot_le_raw(table.knots[0], raw, fractional_bits))
|
|
throw std::out_of_range("principal lut: input is below the principal domain");
|
|
const knot stop = table.knots[table.nparts];
|
|
const __int128 stop_scaled = static_cast<__int128>(stop.num) << fractional_bits;
|
|
const __int128 raw_scaled = static_cast<__int128>(raw) << stop.sh;
|
|
if (raw_scaled > stop_scaled)
|
|
throw std::out_of_range("principal lut: input is above the principal domain");
|
|
|
|
int lo = 0;
|
|
int hi = static_cast<int>(table.nparts);
|
|
while (hi - lo > 1)
|
|
{
|
|
const int mid = (lo + hi) / 2;
|
|
if (knot_le_raw(table.knots[mid], raw, fractional_bits))
|
|
lo = mid;
|
|
else
|
|
hi = mid;
|
|
}
|
|
return lo;
|
|
}
|
|
|
|
inline const table_ref & table_for(principal which, unsigned fractional_bits)
|
|
{
|
|
const auto index = static_cast<unsigned>(which);
|
|
if (which == principal::coth)
|
|
{
|
|
const unsigned slot = fractional_bits / 4u - 2u;
|
|
return *COTH_BY_K[slot];
|
|
}
|
|
if (which == principal::inv || which == principal::rsqrt || which == principal::invsq)
|
|
{
|
|
const unsigned slot = fractional_bits / 4u - 2u;
|
|
const table_ref * const * bank = which == principal::inv
|
|
? INV_BY_K
|
|
: which == principal::rsqrt ? RSQRT_BY_K : INVSQ_BY_K;
|
|
return *bank[slot];
|
|
}
|
|
return *SHARED_TABLE[index];
|
|
}
|
|
|
|
inline w256 w_mul_i64(w256 value, std::int64_t factor)
|
|
{
|
|
if (factor >= 0)
|
|
return w_mul_u64(value, static_cast<std::uint64_t>(factor));
|
|
// `-factor` is undefined at INT64_MIN. The magnitude is the unsigned wrap.
|
|
const auto mag = static_cast<std::uint64_t>(0)
|
|
- static_cast<std::uint64_t>(factor);
|
|
return w_neg(w_mul_u64(value, mag));
|
|
}
|
|
|
|
inline std::int64_t horner(const cubic_bits & piece, unsigned q, std::int64_t raw, unsigned fractional_bits)
|
|
{
|
|
__int128 coeff[4];
|
|
for (int i = 0; i < 4; ++i)
|
|
coeff[i] = unpack_coeff(piece.hi[i], piece.lo[i]);
|
|
const int q_use = static_cast<int>(q) < static_cast<int>(fractional_bits) + 16
|
|
? static_cast<int>(q)
|
|
: static_cast<int>(fractional_bits) + 16;
|
|
const int drop = static_cast<int>(q) - q_use;
|
|
for (int i = 0; i < 4; ++i)
|
|
coeff[i] = rshift_ties_even(coeff[i], drop);
|
|
|
|
// S3 / 2^{q_use + 3k} = p(x), so round(p * 2^k) = round(S3 / 2^{q_use + 2k}).
|
|
w256 acc = w_from_i128(coeff[3]);
|
|
for (int i = 2; i >= 0; --i)
|
|
{
|
|
acc = w_mul_i64(acc, raw);
|
|
w256 term = w_shl(w_from_i128(coeff[i]), fractional_bits * static_cast<unsigned>(3 - i));
|
|
acc = w_add(acc, term);
|
|
}
|
|
const unsigned denom_shift = static_cast<unsigned>(q_use) + 2u * fractional_bits;
|
|
return round_half_away_pow2(acc, denom_shift);
|
|
}
|
|
|
|
} // namespace principal_detail
|
|
|
|
/// @brief Number of cubic pieces at this fractional width.
|
|
/// @param which the party index
|
|
/// @param fractional_bits the number of fractional bits
|
|
/// @return Number of cubic pieces at this fractional width
|
|
/// @throws std::invalid_argument if `precision must be 8, 12, ..., 32`
|
|
inline std::uint16_t principal_parts(principal which, unsigned fractional_bits)
|
|
{
|
|
if (!principal_precision(fractional_bits))
|
|
throw std::invalid_argument("principal lut: precision must be 8, 12, ..., 32");
|
|
return principal_detail::table_for(which, fractional_bits).nparts;
|
|
}
|
|
|
|
/// @brief Evaluate the principal-domain cubic. `raw / 2^fractional_bits` is the input.
|
|
/// @param which the party index
|
|
/// @param fractional_bits the number of fractional bits
|
|
/// @param raw the underlying integer
|
|
/// @return the returned `std::int64_t`
|
|
/// @throws std::invalid_argument if `precision must be 8, 12, ..., 32`
|
|
/// @throws std::out_of_range if `input is below the principal domain`
|
|
inline std::int64_t eval_principal(principal which, unsigned fractional_bits, std::int64_t raw)
|
|
{
|
|
if (!principal_precision(fractional_bits))
|
|
throw std::invalid_argument("principal lut: precision must be 8, 12, ..., 32");
|
|
if (raw < 0)
|
|
throw std::out_of_range("principal lut: input is below the principal domain");
|
|
const principal_detail::table_ref & table = principal_detail::table_for(which, fractional_bits);
|
|
const int index = principal_detail::piece_index(table, raw, fractional_bits);
|
|
return principal_detail::horner(table.pieces[index], table.q, raw, fractional_bits);
|
|
}
|
|
|
|
} // namespace grotto
|
|
|
|
#endif // LIBDPF_INCLUDE_GROTTO_PRINCIPAL_LUT_HPP__
|