2026-09-24 14:08:32 -06:00
/// @file grotto/easy_lut.hpp
/// @brief Exact few-piece polynomials whose knots are obvious integers.
/// @details One program per function. A fractional width only slides knots
/// that sit on an integer (`k` becomes raw `k << F`) and scales
/// constant terms that are themselves integers. Slopes of `0`, `±1`,
/// and `1/2^s` stay exact; `hardsigmoid` / `hardswish` divide by 6
/// and round the final raw encoding to nearest, ties away from zero.
# ifndef LIBDPF_INCLUDE_GROTTO_EASY_LUT_HPP__
# define LIBDPF_INCLUDE_GROTTO_EASY_LUT_HPP__
2026-09-24 20:44:07 -06:00
# include "hedley/hedley.h"
2026-09-24 14:08:32 -06:00
# include <algorithm>
# include <cstdint>
# include <limits>
# include <optional>
# include <stdexcept>
# include <type_traits>
# include <utility>
# include <vector>
namespace grotto
{
2026-09-24 23:18:10 -06:00
/// @brief Piece `y_raw = round((c0 + c1·raw + c2·raw²) / den)`.
/// @tparam Raw underlying representation
2026-09-24 14:08:32 -06:00
template < typename Raw >
struct easy_lut
{
static_assert ( std : : is_integral_v < Raw > & & std : : is_signed_v < Raw > ) ;
using raw_type = Raw ;
std : : vector < Raw > bounds ;
std : : vector < std : : int64_t > c0 ;
std : : vector < std : : int64_t > c1 ;
std : : vector < std : : int64_t > c2 ;
std : : vector < std : : int64_t > den ;
2026-09-24 20:44:07 -06:00
HEDLEY_NO_THROW
2026-09-24 14:08:32 -06:00
std : : size_t parts ( ) const noexcept { return c0 . size ( ) ; }
2026-09-28 05:59:19 -06:00
/// \complexity `upper_bound` on `bounds` (`Θ(log P)` comparisons, `P = parts()`), then three multiplies for the quadratic. Extra space `Θ(1)`.
/// @param x raw fixed-point input
/// @return the rounded piece value
2026-09-24 14:08:32 -06:00
std : : int64_t operator ( ) ( Raw x ) const
{
const auto it = std : : upper_bound ( bounds . begin ( ) , bounds . end ( ) , x ) ;
const auto i = static_cast < std : : size_t > ( it - bounds . begin ( ) ) - 1 ;
const __int128 raw = x ;
const __int128 acc = __int128 ( c0 [ i ] ) + __int128 ( c1 [ i ] ) * raw
+ __int128 ( c2 [ i ] ) * raw * raw ;
return detail_round_div ( acc , den [ i ] ) ;
}
private :
static std : : int64_t detail_round_div ( __int128 num , std : : int64_t den )
{
if ( den < = 0 )
throw std : : invalid_argument ( " easy lut: denominator must be positive " ) ;
if ( den = = 1 )
{
if ( num > std : : numeric_limits < std : : int64_t > : : max ( )
| | num < std : : numeric_limits < std : : int64_t > : : min ( ) )
throw std : : overflow_error ( " easy lut: value does not fit int64 " ) ;
return static_cast < std : : int64_t > ( num ) ;
}
const bool neg = num < 0 ;
const __int128 mag = neg ? - num : num ;
const __int128 d = den ;
const __int128 q = ( mag + d / 2 ) / d ;
const __int128 signed_q = neg ? - q : q ;
if ( signed_q > std : : numeric_limits < std : : int64_t > : : max ( )
| | signed_q < std : : numeric_limits < std : : int64_t > : : min ( ) )
throw std : : overflow_error ( " easy lut: value does not fit int64 " ) ;
return static_cast < std : : int64_t > ( signed_q ) ;
}
} ;
namespace detail
{
struct easy_poly
{
std : : int64_t c0 = 0 ;
std : : int64_t c1 = 0 ;
std : : int64_t c2 = 0 ;
std : : int64_t den = 1 ;
} ;
2026-09-24 20:44:07 -06:00
HEDLEY_NO_THROW
2026-09-24 14:08:32 -06:00
inline bool operator = = ( easy_poly a , easy_poly b ) noexcept
{
return a . c0 = = b . c0 & & a . c1 = = b . c1 & & a . c2 = = b . c2 & & a . den = = b . den ;
}
inline constexpr easy_poly kIdentity { 0 , 1 , 0 , 1 } ;
inline constexpr easy_poly kZero { 0 , 0 , 0 , 1 } ;
template < typename Raw >
std : : optional < std : : int64_t > scaled_integer ( std : : int64_t units , unsigned fractional_bits )
{
using lim = std : : numeric_limits < Raw > ;
if ( fractional_bits > = 63 )
return std : : nullopt ;
const __int128 v = static_cast < __int128 > ( units ) < < fractional_bits ;
if ( v < static_cast < __int128 > ( lim : : min ( ) ) | | v > static_cast < __int128 > ( lim : : max ( ) ) )
return std : : nullopt ;
return static_cast < std : : int64_t > ( v ) ;
}
template < typename Raw , typename PolyAt >
easy_lut < Raw > assemble_easy ( std : : vector < std : : int64_t > cuts , PolyAt & & poly_at )
{
using lim = std : : numeric_limits < Raw > ;
cuts . push_back ( static_cast < std : : int64_t > ( lim : : min ( ) ) ) ;
std : : sort ( cuts . begin ( ) , cuts . end ( ) ) ;
cuts . erase ( std : : unique ( cuts . begin ( ) , cuts . end ( ) ) , cuts . end ( ) ) ;
const std : : int64_t maxv = static_cast < std : : int64_t > ( lim : : max ( ) ) ;
cuts . erase ( std : : remove_if ( cuts . begin ( ) , cuts . end ( ) ,
[ & ] ( std : : int64_t c ) { return c > maxv ; } ) , cuts . end ( ) ) ;
easy_lut < Raw > lut ;
for ( std : : size_t i = 0 ; i < cuts . size ( ) ; + + i )
{
const std : : int64_t start = cuts [ i ] ;
const std : : int64_t last = ( i + 1 < cuts . size ( ) ) ? cuts [ i + 1 ] - 1 : maxv ;
const easy_poly poly = poly_at ( start ) ;
if ( poly . den < = 0 )
throw std : : invalid_argument ( " easy lut: denominator must be positive " ) ;
if ( ! ( poly = = poly_at ( last ) ) )
throw std : : logic_error ( " easy lut span is not one polynomial " ) ;
const __int128 width = static_cast < __int128 > ( last ) - static_cast < __int128 > ( start ) ;
if ( width > 2 )
{
const std : : int64_t mid = static_cast < std : : int64_t > (
static_cast < __int128 > ( start ) + width / 2 ) ;
if ( ! ( poly = = poly_at ( mid ) ) )
throw std : : logic_error ( " easy lut span is not one polynomial " ) ;
}
if ( ! lut . c0 . empty ( ) & & poly = = easy_poly { lut . c0 . back ( ) , lut . c1 . back ( ) , lut . c2 . back ( ) , lut . den . back ( ) } )
continue ;
lut . bounds . push_back ( static_cast < Raw > ( start ) ) ;
lut . c0 . push_back ( poly . c0 ) ;
lut . c1 . push_back ( poly . c1 ) ;
lut . c2 . push_back ( poly . c2 ) ;
lut . den . push_back ( poly . den ) ;
}
return lut ;
}
inline std : : int64_t denom_shift ( unsigned shift )
{
if ( shift > = 63 )
throw std : : invalid_argument ( " easy lut: dyadic slope is too small " ) ;
return std : : int64_t { 1 } < < shift ;
}
} // namespace detail
2026-09-28 05:59:19 -06:00
/// \complexity Assembles a constant number of pieces (the knots are a fixed integer pattern shifted by `fractional_bits`). `Θ(1)` time and extra space, aside from the returned vectors of that constant length.
/// @see grotto::easy_lut
/// @see grotto::eval_window
/// @param fractional_bits fractional bits; integer knots become `k << fractional_bits`
/// @tparam Raw signed raw word
2026-09-24 14:08:32 -06:00
template < typename Raw >
easy_lut < Raw > make_abs_lut ( unsigned fractional_bits = 0 )
{
( void ) fractional_bits ;
return detail : : assemble_easy < Raw > ( { 0 } , [ ] ( std : : int64_t raw ) {
detail : : easy_poly p = detail : : kIdentity ;
if ( raw < 0 )
p . c1 = - 1 ;
return p ;
} ) ;
}
2026-09-28 05:59:19 -06:00
/// \complexity Assembles a constant number of pieces (the knots are a fixed integer pattern shifted by `fractional_bits`). `Θ(1)` time and extra space, aside from the returned vectors of that constant length.
/// @see grotto::easy_lut
/// @see grotto::eval_window
/// @param fractional_bits fractional bits; integer knots become `k << fractional_bits`
/// @tparam Raw signed raw word
2026-09-24 14:08:32 -06:00
template < typename Raw >
easy_lut < Raw > make_relu_lut ( unsigned fractional_bits = 0 )
{
( void ) fractional_bits ;
return detail : : assemble_easy < Raw > ( { 0 } , [ ] ( std : : int64_t raw ) {
return raw < 0 ? detail : : kZero : detail : : kIdentity ;
} ) ;
}
2026-09-24 23:18:10 -06:00
/// @brief Negative side is `x / 2^shift`, rounded to nearest, ties away from zero.
/// @details `shift == 0` is the identity. The slope does not depend on fractional width.
/// @tparam Raw underlying representation
/// @param shift the bit shift
/// @return Negative side is `x / 2^shift`, rounded to nearest, ties away from zero
2026-09-28 05:59:19 -06:00
/// \complexity Assembles a constant number of pieces. `Θ(1)` time and extra space, aside from the returned vectors of that constant length.
/// @see grotto::easy_lut
/// @see grotto::eval_window
2026-09-24 14:08:32 -06:00
template < typename Raw >
easy_lut < Raw > make_leaky_relu_lut ( unsigned shift )
{
const std : : int64_t den = detail : : denom_shift ( shift ) ;
return detail : : assemble_easy < Raw > ( { 0 } , [ = ] ( std : : int64_t raw ) {
if ( raw > = 0 )
return detail : : kIdentity ;
return detail : : easy_poly { 0 , 1 , 0 , den } ;
} ) ;
}
2026-09-28 05:59:19 -06:00
/// \complexity Assembles a constant number of pieces (the knots are a fixed integer pattern shifted by `fractional_bits`). `Θ(1)` time and extra space, aside from the returned vectors of that constant length.
/// @see grotto::easy_lut
/// @see grotto::eval_window
/// @param fractional_bits fractional bits; integer knots become `k << fractional_bits`
/// @tparam Raw signed raw word
2026-09-24 14:08:32 -06:00
template < typename Raw >
easy_lut < Raw > make_squared_relu_lut ( unsigned fractional_bits )
{
const std : : int64_t den = detail : : denom_shift ( fractional_bits ) ;
return detail : : assemble_easy < Raw > ( { 0 } , [ = ] ( std : : int64_t raw ) {
if ( raw < 0 )
return detail : : kZero ;
return detail : : easy_poly { 0 , 0 , 1 , den } ;
} ) ;
}
namespace detail
{
template < typename Raw >
easy_lut < Raw > clip_to ( unsigned fractional_bits , std : : int64_t low_units , std : : int64_t high_units )
{
if ( low_units > high_units )
throw std : : invalid_argument ( " clip lut: low > high " ) ;
const auto low = scaled_integer < Raw > ( low_units , fractional_bits ) ;
const auto high = scaled_integer < Raw > ( high_units , fractional_bits ) ;
std : : vector < std : : int64_t > cuts { 0 } ;
if ( low )
cuts . push_back ( * low ) ;
if ( high & & * high < std : : numeric_limits < Raw > : : max ( ) )
cuts . push_back ( * high + 1 ) ;
const std : : int64_t low_raw = low ? * low : std : : numeric_limits < std : : int64_t > : : min ( ) ;
const std : : int64_t high_raw = high ? * high : std : : numeric_limits < std : : int64_t > : : max ( ) ;
const std : : int64_t low_level = low ? * low : 0 ;
const std : : int64_t high_level = high ? * high : 0 ;
return assemble_easy < Raw > ( std : : move ( cuts ) , [ = ] ( std : : int64_t raw ) {
if ( low & & raw < low_raw )
return easy_poly { low_level , 0 , 0 , 1 } ;
if ( high & & raw > high_raw )
return easy_poly { high_level , 0 , 0 , 1 } ;
return kIdentity ;
} ) ;
}
} // namespace detail
2026-09-28 05:59:19 -06:00
/// @brief Clip to `[low_units, high_units]` in raw units, then shift the knots by `fractional_bits`.
/// @tparam Raw signed raw word
/// @param fractional_bits fractional bits; integer knots become `k << fractional_bits`
/// @param low_units inclusive lower clip, in integer units before the fractional shift
/// @param high_units inclusive upper clip, in integer units before the fractional shift
/// \complexity Assembles a constant number of pieces. `Θ(1)` time and extra space.
/// @see grotto::easy_lut
/// @see grotto::eval_window
2026-09-24 14:08:32 -06:00
template < typename Raw >
easy_lut < Raw > make_clip_lut ( unsigned fractional_bits , std : : int64_t low_units , std : : int64_t high_units )
{
return detail : : clip_to < Raw > ( fractional_bits , low_units , high_units ) ;
}
2026-09-28 05:59:19 -06:00
/// \complexity Assembles a constant number of pieces (the knots are a fixed integer pattern shifted by `fractional_bits`). `Θ(1)` time and extra space, aside from the returned vectors of that constant length.
/// @see grotto::easy_lut
/// @see grotto::eval_window
/// @param fractional_bits fractional bits; integer knots become `k << fractional_bits`
/// @tparam Raw signed raw word
2026-09-24 14:08:32 -06:00
template < typename Raw >
easy_lut < Raw > make_relu6_lut ( unsigned fractional_bits )
{
return make_clip_lut < Raw > ( fractional_bits , 0 , 6 ) ;
}
2026-09-28 05:59:19 -06:00
/// \complexity Assembles a constant number of pieces (the knots are a fixed integer pattern shifted by `fractional_bits`). `Θ(1)` time and extra space, aside from the returned vectors of that constant length.
/// @see grotto::easy_lut
/// @see grotto::eval_window
/// @param fractional_bits fractional bits; integer knots become `k << fractional_bits`
/// @tparam Raw signed raw word
2026-09-24 14:08:32 -06:00
template < typename Raw >
easy_lut < Raw > make_hardtanh_lut ( unsigned fractional_bits )
{
return make_clip_lut < Raw > ( fractional_bits , - 1 , 1 ) ;
}
2026-09-24 23:18:10 -06:00
/// @brief `0` on `[-1, 1]`, `x - 1` above, `x + 1` below.
/// @tparam Raw underlying representation
/// @param fractional_bits the number of fractional bits
/// @return `0` on `[-1, 1]`, `x - 1` above, `x + 1` below
2026-09-28 05:59:19 -06:00
/// \complexity Assembles a constant number of pieces (the knots are a fixed integer pattern shifted by `fractional_bits`). `Θ(1)` time and extra space, aside from the returned vectors of that constant length.
/// @see grotto::easy_lut
/// @see grotto::eval_window
2026-09-24 14:08:32 -06:00
template < typename Raw >
easy_lut < Raw > make_softshrink_lut ( unsigned fractional_bits )
{
const auto knot = detail : : scaled_integer < Raw > ( 1 , fractional_bits ) ;
std : : vector < std : : int64_t > cuts { 0 } ;
if ( knot )
{
cuts . push_back ( - * knot ) ;
if ( * knot < std : : numeric_limits < Raw > : : max ( ) )
cuts . push_back ( * knot + 1 ) ;
}
const std : : int64_t k = knot ? * knot : 0 ;
return detail : : assemble_easy < Raw > ( std : : move ( cuts ) , [ = ] ( std : : int64_t raw ) {
if ( ! knot | | ( raw > = - k & & raw < = k ) )
return detail : : kZero ;
if ( raw > k )
return detail : : easy_poly { - k , 1 , 0 , 1 } ;
return detail : : easy_poly { k , 1 , 0 , 1 } ;
} ) ;
}
2026-09-24 23:18:10 -06:00
/// @brief `0` on `[-1, 1]`, identity outside. Lambda is the integer 1.
/// @tparam Raw underlying representation
/// @param fractional_bits the number of fractional bits
/// @return `0` on `[-1, 1]`, identity outside
2026-09-28 05:59:19 -06:00
/// \complexity Assembles a constant number of pieces (the knots are a fixed integer pattern shifted by `fractional_bits`). `Θ(1)` time and extra space, aside from the returned vectors of that constant length.
/// @see grotto::easy_lut
/// @see grotto::eval_window
2026-09-24 14:08:32 -06:00
template < typename Raw >
easy_lut < Raw > make_hardshrink_lut ( unsigned fractional_bits )
{
const auto knot = detail : : scaled_integer < Raw > ( 1 , fractional_bits ) ;
std : : vector < std : : int64_t > cuts { 0 } ;
if ( knot )
{
cuts . push_back ( - * knot ) ;
if ( * knot < std : : numeric_limits < Raw > : : max ( ) )
cuts . push_back ( * knot + 1 ) ;
}
const std : : int64_t k = knot ? * knot : 0 ;
return detail : : assemble_easy < Raw > ( std : : move ( cuts ) , [ = ] ( std : : int64_t raw ) {
if ( ! knot | | ( raw > = - k & & raw < = k ) )
return detail : : kZero ;
return detail : : kIdentity ;
} ) ;
}
2026-09-24 23:18:10 -06:00
/// @brief `0` left of `-3`, `1` right of `3`, `(x + 3) / 6` between, rounded.
/// @tparam Raw underlying representation
/// @param fractional_bits the number of fractional bits
/// @return `0` left of `-3`, `1` right of `3`, `(x + 3) / 6` between, rounded
/// @throws std::invalid_argument if `fractional width does not fit`
2026-09-28 05:59:19 -06:00
/// \complexity Assembles a constant number of pieces (the knots are a fixed integer pattern shifted by `fractional_bits`). `Θ(1)` time and extra space, aside from the returned vectors of that constant length.
/// @see grotto::easy_lut
/// @see grotto::eval_window
2026-09-24 14:08:32 -06:00
template < typename Raw >
easy_lut < Raw > make_hardsigmoid_lut ( unsigned fractional_bits )
{
if ( fractional_bits > = 62 )
throw std : : invalid_argument ( " hardsigmoid: fractional width does not fit " ) ;
const std : : int64_t three = std : : int64_t { 3 } < < fractional_bits ;
const std : : int64_t one = std : : int64_t { 1 } < < fractional_bits ;
const auto knot = detail : : scaled_integer < Raw > ( 3 , fractional_bits ) ;
std : : vector < std : : int64_t > cuts ;
if ( knot )
{
cuts . push_back ( - * knot ) ;
if ( * knot < std : : numeric_limits < Raw > : : max ( ) )
cuts . push_back ( * knot + 1 ) ;
}
const std : : int64_t k = knot ? * knot : three ;
return detail : : assemble_easy < Raw > ( std : : move ( cuts ) , [ = ] ( std : : int64_t raw ) {
if ( knot & & raw < - k )
return detail : : kZero ;
if ( knot & & raw > k )
return detail : : easy_poly { one , 0 , 0 , 1 } ;
return detail : : easy_poly { three , 1 , 0 , 6 } ;
} ) ;
}
2026-09-24 23:18:10 -06:00
/// @brief `0` left of `-3`, `x` right of `3`, `x(x + 3) / 6` between, rounded.
/// @tparam Raw underlying representation
/// @param fractional_bits the number of fractional bits
/// @return `0` left of `-3`, `x` right of `3`, `x(x + 3) / 6` between, rounded
/// @throws std::invalid_argument if `fractional width does not fit the denominator`
2026-09-28 05:59:19 -06:00
/// \complexity Assembles a constant number of pieces (the knots are a fixed integer pattern shifted by `fractional_bits`). `Θ(1)` time and extra space, aside from the returned vectors of that constant length.
/// @see grotto::easy_lut
/// @see grotto::eval_window
2026-09-24 14:08:32 -06:00
template < typename Raw >
easy_lut < Raw > make_hardswish_lut ( unsigned fractional_bits )
{
if ( fractional_bits > = 61 )
throw std : : invalid_argument ( " hardswish: fractional width does not fit the denominator " ) ;
const std : : int64_t three = std : : int64_t { 3 } < < fractional_bits ;
const std : : int64_t den = 6 * ( std : : int64_t { 1 } < < fractional_bits ) ;
const auto knot = detail : : scaled_integer < Raw > ( 3 , fractional_bits ) ;
std : : vector < std : : int64_t > cuts ;
if ( knot )
{
cuts . push_back ( - * knot ) ;
if ( * knot < std : : numeric_limits < Raw > : : max ( ) )
cuts . push_back ( * knot + 1 ) ;
}
const std : : int64_t k = knot ? * knot : three ;
return detail : : assemble_easy < Raw > ( std : : move ( cuts ) , [ = ] ( std : : int64_t raw ) {
if ( knot & & raw < - k )
return detail : : kZero ;
if ( knot & & raw > k )
return detail : : kIdentity ;
return detail : : easy_poly { 0 , three , 1 , den } ;
} ) ;
}
2026-09-28 05:59:19 -06:00
/// @brief Appendix D leaky ReLU: identity on the right, `x/100` on the left.
/// @details `make_leaky_relu_lut` is the dyadic slope `1/2^shift`. This one is
/// the paper's slope `1/100`, rounded to nearest, ties away from zero.
/// The fractional scale cancels, so the denominator does not depend
/// on `fractional_bits`.
/// @tparam Raw underlying representation
/// @return two pieces, degree 1
/// \complexity Assembles two pieces. `Θ(1)` time and extra space, aside from the returned vectors.
/// @see grotto::make_leaky_relu_lut
template < typename Raw >
easy_lut < Raw > make_leaky_relu_hundredth_lut ( )
{
return detail : : assemble_easy < Raw > ( { 0 } , [ ] ( std : : int64_t raw ) {
if ( raw > = 0 )
return detail : : kIdentity ;
return detail : : easy_poly { 0 , 1 , 0 , 100 } ;
} ) ;
}
2026-09-24 14:08:32 -06:00
} // namespace grotto
# endif // LIBDPF_INCLUDE_GROTTO_EASY_LUT_HPP__