libdpf/include/grotto/window_lut.hpp
Ryan Henry e4e666f459 Initial import of libdpf.
Co-authored-by: Cursor <cursoragent@cursor.com>
2026-09-24 14:08:32 -06:00

253 lines
8.8 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.
#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);
}
/// `round(sqrt(v / 2^{k+1}) * 2^k)`, `v > 0`.
inline std::int64_t sqrt_half_scale(unsigned fractional_bits, std::int64_t magnitude)
{
const int log = 63 - __builtin_clzll(static_cast<unsigned long long>(magnitude));
const std::int64_t mant = magnitude << (fractional_bits - static_cast<unsigned>(log + 1));
const std::int64_t root = eval_principal(principal::sqrt, fractional_bits, mant);
const int exp2 = log - static_cast<int>(fractional_bits);
if ((exp2 & 1) == 0)
return round_half_away_i128(root, static_cast<unsigned>(-exp2) / 2u);
const unsigned t = static_cast<unsigned>(-exp2 - 1) / 2u;
// sqrt(2) rounded onto 62 fractional bits.
constexpr __int128 sqrt2_62 = 6521908912666391106LL;
return round_half_away_i128(__int128(root) * sqrt2_62, 62u + t + 1u);
}
inline std::int64_t eval_asin_abs(unsigned fractional_bits, std::int64_t magnitude)
{
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;
const std::int64_t gap = one - magnitude;
const auto pi = HALF_PI_RAW[slot_of(fractional_bits)];
if (gap <= 0)
return pi;
std::int64_t reduced = sqrt_half_scale(fractional_bits, gap);
if (reduced > half)
reduced = half;
const std::int64_t inner = eval_table(table, fractional_bits, reduced);
const std::int64_t lifted = pi - 2 * inner;
return lifted < 0 ? 0 : lifted;
}
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);
return eval_table(at(PROBIT_TAIL, fractional_bits), fractional_bits, probability);
}
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__