libdpf/include/grotto/principal_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

365 lines
11 KiB
C++

/// @file grotto/principal_lut.hpp
/// @brief Principal-domain cubics at fractional precisions 8, 12, ..., 32.
/// @details Eleven maps reuse the elementary2 tournament knots (the
/// arbitrary-breakpoint cubics chosen among Mathematica, Maple, and
/// Chebfun). A requested precision only rounds those coefficients
/// down to `k + 16` fractional bits. `coth` is different: its
/// principal function depends on `k` through
/// `beta = ln(2^{k+1}+1)/2`, so each precision has its own
/// partition. `inv` (1/x), `rsqrt` (1/sqrt(x)), and `invsq`
/// (1/x^2) are also per precision: each is a longest-feasible
/// cubic march on the closed principal interval [1/2, 1], with
/// absolute error at most half an ulp at that precision. They do
/// not apply an exponent lift. The returned raw value is
/// `round_half_away(p(x) * 2^k)`. On the closed principal domain,
/// `p` stays within one unit in the last place of that precision.
#ifndef LIBDPF_INCLUDE_GROTTO_PRINCIPAL_LUT_HPP__
#define LIBDPF_INCLUDE_GROTTO_PRINCIPAL_LUT_HPP__
#include <cstdint>
#include <stdexcept>
namespace grotto
{
enum class principal : unsigned
{
ln = 0,
exp,
sin,
tanf,
tang,
sinh,
cosh,
sqrt,
coth,
sec,
gsec,
csch,
inv,
rsqrt,
invsq,
};
inline constexpr unsigned principal_precisions[] = {8u, 12u, 16u, 20u, 24u, 28u, 32u};
inline constexpr bool principal_precision(unsigned fractional_bits) noexcept
{
for (unsigned k : principal_precisions)
{
if (k == fractional_bits)
return true;
}
return false;
}
namespace principal_detail
{
struct cubic_bits
{
std::int64_t hi[4];
std::uint64_t lo[4];
};
struct knot
{
std::int64_t num;
std::uint8_t sh;
};
struct table_ref
{
const knot * knots;
const cubic_bits * pieces;
std::uint16_t nparts;
std::uint16_t q;
};
#include "grotto/principal_tables.inc"
#include "grotto/principal_recip_tables.inc"
struct w256
{
unsigned __int128 lo;
__int128 hi;
};
inline w256 w_from_i128(__int128 value)
{
w256 out;
out.lo = static_cast<unsigned __int128>(value);
out.hi = value < 0 ? __int128(-1) : __int128(0);
return out;
}
inline w256 w_add(w256 lhs, w256 rhs)
{
w256 out;
out.lo = lhs.lo + rhs.lo;
const unsigned __int128 carry = out.lo < lhs.lo ? 1u : 0u;
const unsigned __int128 hi = static_cast<unsigned __int128>(lhs.hi)
+ static_cast<unsigned __int128>(rhs.hi) + carry;
out.hi = static_cast<__int128>(hi);
return out;
}
inline w256 w_shl(w256 value, unsigned shift)
{
if (shift == 0)
return value;
unsigned __int128 hi = static_cast<unsigned __int128>(value.hi);
w256 out;
if (shift >= 128)
{
const unsigned off = shift - 128;
out.lo = 0;
out.hi = static_cast<__int128>(off >= 128 ? 0 : value.lo << off);
return out;
}
out.lo = value.lo << shift;
const unsigned __int128 spill = value.lo >> (128u - shift);
out.hi = static_cast<__int128>((hi << shift) | spill);
return out;
}
inline w256 w_mul_u64(w256 value, std::uint64_t factor)
{
const unsigned __int128 hi = static_cast<unsigned __int128>(value.hi);
const std::uint64_t limbs[4] = {
static_cast<std::uint64_t>(value.lo),
static_cast<std::uint64_t>(value.lo >> 64),
static_cast<std::uint64_t>(hi),
static_cast<std::uint64_t>(hi >> 64),
};
std::uint64_t out_limbs[4];
unsigned __int128 carry = 0;
for (int i = 0; i < 4; ++i)
{
const unsigned __int128 prod = static_cast<unsigned __int128>(limbs[i]) * factor + carry;
out_limbs[i] = static_cast<std::uint64_t>(prod);
carry = prod >> 64;
}
w256 out;
out.lo = static_cast<unsigned __int128>(out_limbs[0])
| (static_cast<unsigned __int128>(out_limbs[1]) << 64);
const unsigned __int128 new_hi = static_cast<unsigned __int128>(out_limbs[2])
| (static_cast<unsigned __int128>(out_limbs[3]) << 64);
out.hi = static_cast<__int128>(new_hi);
return out;
}
inline bool w_negative(w256 value) noexcept
{
return value.hi < 0;
}
inline w256 w_neg(w256 value)
{
unsigned __int128 lo = ~value.lo + 1;
unsigned __int128 hi = ~static_cast<unsigned __int128>(value.hi);
if (lo == 0)
hi += 1;
w256 out;
out.lo = lo;
out.hi = static_cast<__int128>(hi);
return out;
}
inline w256 w_add_pow2(w256 value, unsigned shift)
{
w256 add;
add.lo = 0;
add.hi = 0;
if (shift < 128)
add.lo = static_cast<unsigned __int128>(1) << shift;
else if (shift < 256)
add.hi = static_cast<__int128>(static_cast<unsigned __int128>(1) << (shift - 128));
return w_add(value, add);
}
inline w256 w_shr(w256 value, unsigned shift)
{
if (shift == 0)
return value;
const unsigned __int128 hi = static_cast<unsigned __int128>(value.hi);
w256 out;
if (shift >= 256)
{
out.lo = 0;
out.hi = 0;
return out;
}
if (shift >= 128)
{
out.hi = 0;
out.lo = hi >> (shift - 128);
return out;
}
out.lo = (value.lo >> shift) | (hi << (128u - shift));
out.hi = static_cast<__int128>(hi >> shift);
return out;
}
inline std::int64_t w_to_i64(w256 value)
{
const unsigned __int128 hi = static_cast<unsigned __int128>(value.hi);
if (hi == 0)
{
if (value.lo > static_cast<unsigned __int128>(INT64_MAX))
throw std::overflow_error("principal lut: value does not fit int64");
return static_cast<std::int64_t>(value.lo);
}
if (value.hi == __int128(-1))
{
const auto hi_lo = static_cast<std::uint64_t>(value.lo >> 64);
if (hi_lo != ~std::uint64_t{0})
throw std::overflow_error("principal lut: value does not fit int64");
return static_cast<std::int64_t>(static_cast<std::uint64_t>(value.lo));
}
throw std::overflow_error("principal lut: value does not fit int64");
}
inline std::int64_t round_half_away_pow2(w256 numerator, unsigned shift)
{
if (shift == 0)
return w_to_i64(numerator);
const bool neg = w_negative(numerator);
w256 mag = neg ? w_neg(numerator) : numerator;
mag = w_add_pow2(mag, shift - 1);
mag = w_shr(mag, shift);
if (neg)
mag = w_neg(mag);
return w_to_i64(mag);
}
inline __int128 unpack_coeff(std::int64_t hi, std::uint64_t lo)
{
__int128 value = hi;
value <<= 64;
value |= static_cast<__int128>(lo);
return value;
}
inline __int128 rshift_ties_even(__int128 value, int shift)
{
if (shift <= 0)
return value;
const bool neg = value < 0;
const auto mag = static_cast<unsigned __int128>(neg ? -value : value);
const auto half = static_cast<unsigned __int128>(1) << (shift - 1);
const auto mask = (static_cast<unsigned __int128>(1) << shift) - 1;
unsigned __int128 quot = mag >> shift;
const unsigned __int128 rem = mag & mask;
if (rem > half || (rem == half && (quot & 1u)))
++quot;
const auto out = static_cast<__int128>(quot);
return neg ? -out : out;
}
inline bool knot_le_raw(knot endpoint, std::int64_t raw, unsigned fractional_bits)
{
const __int128 left = static_cast<__int128>(endpoint.num) << fractional_bits;
const __int128 right = static_cast<__int128>(raw) << endpoint.sh;
return left <= right;
}
inline int piece_index(const table_ref & table, std::int64_t raw, unsigned fractional_bits)
{
if (!knot_le_raw(table.knots[0], raw, fractional_bits))
throw std::out_of_range("principal lut: input is below the principal domain");
const knot stop = table.knots[table.nparts];
const __int128 stop_scaled = static_cast<__int128>(stop.num) << fractional_bits;
const __int128 raw_scaled = static_cast<__int128>(raw) << stop.sh;
if (raw_scaled > stop_scaled)
throw std::out_of_range("principal lut: input is above the principal domain");
int lo = 0;
int hi = static_cast<int>(table.nparts);
while (hi - lo > 1)
{
const int mid = (lo + hi) / 2;
if (knot_le_raw(table.knots[mid], raw, fractional_bits))
lo = mid;
else
hi = mid;
}
return lo;
}
inline const table_ref & table_for(principal which, unsigned fractional_bits)
{
const auto index = static_cast<unsigned>(which);
if (which == principal::coth)
{
const unsigned slot = fractional_bits / 4u - 2u;
return *COTH_BY_K[slot];
}
if (which == principal::inv || which == principal::rsqrt || which == principal::invsq)
{
const unsigned slot = fractional_bits / 4u - 2u;
const table_ref * const * bank = which == principal::inv
? INV_BY_K
: which == principal::rsqrt ? RSQRT_BY_K : INVSQ_BY_K;
return *bank[slot];
}
return *SHARED_TABLE[index];
}
inline w256 w_mul_i64(w256 value, std::int64_t factor)
{
if (factor >= 0)
return w_mul_u64(value, static_cast<std::uint64_t>(factor));
return w_neg(w_mul_u64(value, static_cast<std::uint64_t>(-factor)));
}
inline std::int64_t horner(const cubic_bits & piece, unsigned q, std::int64_t raw, unsigned fractional_bits)
{
__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);
// S3 / 2^{q_use + 3k} = p(x), so round(p * 2^k) = round(S3 / 2^{q_use + 2k}).
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;
return round_half_away_pow2(acc, denom_shift);
}
} // namespace principal_detail
/// Number of cubic pieces at this fractional width.
inline std::uint16_t principal_parts(principal which, unsigned fractional_bits)
{
if (!principal_precision(fractional_bits))
throw std::invalid_argument("principal lut: precision must be 8, 12, ..., 32");
return principal_detail::table_for(which, fractional_bits).nparts;
}
/// Evaluate the principal-domain cubic. `raw / 2^fractional_bits` is the input.
inline std::int64_t eval_principal(principal which, unsigned fractional_bits, std::int64_t raw)
{
if (!principal_precision(fractional_bits))
throw std::invalid_argument("principal lut: precision must be 8, 12, ..., 32");
if (raw < 0)
throw std::out_of_range("principal lut: input is below the principal domain");
const principal_detail::table_ref & table = principal_detail::table_for(which, fractional_bits);
const int index = principal_detail::piece_index(table, raw, fractional_bits);
return principal_detail::horner(table.pieces[index], table.q, raw, fractional_bits);
}
} // namespace grotto
#endif // LIBDPF_INCLUDE_GROTTO_PRINCIPAL_LUT_HPP__