libdpf/include/grotto/window_lut.hpp

392 lines
15 KiB
C++
Raw Normal View History

/// @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;
}
/// @brief `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.
/// @param fractional_bits the number of fractional bits
/// @param magnitude the magnitude
/// @return `round(sqrt(v / 2^{k+1}) * 2^{k+extra})`, `v > 0`
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;
}
/// @brief 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,
};
/// @brief `round_half_away(ln(probability / 2^k) * 2^{k+10})`.
/// @param fractional_bits the number of fractional bits
/// @param probability the `probability`
/// @return `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);
}
/// @brief Horner, then one extra right shift so a tail argument at scale `k+10` rounds onto scale `k`.
/// @param piece the `piece`
/// @param q the `q`
/// @param raw the underlying integer
/// @param fractional_bits the number of fractional bits
/// @param extra_shift the `extra_shift`
/// @return 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__