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
|
|
|
|
|
{
|
|
|
|
|
|
|
|
|
|
/// Piece `y_raw = round((c0 + c1·raw + c2·raw²) / den)`.
|
|
|
|
|
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(); }
|
|
|
|
|
|
|
|
|
|
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
|
|
|
|
|
|
|
|
|
|
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;
|
|
|
|
|
});
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
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;
|
|
|
|
|
});
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
/// Negative side is `x / 2^shift`, rounded to nearest, ties away from zero.
|
|
|
|
|
/// `shift == 0` is the identity. The slope does not depend on fractional width.
|
|
|
|
|
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};
|
|
|
|
|
});
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
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
|
|
|
|
|
|
|
|
|
|
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);
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
template <typename Raw>
|
|
|
|
|
easy_lut<Raw> make_relu6_lut(unsigned fractional_bits)
|
|
|
|
|
{
|
|
|
|
|
return make_clip_lut<Raw>(fractional_bits, 0, 6);
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
template <typename Raw>
|
|
|
|
|
easy_lut<Raw> make_hardtanh_lut(unsigned fractional_bits)
|
|
|
|
|
{
|
|
|
|
|
return make_clip_lut<Raw>(fractional_bits, -1, 1);
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
/// `0` on `[-1, 1]`, `x - 1` above, `x + 1` below.
|
|
|
|
|
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};
|
|
|
|
|
});
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
/// `0` on `[-1, 1]`, identity outside. Lambda is the integer 1.
|
|
|
|
|
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;
|
|
|
|
|
});
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
/// `0` left of `-3`, `1` right of `3`, `(x + 3) / 6` between, rounded.
|
|
|
|
|
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};
|
|
|
|
|
});
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
/// `0` left of `-3`, `x` right of `3`, `x(x + 3) / 6` between, rounded.
|
|
|
|
|
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};
|
|
|
|
|
});
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
} // namespace grotto
|
|
|
|
|
|
|
|
|
|
#endif // LIBDPF_INCLUDE_GROTTO_EASY_LUT_HPP__
|