libdpf/include/grotto/closed_form.hpp

519 lines
19 KiB
C++
Raw Normal View History

/// @file grotto/closed_form.hpp
/// @brief Closed forms of the principal, range, and window maps.
/// @details Inverse hyperbolics, inverse trig, SELU / ELU / CELU, softsign,
/// tanhshrink, the logistic / exponential / Laplace / Cauchy
/// quantiles, sinc, and the extra powers are compositions of
/// `eval_reduced` and `eval_window`. `atan` on `[0, tan(π/8)]` is a
/// odd power series; larger arguments reduce by `π/4` and `π/2`.
/// Roots and `x^{-0.1}` go through `ln` and `exp`. `x^{1.5}` is
/// `x √x` and `x^{-3}` is the reciprocal of an exact cube.
/// Precision is one of 8, 12, ..., 32. These compositions inherit
/// the range-reduction error, so they are not a 1-ulp claim.
#ifndef LIBDPF_INCLUDE_GROTTO_CLOSED_FORM_HPP__
#define LIBDPF_INCLUDE_GROTTO_CLOSED_FORM_HPP__
#include "grotto/range_lut.hpp"
#include "grotto/window_lut.hpp"
#include <cstdint>
#include <stdexcept>
namespace grotto
{
enum class closed : unsigned
{
atanh = 0,
asinh,
acosh,
atan,
acot,
asec,
acsc,
asech,
acsch,
acoth,
selu,
elu,
celu,
softsign,
tanhshrink,
logistic,
exponential,
laplace,
cauchy,
sinc,
cbrt,
qtrt,
icbrt,
iqtrt,
pow_m01,
pow_p15,
pow_m3,
};
/// @brief Evaluate one closed form at a fixed-point raw value.
/// @param which the closed-form function
/// @param fractional_bits precision in {8, 12, ..., 32}
/// @param raw the fixed-point argument
/// @return the fixed-point result
/// @throws std::invalid_argument if the precision is not a principal step
/// @throws std::domain_error at a pole or outside the function's domain
/// \complexity A constant number of `eval_reduced` / `eval_window` calls, plus the loops in this file:
/// `atan_series` runs `n = 1 .. 47` (stops on a zero term), `sinc_series` runs `n = 1 .. 16`, `newton_sqrt_raw` runs 4 steps, `newton_icbrt_raw` runs 6 steps after a right-shift log of the magnitude.
/// Extra space `Θ(1)`. These compositions are not a 1-ulp claim; the file comment says they inherit range-reduction error.
/// @see grotto::eval_reduced
/// @see grotto::eval_window
/// @note Not a 1-ulp claim. Poles and domain exits throw `std::domain_error`.
inline std::int64_t eval_closed(closed which, unsigned fractional_bits, std::int64_t raw);
namespace closed_detail
{
using range_detail::div_raw;
using range_detail::mul_raw;
using range_detail::one_raw;
using range_detail::scale_unit;
using range_detail::u128;
inline constexpr u128 selu_alpha_64 = (u128{1} << 64) | 12419514725947086766ULL;
inline constexpr u128 selu_scale_64 = (u128{1} << 64) | 935268138030932704ULL;
inline constexpr u128 tenth_64 = 1844674407370955162ULL;
inline constexpr u128 three_halves_64 = (u128{1} << 64) | 9223372036854775808ULL;
inline std::int64_t require_precision(unsigned fractional_bits)
{
if (!principal_precision(fractional_bits))
throw std::invalid_argument("closed form: precision must be 8, 12, ..., 32");
return one_raw(fractional_bits);
}
inline std::int64_t abs_raw(std::int64_t raw)
{
if (raw >= 0)
return raw;
const auto mag = -static_cast<__int128>(raw);
if (mag > INT64_MAX)
throw std::overflow_error("closed form: magnitude does not fit int64");
return static_cast<std::int64_t>(mag);
}
inline std::int64_t half_of(std::int64_t raw)
{
const bool neg = raw < 0;
auto mag = static_cast<std::uint64_t>(neg ? -static_cast<__int128>(raw) : raw);
const std::uint64_t bit = mag & 1u;
mag >>= 1;
if (bit)
++mag;
if (mag > static_cast<std::uint64_t>(INT64_MAX))
throw std::overflow_error("closed form: value does not fit int64");
const auto out = static_cast<std::int64_t>(mag);
return neg ? -out : out;
}
inline std::int64_t div_int(std::int64_t raw, int denominator)
{
if (denominator <= 0)
throw std::invalid_argument("closed form: divisor must be positive");
const bool neg = raw < 0;
auto mag = static_cast<__int128>(neg ? -static_cast<__int128>(raw) : raw);
const __int128 den = denominator;
__int128 quot = mag / den;
if ((mag % den) * 2 >= den)
++quot;
if (quot > INT64_MAX)
throw std::overflow_error("closed form: value does not fit int64");
const auto out = static_cast<std::int64_t>(quot);
return neg ? -out : out;
}
inline std::int64_t pi_over_4(unsigned fractional_bits)
{
return scale_unit(range_detail::pi_over_4_64, fractional_bits);
}
inline std::int64_t pi_over_2(unsigned fractional_bits)
{
const __int128 wide = static_cast<__int128>(pi_over_4(fractional_bits)) << 1;
if (wide > INT64_MAX)
throw std::overflow_error("closed form: pi/2 does not fit");
return static_cast<std::int64_t>(wide);
}
inline std::int64_t atan_series(unsigned fractional_bits, std::int64_t x)
{
__int128 acc = x;
std::int64_t power = x;
const std::int64_t x2 = mul_raw(x, x, fractional_bits);
int sign = -1;
for (int n = 1; n < 48; ++n)
{
power = mul_raw(power, x2, fractional_bits);
const std::int64_t term = div_int(power, 2 * n + 1);
if (term == 0)
break;
acc += sign < 0 ? -static_cast<__int128>(term) : term;
sign = -sign;
}
if (acc > INT64_MAX || acc < INT64_MIN)
throw std::overflow_error("closed form: atan does not fit");
return static_cast<std::int64_t>(acc);
}
inline std::int64_t atan_positive(unsigned fractional_bits, std::int64_t mag)
{
const std::int64_t one = one_raw(fractional_bits);
if (mag > one)
return pi_over_2(fractional_bits) - atan_positive(fractional_bits, div_raw(one, mag, fractional_bits));
const std::int64_t bound = scale_unit(range_detail::sqrt2_64, fractional_bits) - one;
if (mag > bound)
{
const std::int64_t t = div_raw(mag - one, mag + one, fractional_bits);
return pi_over_4(fractional_bits) + atan_series(fractional_bits, t);
}
return atan_series(fractional_bits, mag);
}
inline std::int64_t eval_atan(unsigned fractional_bits, std::int64_t raw)
{
require_precision(fractional_bits);
if (raw < 0)
return -atan_positive(fractional_bits, abs_raw(raw));
return atan_positive(fractional_bits, raw);
}
inline std::int64_t eval_ln_abs(unsigned fractional_bits, std::int64_t raw)
{
return eval_reduced(reduced::ln, fractional_bits, abs_raw(raw));
}
inline std::int64_t eval_cbrt_abs(unsigned fractional_bits, std::int64_t mag)
{
if (mag == 0)
return 0;
const std::int64_t ln = eval_ln_abs(fractional_bits, mag);
return eval_reduced(reduced::exp, fractional_bits, div_int(ln, 3));
}
inline range_detail::u256 shl_u128(range_detail::u128 value, unsigned shift)
{
if (shift == 0)
return range_detail::u256{value, 0};
if (shift < 128)
return range_detail::u256{value << shift, value >> (128u - shift)};
return range_detail::u256{0, value << (shift - 128u)};
}
inline bool u256_less(range_detail::u256 a, range_detail::u256 b)
{
if (a.hi != b.hi)
return a.hi < b.hi;
return a.lo < b.lo;
}
/// @brief Rounded `(num << shift) / den`.
inline range_detail::u128 div_shifted(range_detail::u128 num, range_detail::u128 den, unsigned shift)
{
if (den == 0)
throw std::domain_error("closed form: division by zero");
const range_detail::u256 target = shl_u128(num, shift);
range_detail::u128 lo = 0;
range_detail::u128 hi = 1;
while (u256_less(range_detail::mul_u128(hi, den), target) ||
(!u256_less(target, range_detail::mul_u128(hi, den)) && hi < (range_detail::u128{1} << 80)))
{
if (hi > (range_detail::u128{1} << 100))
break;
hi <<= 1;
}
while (lo + 1 < hi)
{
const range_detail::u128 mid = lo + (hi - lo) / 2;
if (u256_less(target, range_detail::mul_u128(mid, den)))
hi = mid;
else
lo = mid;
}
const range_detail::u256 half = range_detail::mul_u128(lo + lo + 1, den);
if (!u256_less(target, half))
return lo + 1;
return lo;
}
inline std::int64_t newton_sqrt_raw(unsigned fractional_bits, std::int64_t raw)
{
const unsigned K = fractional_bits * 2u;
range_detail::u128 y = static_cast<range_detail::u128>(
std::max<std::int64_t>(eval_reduced(reduced::sqrt, fractional_bits, raw), 1))
<< fractional_bits;
const range_detail::u128 x = static_cast<range_detail::u128>(raw) << fractional_bits;
for (int step = 0; step < 4; ++step)
{
const range_detail::u128 quot = div_shifted(x, y, K);
y = (y + quot + 1) >> 1;
}
const range_detail::u128 prod = range_detail::round_u256(
range_detail::mul_u128(static_cast<range_detail::u128>(raw), y), K);
if (prod > static_cast<range_detail::u128>(INT64_MAX))
throw std::overflow_error("closed form: power does not fit int64");
return static_cast<std::int64_t>(prod);
}
inline std::int64_t newton_icbrt_raw(unsigned fractional_bits, std::int64_t mag)
{
const unsigned K = fractional_bits * 2u;
int log = 0;
auto bits = static_cast<std::uint64_t>(mag);
while (bits > 1)
{
bits >>= 1;
++log;
}
const int exp = log - static_cast<int>(fractional_bits);
const int third = exp >= 0 ? exp / 3 : -(( -exp + 2) / 3);
range_detail::u128 y = range_detail::u128{1} << static_cast<unsigned>(K + std::max(third, 0));
if (third < 0)
y >>= static_cast<unsigned>(-third);
if (y == 0)
y = 1;
const range_detail::u128 x = static_cast<range_detail::u128>(mag) << fractional_bits;
for (int step = 0; step < 6; ++step)
{
const range_detail::u128 y2 = range_detail::round_u256(range_detail::mul_u128(y, y), K);
if (y2 == 0)
break;
const range_detail::u128 quot = div_shifted(x, y2, K);
y = (y + y + quot) / 3;
if (y == 0)
y = 1;
}
const range_detail::u128 numer = range_detail::u128{1} << (K + fractional_bits);
const range_detail::u128 inv = (numer + y / 2) / y;
if (inv > static_cast<range_detail::u128>(INT64_MAX))
throw std::overflow_error("closed form: inverse cube root does not fit int64");
return static_cast<std::int64_t>(inv);
}
inline std::int64_t sinc_series(unsigned fractional_bits, std::int64_t raw)
{
constexpr unsigned extra = 16;
__int128 acc = __int128{1} << (fractional_bits + extra);
__int128 term = acc;
const std::int64_t mag = raw < 0 ? -raw : raw;
for (int n = 1; n <= 16; ++n)
{
term = range_detail::shr_round_i128(term * mag, fractional_bits);
term = range_detail::shr_round_i128(term * mag, fractional_bits);
term = range_detail::div_round_i128(term, (2 * n) * (2 * n + 1));
if (term == 0)
break;
acc += (n % 2) != 0 ? -term : term;
}
return range_detail::round_i128(acc, extra);
}
inline std::int64_t square_plus(unsigned fractional_bits, std::int64_t raw, int sign)
{
const std::int64_t one = one_raw(fractional_bits);
const std::int64_t sq = mul_raw(raw, raw, fractional_bits);
const __int128 sum = static_cast<__int128>(sq) + (sign < 0 ? -one : one);
if (sum < 0)
throw std::domain_error("closed form: square is below the root domain");
if (sum > INT64_MAX)
throw std::overflow_error("closed form: square does not fit");
return eval_reduced(reduced::sqrt, fractional_bits, static_cast<std::int64_t>(sum));
}
} // namespace closed_detail
/// \complexity A constant number of `eval_reduced` / `eval_window` calls, plus the loops in this file:
/// `atan_series` runs `n = 1 .. 47` (stops on a zero term), `sinc_series` runs `n = 1 .. 16`, `newton_sqrt_raw` runs 4 steps, `newton_icbrt_raw` runs 6 steps after a right-shift log of the magnitude.
/// Extra space `Θ(1)`. These compositions are not a 1-ulp claim; the file comment says they inherit range-reduction error.
/// @see grotto::eval_reduced
/// @see grotto::eval_window
/// @note Not a 1-ulp claim. Poles and domain exits throw `std::domain_error`.
inline std::int64_t eval_closed(closed which, unsigned fractional_bits, std::int64_t raw)
{
using namespace closed_detail;
const std::int64_t one = require_precision(fractional_bits);
switch (which)
{
case closed::atan:
return eval_atan(fractional_bits, raw);
case closed::acot:
return pi_over_2(fractional_bits) - eval_atan(fractional_bits, raw);
case closed::asec:
case closed::acsc:
{
if (abs_raw(raw) < one)
throw std::domain_error("closed form: inverse secant requires |x| >= 1");
const std::int64_t inv = div_raw(one, raw, fractional_bits);
return which == closed::asec
? eval_window(window::acos, fractional_bits, inv)
: eval_window(window::asin, fractional_bits, inv);
}
case closed::atanh:
{
if (abs_raw(raw) >= one)
throw std::domain_error("closed form: atanh domain is (-1, 1)");
const std::int64_t plus = eval_reduced(reduced::log1p, fractional_bits, raw);
const std::int64_t minus = eval_reduced(reduced::log1p, fractional_bits, -raw);
return half_of(plus - minus);
}
case closed::asinh:
{
const std::int64_t mag = abs_raw(raw);
std::int64_t sum = 0;
try
{
sum = mag + square_plus(fractional_bits, mag, +1);
}
catch (const std::overflow_error &)
{
sum = 0;
}
const std::int64_t value = sum == 0
? eval_ln_abs(fractional_bits, mag) + range_detail::ln2_raw(fractional_bits)
: eval_reduced(reduced::ln, fractional_bits, sum);
return raw < 0 ? -value : value;
}
case closed::acosh:
{
if (raw < one)
throw std::domain_error("closed form: acosh domain is [1, inf)");
return eval_reduced(reduced::ln, fractional_bits,
raw + square_plus(fractional_bits, raw, -1));
}
case closed::asech:
{
if (raw <= 0 || raw > one)
throw std::domain_error("closed form: asech domain is (0, 1]");
return eval_closed(closed::acosh, fractional_bits, div_raw(one, raw, fractional_bits));
}
case closed::acsch:
{
if (raw == 0)
throw std::domain_error("closed form: acsch pole");
return eval_closed(closed::asinh, fractional_bits, div_raw(one, raw, fractional_bits));
}
case closed::acoth:
{
if (abs_raw(raw) <= one)
throw std::domain_error("closed form: acoth domain is |x| > 1");
return eval_closed(closed::atanh, fractional_bits, div_raw(one, raw, fractional_bits));
}
case closed::elu:
return raw > 0 ? raw : eval_reduced(reduced::expm1, fractional_bits, raw);
case closed::celu:
return eval_closed(closed::elu, fractional_bits, raw);
case closed::selu:
{
const std::int64_t body = raw > 0
? raw
: mul_raw(scale_unit(selu_alpha_64, fractional_bits),
eval_reduced(reduced::expm1, fractional_bits, raw), fractional_bits);
return mul_raw(scale_unit(selu_scale_64, fractional_bits), body, fractional_bits);
}
case closed::softsign:
{
const std::int64_t mag = abs_raw(raw);
return div_raw(raw, one + mag, fractional_bits);
}
case closed::tanhshrink:
return raw - eval_window(window::tanh, fractional_bits, raw);
case closed::logistic:
{
if (raw <= 0 || raw >= one)
throw std::domain_error("closed form: quantile domain is (0, 1)");
const std::int64_t num = eval_reduced(reduced::ln, fractional_bits, raw);
const std::int64_t den = eval_reduced(reduced::ln, fractional_bits, one - raw);
return num - den;
}
case closed::exponential:
{
if (raw <= 0 || raw >= one)
throw std::domain_error("closed form: quantile domain is (0, 1)");
return -eval_reduced(reduced::ln, fractional_bits, one - raw);
}
case closed::laplace:
{
if (raw <= 0 || raw >= one)
throw std::domain_error("closed form: quantile domain is (0, 1)");
const std::int64_t half = one >> 1;
if (raw <= half)
return eval_reduced(reduced::ln, fractional_bits, raw << 1);
return -eval_reduced(reduced::ln, fractional_bits, (one - raw) << 1);
}
case closed::cauchy:
{
if (raw <= 0 || raw >= one)
throw std::domain_error("closed form: quantile domain is (0, 1)");
const std::int64_t half = one >> 1;
const std::int64_t pi = pi_over_2(fractional_bits) << 1;
const std::int64_t angle = mul_raw(pi, raw - half, fractional_bits);
return eval_reduced(reduced::tan, fractional_bits, angle);
}
case closed::sinc:
if (raw == 0)
return one;
if (abs_raw(raw) <= one)
return sinc_series(fractional_bits, raw);
return div_raw(eval_reduced(reduced::sin, fractional_bits, raw), raw, fractional_bits);
case closed::cbrt:
case closed::icbrt:
{
if (which == closed::icbrt)
{
if (raw == 0)
throw std::domain_error("closed form: inverse cube root of zero");
const std::int64_t value = newton_icbrt_raw(fractional_bits, abs_raw(raw));
return raw < 0 ? -value : value;
}
const std::int64_t root = eval_cbrt_abs(fractional_bits, abs_raw(raw));
return raw < 0 ? -root : root;
}
case closed::qtrt:
case closed::iqtrt:
{
if (raw < 0)
throw std::domain_error("closed form: fourth root requires x >= 0");
if (raw == 0)
return which == closed::qtrt ? 0 : throw std::domain_error("closed form: inverse fourth root of zero"), 0;
const std::int64_t ln = eval_ln_abs(fractional_bits, raw);
const std::int64_t root = eval_reduced(reduced::exp, fractional_bits, div_int(ln, 4));
return which == closed::qtrt ? root : div_raw(one, root, fractional_bits);
}
case closed::pow_m01:
{
if (raw <= 0)
throw std::domain_error("closed form: x^{-0.1} requires x > 0");
const std::int64_t ln = eval_ln_abs(fractional_bits, raw);
const std::int64_t scaled = mul_raw(ln, scale_unit(tenth_64, fractional_bits), fractional_bits);
return eval_reduced(reduced::exp, fractional_bits, -scaled);
}
case closed::pow_p15:
{
if (raw < 0)
throw std::domain_error("closed form: x^{1.5} requires x >= 0");
if (raw == 0)
return 0;
return newton_sqrt_raw(fractional_bits, raw);
}
case closed::pow_m3:
{
if (raw == 0)
throw std::domain_error("closed form: x^{-3} pole");
const std::int64_t sq = mul_raw(raw, raw, fractional_bits);
const std::int64_t cube = mul_raw(sq, raw, fractional_bits);
return div_raw(one, cube, fractional_bits);
}
}
throw std::invalid_argument("closed form: unknown map");
}
} // namespace grotto
#endif // LIBDPF_INCLUDE_GROTTO_CLOSED_FORM_HPP__