libdpf/include/grotto/principal_lut.hpp
Ryan Henry 0d22946a0e Checkpoint the party/runtime stack before share-program and malicious-mode work.
Ship the TLS mesh, composer, Beaver/Yao/leaf MPC, prep/online paths, apps, and docs so the tree is pushable before elevating share_expr, security_mode, and prep resume.

Co-authored-by: Cursor <cursoragent@cursor.com>
2026-09-28 05:59:19 -06:00

388 lines
12 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 "hedley/hedley.h"
#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};
HEDLEY_NO_THROW
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;
}
HEDLEY_NO_THROW
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));
// `-factor` is undefined at INT64_MIN. The magnitude is the unsigned wrap.
const auto mag = static_cast<std::uint64_t>(0)
- static_cast<std::uint64_t>(factor);
return w_neg(w_mul_u64(value, mag));
}
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
/// @brief Number of cubic pieces at this fractional width.
/// @param which the principal map, not a party index
/// @param fractional_bits the number of fractional bits
/// @return Number of cubic pieces at this fractional width
/// @throws std::invalid_argument if `precision must be 8, 12, ..., 32`
/// \complexity One table lookup of `nparts`. `Θ(1)`.
/// @see grotto::eval_principal
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;
}
/// @brief Evaluate the principal-domain cubic. `raw / 2^fractional_bits` is the input.
/// @param which the principal map, not a party index
/// @param fractional_bits the number of fractional bits
/// @param raw the underlying integer
/// @return the returned `std::int64_t`
/// @throws std::invalid_argument if `precision must be 8, 12, ..., 32`
/// @throws std::out_of_range if `input is below the principal domain`
/// \complexity `piece_index` binary-searches `nparts` knots, then `horner` loops the four cubic coefficients (`i` from 3 down to 0).
/// Time `Θ(log P)` knot comparisons plus `Θ(1)` arithmetic, `P = principal_parts(which, fractional_bits)`. Extra space `Θ(1)`.
/// @see grotto::eval_reduced
/// @see grotto::eval_window
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__