libdpf/include/grotto/easy_lut.hpp

348 lines
12 KiB
C++
Raw Normal View History

/// @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 <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;
HEDLEY_NO_THROW
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;
};
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 <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__