libdpf/include/grotto/window_lut.hpp
Ryan Henry 875f09fec1 Record Grotto half-ulp tables and comparison geneval, and factor shared beaver terms before the quotient.
Horner and window evaluation need those tables in the tree. Comparison geneval opens the same value words as a Doerner–Shelat key. A factor common to every polynomial term is multiplied first so that preprocessing stays smaller.

Co-authored-by: Cursor <cursoragent@cursor.com>
2026-09-24 15:16:21 -06:00

378 lines
14 KiB
C++
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

/// @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 <cstdint>
#include <stdexcept>
#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<int>(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<std::int64_t>(number);
}
const bool neg = number < 0;
const auto mag = static_cast<unsigned __int128>(neg ? -number : number);
const unsigned __int128 quot = (mag + (static_cast<unsigned __int128>(1) << (shift - 1))) >> shift;
const auto out = static_cast<__int128>(quot);
return static_cast<std::int64_t>(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<unsigned>(__builtin_clzll(static_cast<unsigned long long>(n >> 64)))
: 64u - static_cast<unsigned>(__builtin_clzll(static_cast<unsigned long long>(n)));
unsigned __int128 x = static_cast<unsigned __int128>(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<unsigned __int128>(static_cast<std::uint64_t>(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<unsigned __int128>(1) << 64) | static_cast<unsigned __int128>(7640891576956012809ULL);
const unsigned __int128 scaled = (root * sqrt2_64 + (static_cast<unsigned __int128>(1) << 64)) >> 65;
return static_cast<std::int64_t>(scaled);
}
inline int piece_of_scaled(const window_table & table, std::int64_t raw, unsigned extra)
{
int lo = 0;
int hi = static_cast<int>(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<unsigned __int128>(1) << 64) | static_cast<unsigned __int128>(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<unsigned long long>(probability);
const int e = 63 - __builtin_clzll(bits);
const unsigned shift_in = static_cast<unsigned>(e + 1);
const auto wide = static_cast<unsigned __int128>(bits) << (64u - shift_in);
const auto m64 = static_cast<std::uint64_t>(wide);
const unsigned idx = static_cast<unsigned>((m64 - (1ull << 63)) >> 58);
const unsigned b_num = 32u + idx;
const unsigned __int128 t_scaled = (static_cast<unsigned __int128>(m64) * 64u) / b_num;
__int128 t = static_cast<__int128>(t_scaled - (static_cast<unsigned __int128>(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<int>(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<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);
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 + 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<std::int64_t>(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__