/// @file grotto/window_lut.hpp /// @brief Direct cubics for Grotto maps that do not want mantissa reduction. /// @details Sollya minimax cubics cover the bend. Outside it the value is an /// exact tail: 0, ±1, or the identity. `smoothstep` is the exact /// cubic. `erfc`, `softminus`, `logsigmoid`, and `acos` are integer /// rewrites of `erf`, `softplus`, and `asin`. `asin` on `(1/2, 1]` /// uses `π/2 − 2 asin(sqrt((1−x)/2))` with the principal square-root /// table. `probit` is stored on `(0, 1/2]` and mirrored. The tail /// below 1/20 is a cubic in `ln(p)` (knots at scale `2^{k+10}`), /// because a cubic in `p` cannot meet half an ulp on the first /// input step once `k` is large. #ifndef LIBDPF_INCLUDE_GROTTO_WINDOW_LUT_HPP__ #define LIBDPF_INCLUDE_GROTTO_WINDOW_LUT_HPP__ #include #include #include "grotto/principal_lut.hpp" namespace grotto { enum class window : unsigned { smoothstep = 0, sigmoid, tanh, erf, erfc, softplus, softminus, logsigmoid, gelu, silu, mish, elish, serf, tanhexp, asin, acos, probit, }; namespace window_detail { using grotto::principal_detail::cubic_bits; using grotto::principal_detail::horner; struct window_table { const std::int64_t * knots; const cubic_bits * pieces; std::uint16_t nparts; std::uint16_t q; }; #include "grotto/window_tables.inc" inline unsigned slot_of(unsigned fractional_bits) { return fractional_bits / 4u - 2u; } inline const window_table & at(window_table const * const * tables, unsigned fractional_bits) { return *tables[slot_of(fractional_bits)]; } inline int piece_of(const window_table & table, std::int64_t raw) { int lo = 0; int hi = static_cast(table.nparts); while (hi - lo > 1) { const int mid = (lo + hi) / 2; if (table.knots[mid] <= raw) lo = mid; else hi = mid; } return lo; } inline std::int64_t eval_table(const window_table & table, unsigned fractional_bits, std::int64_t raw) { if (raw < table.knots[0] || raw > table.knots[table.nparts]) throw std::out_of_range("window lut: input is outside this piece table"); return horner(table.pieces[piece_of(table, raw)], table.q, raw, fractional_bits); } enum tail_kind { tail_zero = 0, tail_one = 1, tail_neg = 2, tail_id = 3 }; inline std::int64_t apply_tail(tail_kind kind, unsigned fractional_bits, std::int64_t raw) { switch (kind) { case tail_zero: return 0; case tail_one: return std::int64_t{1} << fractional_bits; case tail_neg: return -(std::int64_t{1} << fractional_bits); case tail_id: return raw; } throw std::invalid_argument("window lut: bad tail"); } inline std::int64_t eval_tailed( window_table const * const * tables, tail_kind lo, tail_kind hi, unsigned fractional_bits, std::int64_t raw) { const window_table & table = at(tables, fractional_bits); if (raw < table.knots[0]) return apply_tail(lo, fractional_bits, raw); if (raw > table.knots[table.nparts]) return apply_tail(hi, fractional_bits, raw); return eval_table(table, fractional_bits, raw); } inline std::int64_t round_half_away_i128(__int128 number, unsigned shift) { if (shift == 0) { if (number > INT64_MAX || number < INT64_MIN) throw std::overflow_error("window lut: value does not fit int64"); return static_cast(number); } const bool neg = number < 0; const auto mag = static_cast(neg ? -number : number); const unsigned __int128 quot = (mag + (static_cast(1) << (shift - 1))) >> shift; const auto out = static_cast<__int128>(quot); return static_cast(neg ? -out : out); } inline unsigned __int128 isqrt_floor(unsigned __int128 n) { if (n == 0) return 0; const unsigned bits = (n >> 64) != 0 ? 128u - static_cast(__builtin_clzll(static_cast(n >> 64))) : 64u - static_cast(__builtin_clzll(static_cast(n))); unsigned __int128 x = static_cast(1) << ((bits + 1u) / 2u); for (;;) { const unsigned __int128 y = (x + n / x) >> 1; if (y >= x) break; x = y; } while (x > 0 && x > n / x) --x; return x; } /// `round(sqrt(v / 2^{k+1}) * 2^{k+extra})`, `v > 0`. Eight extra bits so the /// half-angle identity can absorb the square root before the final rounding. inline std::int64_t sqrt_half_scale_fine(unsigned fractional_bits, std::int64_t magnitude) { constexpr unsigned extra = 8; const unsigned shift = fractional_bits + 2u * extra; const unsigned __int128 radicand = static_cast(static_cast(magnitude)) << shift; const unsigned __int128 root = isqrt_floor(radicand); // sqrt(gap << (k+2*extra)) / sqrt(2) = sqrt(gap / 2^{k+1}) * 2^{k+extra} static constexpr unsigned __int128 sqrt2_64 = (static_cast(1) << 64) | static_cast(7640891576956012809ULL); const unsigned __int128 scaled = (root * sqrt2_64 + (static_cast(1) << 64)) >> 65; return static_cast(scaled); } inline int piece_of_scaled(const window_table & table, std::int64_t raw, unsigned extra) { int lo = 0; int hi = static_cast(table.nparts); while (hi - lo > 1) { const int mid = (lo + hi) / 2; if ((table.knots[mid] << extra) <= raw) lo = mid; else hi = mid; } return lo; } inline std::int64_t eval_asin_abs(unsigned fractional_bits, std::int64_t magnitude) { constexpr unsigned extra = 8; const auto half = std::int64_t{1} << (fractional_bits - 1); const window_table & table = at(ASIN, fractional_bits); if (magnitude <= half) return eval_table(table, fractional_bits, magnitude); const std::int64_t one = std::int64_t{1} << fractional_bits; if (magnitude >= one) return HALF_PI_RAW[slot_of(fractional_bits)]; const std::int64_t gap = one - magnitude; std::int64_t reduced = sqrt_half_scale_fine(fractional_bits, gap); const std::int64_t half_fine = half << extra; if (reduced > half_fine) reduced = half_fine; const unsigned scale = fractional_bits + extra; const std::int64_t inner = horner( table.pieces[piece_of_scaled(table, reduced, extra)], table.q, reduced, scale); // pi/2 at 64 fractional bits, then onto scale k+extra in one rounding. static constexpr unsigned __int128 half_pi_64 = (static_cast(1) << 64) | static_cast(10529333758598939754ULL); const __int128 pi_fine = round_half_away_i128( static_cast<__int128>(half_pi_64), 64u - scale); const __int128 lifted = pi_fine - 2 * static_cast<__int128>(inner); const auto out = round_half_away_i128(lifted, extra); return out < 0 ? 0 : out; } /// Surplus fractional bits on probit-tail knots. `u = ln(p)` is stored as /// `round(u * 2^{k+probit_tail_extra})`. inline constexpr unsigned probit_tail_extra = 10; // ln((32+i)/64) * 2^64, stored as a positive magnitude. Every anchor is in (0, ln 2]. static constexpr std::uint64_t probit_ln2_64 = 12786308645202655660ull; static constexpr std::uint64_t probit_ln_anchor_mag[32] = { 12786308645202655660ull, 12218671733053503897ull, 11667981761989453435ull, 11133256087961349648ull, 10613595130224743362ull, 10108173265494422292ull, 9616230936675340827ull, 9137067786804269247ull, 8670036662410753619ull, 8214538357444912273ull, 7770016990662967709ull, 7335955927010031419ull, 6911874167941132216ull, 6497323147432841322ull, 6091883880171659064ull, 5695164416463867605ull, 5306797565112371681ull, 4926438851101192057ull, 4553764679618851579ull, 4188470681899169456ull, 3830270221691897566ull, 3478893044001375095ull, 3134084050134459383ull, 2795602185149230175ull, 2463219425550596028ull, 2136719856585056848ull, 1815898829783402670ull, 1500562192519310430ull, 1190525582320469641ull, 885613779509420443ull, 585660112482476600ull, 290505910572683730ull, }; /// `round_half_away(ln(probability / 2^k) * 2^{k+10})`. inline std::int64_t probit_ln_argument(unsigned fractional_bits, std::int64_t probability) { const auto bits = static_cast(probability); const int e = 63 - __builtin_clzll(bits); const unsigned shift_in = static_cast(e + 1); const auto wide = static_cast(bits) << (64u - shift_in); const auto m64 = static_cast(wide); const unsigned idx = static_cast((m64 - (1ull << 63)) >> 58); const unsigned b_num = 32u + idx; const unsigned __int128 t_scaled = (static_cast(m64) * 64u) / b_num; __int128 t = static_cast<__int128>(t_scaled - (static_cast(1) << 64)); __int128 p = t; __int128 acc = 0; for (int n = 1; n <= 14; ++n) { const __int128 term = p / n; acc += (n & 1) ? term : -term; p = (p * t) >> 64; } const __int128 ln_m = -static_cast<__int128>(probit_ln_anchor_mag[idx]) + acc; const int exp_fix = e + 1 - static_cast(fractional_bits); const __int128 ln_x = ln_m + static_cast<__int128>(exp_fix) * static_cast<__int128>(probit_ln2_64); return round_half_away_i128(ln_x, 54u - fractional_bits); } /// Horner, then one extra right shift so a tail argument at scale `k+10` rounds onto scale `k`. inline std::int64_t eval_cubic_extra( const cubic_bits & piece, unsigned q, std::int64_t raw, unsigned fractional_bits, unsigned extra_shift) { using namespace principal_detail; __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(q) < static_cast(fractional_bits) + 16 ? static_cast(q) : static_cast(fractional_bits) + 16; const int drop = static_cast(q) - q_use; for (int i = 0; i < 4; ++i) coeff[i] = rshift_ties_even(coeff[i], drop); 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(3 - i)); acc = w_add(acc, term); } const unsigned denom_shift = static_cast(q_use) + 2u * fractional_bits + extra_shift; return round_half_away_pow2(acc, denom_shift); } inline std::int64_t eval_probit_abs(unsigned fractional_bits, std::int64_t probability) { const window_table & mid = at(PROBIT_MID, fractional_bits); if (probability >= mid.knots[0]) return eval_table(mid, fractional_bits, probability); const window_table & tail = at(PROBIT_TAIL, fractional_bits); std::int64_t u = probit_ln_argument(fractional_bits, probability); if (u < tail.knots[0]) u = tail.knots[0]; if (u > tail.knots[tail.nparts]) u = tail.knots[tail.nparts]; const unsigned scale = fractional_bits + probit_tail_extra; return eval_cubic_extra(tail.pieces[piece_of(tail, u)], tail.q, u, scale, probit_tail_extra); } inline std::int64_t eval_smoothstep(unsigned fractional_bits, std::int64_t raw) { const auto half = std::int64_t{1} << (fractional_bits - 1); if (raw <= -half) return 0; if (raw >= half) return std::int64_t{1} << fractional_bits; // -2 x^3 + (3/2) x + 1/2, with x = raw / 2^k. const __int128 x = raw; const __int128 cubic = round_half_away_i128(-(x * x * x), 2u * fractional_bits - 1u); const __int128 linear = round_half_away_i128(3 * x, 1); return static_cast(cubic + linear + half); } } // namespace window_detail inline std::int64_t eval_window(window which, unsigned fractional_bits, std::int64_t raw) { if (!principal_precision(fractional_bits)) throw std::invalid_argument("window lut: precision must be 8, 12, ..., 32"); using namespace window_detail; switch (which) { case window::smoothstep: return eval_smoothstep(fractional_bits, raw); case window::sigmoid: return eval_tailed(SIGMOID, tail_zero, tail_one, fractional_bits, raw); case window::tanh: return eval_tailed(TANH, tail_neg, tail_one, fractional_bits, raw); case window::erf: return eval_tailed(ERF, tail_neg, tail_one, fractional_bits, raw); case window::erfc: return (std::int64_t{1} << fractional_bits) - eval_tailed(ERF, tail_neg, tail_one, fractional_bits, raw); case window::softplus: return eval_tailed(SOFTPLUS, tail_zero, tail_id, fractional_bits, raw); case window::softminus: return raw - eval_tailed(SOFTPLUS, tail_zero, tail_id, fractional_bits, raw); case window::logsigmoid: return -eval_tailed(SOFTPLUS, tail_zero, tail_id, fractional_bits, -raw); case window::gelu: return eval_tailed(GELU, tail_zero, tail_id, fractional_bits, raw); case window::silu: return eval_tailed(SILU, tail_zero, tail_id, fractional_bits, raw); case window::mish: return eval_tailed(MISH, tail_zero, tail_id, fractional_bits, raw); case window::elish: return eval_tailed(ELISH, tail_zero, tail_id, fractional_bits, raw); case window::serf: return eval_tailed(SERF, tail_zero, tail_id, fractional_bits, raw); case window::tanhexp: return eval_tailed(TANHEXP, tail_zero, tail_id, fractional_bits, raw); case window::asin: case window::acos: { const auto one = std::int64_t{1} << fractional_bits; if (raw < -one || raw > one) throw std::out_of_range("window lut: asin/acos domain is [-1, 1]"); const std::int64_t positive = eval_asin_abs(fractional_bits, raw < 0 ? -raw : raw); const std::int64_t signed_asin = raw < 0 ? -positive : positive; if (which == window::asin) return signed_asin; return HALF_PI_RAW[slot_of(fractional_bits)] - signed_asin; } case window::probit: { const auto one = std::int64_t{1} << fractional_bits; if (raw <= 0 || raw >= one) throw std::out_of_range("window lut: probit domain is (0, 1)"); const auto half = one >> 1; if (raw > half) return -eval_probit_abs(fractional_bits, one - raw); return eval_probit_abs(fractional_bits, raw); } } throw std::invalid_argument("window lut: unknown function"); } } // namespace grotto #endif // LIBDPF_INCLUDE_GROTTO_WINDOW_LUT_HPP__