/// @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__ #include "hedley/hedley.h" #include #include #include #include #include #include #include #include namespace grotto { /// @brief Piece `y_raw = round((c0 + c1·raw + c2·raw²) / den)`. /// @tparam Raw underlying representation template struct easy_lut { static_assert(std::is_integral_v && std::is_signed_v); using raw_type = Raw; std::vector bounds; std::vector c0; std::vector c1; std::vector c2; std::vector den; HEDLEY_NO_THROW std::size_t parts() const noexcept { return c0.size(); } /// \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 std::int64_t operator()(Raw x) const { const auto it = std::upper_bound(bounds.begin(), bounds.end(), x); const auto i = static_cast(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::max() || num < std::numeric_limits::min()) throw std::overflow_error("easy lut: value does not fit int64"); return static_cast(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::max() || signed_q < std::numeric_limits::min()) throw std::overflow_error("easy lut: value does not fit int64"); return static_cast(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; }; HEDLEY_NO_THROW 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 std::optional scaled_integer(std::int64_t units, unsigned fractional_bits) { using lim = std::numeric_limits; 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(v); } template easy_lut assemble_easy(std::vector cuts, PolyAt && poly_at) { using lim = std::numeric_limits; cuts.push_back(static_cast(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(lim::max()); cuts.erase(std::remove_if(cuts.begin(), cuts.end(), [&](std::int64_t c) { return c > maxv; }), cuts.end()); easy_lut 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( 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(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 /// \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 template easy_lut make_abs_lut(unsigned fractional_bits = 0) { (void)fractional_bits; return detail::assemble_easy({0}, [](std::int64_t raw) { detail::easy_poly p = detail::kIdentity; if (raw < 0) p.c1 = -1; return p; }); } /// \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 template easy_lut make_relu_lut(unsigned fractional_bits = 0) { (void)fractional_bits; return detail::assemble_easy({0}, [](std::int64_t raw) { return raw < 0 ? detail::kZero : detail::kIdentity; }); } /// @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 /// \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 template easy_lut make_leaky_relu_lut(unsigned shift) { const std::int64_t den = detail::denom_shift(shift); return detail::assemble_easy({0}, [=](std::int64_t raw) { if (raw >= 0) return detail::kIdentity; return detail::easy_poly{0, 1, 0, den}; }); } /// \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 template easy_lut make_squared_relu_lut(unsigned fractional_bits) { const std::int64_t den = detail::denom_shift(fractional_bits); return detail::assemble_easy({0}, [=](std::int64_t raw) { if (raw < 0) return detail::kZero; return detail::easy_poly{0, 0, 1, den}; }); } namespace detail { template easy_lut 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(low_units, fractional_bits); const auto high = scaled_integer(high_units, fractional_bits); std::vector cuts{0}; if (low) cuts.push_back(*low); if (high && *high < std::numeric_limits::max()) cuts.push_back(*high + 1); const std::int64_t low_raw = low ? *low : std::numeric_limits::min(); const std::int64_t high_raw = high ? *high : std::numeric_limits::max(); const std::int64_t low_level = low ? *low : 0; const std::int64_t high_level = high ? *high : 0; return assemble_easy(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 /// @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 template easy_lut make_clip_lut(unsigned fractional_bits, std::int64_t low_units, std::int64_t high_units) { return detail::clip_to(fractional_bits, low_units, high_units); } /// \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 template easy_lut make_relu6_lut(unsigned fractional_bits) { return make_clip_lut(fractional_bits, 0, 6); } /// \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 template easy_lut make_hardtanh_lut(unsigned fractional_bits) { return make_clip_lut(fractional_bits, -1, 1); } /// @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 /// \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 template easy_lut make_softshrink_lut(unsigned fractional_bits) { const auto knot = detail::scaled_integer(1, fractional_bits); std::vector cuts{0}; if (knot) { cuts.push_back(-*knot); if (*knot < std::numeric_limits::max()) cuts.push_back(*knot + 1); } const std::int64_t k = knot ? *knot : 0; return detail::assemble_easy(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}; }); } /// @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 /// \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 template easy_lut make_hardshrink_lut(unsigned fractional_bits) { const auto knot = detail::scaled_integer(1, fractional_bits); std::vector cuts{0}; if (knot) { cuts.push_back(-*knot); if (*knot < std::numeric_limits::max()) cuts.push_back(*knot + 1); } const std::int64_t k = knot ? *knot : 0; return detail::assemble_easy(std::move(cuts), [=](std::int64_t raw) { if (!knot || (raw >= -k && raw <= k)) return detail::kZero; return detail::kIdentity; }); } /// @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` /// \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 template easy_lut 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(3, fractional_bits); std::vector cuts; if (knot) { cuts.push_back(-*knot); if (*knot < std::numeric_limits::max()) cuts.push_back(*knot + 1); } const std::int64_t k = knot ? *knot : three; return detail::assemble_easy(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}; }); } /// @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` /// \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 template easy_lut 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(3, fractional_bits); std::vector cuts; if (knot) { cuts.push_back(-*knot); if (*knot < std::numeric_limits::max()) cuts.push_back(*knot + 1); } const std::int64_t k = knot ? *knot : three; return detail::assemble_easy(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}; }); } /// @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 easy_lut make_leaky_relu_hundredth_lut() { return detail::assemble_easy({0}, [](std::int64_t raw) { if (raw >= 0) return detail::kIdentity; return detail::easy_poly{0, 1, 0, 100}; }); } } // namespace grotto #endif // LIBDPF_INCLUDE_GROTTO_EASY_LUT_HPP__