Initial import of libdpf.
Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
commit
e4e666f459
4563 changed files with 1690372 additions and 0 deletions
253
include/grotto/window_lut.hpp
Normal file
253
include/grotto/window_lut.hpp
Normal file
|
|
@ -0,0 +1,253 @@
|
|||
/// @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__
|
||||
Loading…
Add table
Add a link
Reference in a new issue