/// @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 #include 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(mag); } inline std::int64_t half_of(std::int64_t raw) { const bool neg = raw < 0; auto mag = static_cast(neg ? -static_cast<__int128>(raw) : raw); const std::uint64_t bit = mag & 1u; mag >>= 1; if (bit) ++mag; if (mag > static_cast(INT64_MAX)) throw std::overflow_error("closed form: value does not fit int64"); const auto out = static_cast(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(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(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(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( std::max(eval_reduced(reduced::sqrt, fractional_bits, raw), 1)) << fractional_bits; const range_detail::u128 x = static_cast(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(raw), y), K); if (prod > static_cast(INT64_MAX)) throw std::overflow_error("closed form: power does not fit int64"); return static_cast(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(mag); while (bits > 1) { bits >>= 1; ++log; } const int exp = log - static_cast(fractional_bits); const int third = exp >= 0 ? exp / 3 : -(( -exp + 2) / 3); range_detail::u128 y = range_detail::u128{1} << static_cast(K + std::max(third, 0)); if (third < 0) y >>= static_cast(-third); if (y == 0) y = 1; const range_detail::u128 x = static_cast(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(INT64_MAX)) throw std::overflow_error("closed form: inverse cube root does not fit int64"); return static_cast(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(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__