libdpf/test/tests/easy_lut_test.cpp

172 lines
6 KiB
C++
Raw Normal View History

#include <gtest/gtest.h>
#include "grotto/easy_lut.hpp"
#include <cstdint>
#include <limits>
namespace
{
std::int64_t round_div(__int128 num, __int128 den)
{
const bool neg = num < 0;
const __int128 mag = neg ? -num : num;
const __int128 q = (mag + den / 2) / den;
return static_cast<std::int64_t>(neg ? -q : q);
}
std::int64_t ref_abs(std::int64_t raw) { return raw < 0 ? -raw : raw; }
std::int64_t ref_relu(std::int64_t raw) { return raw < 0 ? 0 : raw; }
std::int64_t ref_clip(std::int64_t raw, std::int64_t low, std::int64_t high)
{
if (raw < low)
return low;
if (raw > high)
return high;
return raw;
}
std::int64_t ref_hardsigmoid(std::int64_t raw, std::int64_t three, std::int64_t one)
{
if (raw <= -three)
return 0;
if (raw >= three)
return one;
return round_div(raw + three, 6);
}
std::int64_t ref_hardswish(std::int64_t raw, std::int64_t three, std::int64_t den)
{
if (raw <= -three)
return 0;
if (raw >= three)
return raw;
return round_div(__int128(raw) * (raw + three), den);
}
std::int64_t ref_sqrelu(std::int64_t raw, std::int64_t den)
{
if (raw < 0)
return 0;
return round_div(__int128(raw) * raw, den);
}
std::int64_t ref_soft(std::int64_t raw, std::int64_t knot)
{
if (raw > knot)
return raw - knot;
if (raw < -knot)
return raw + knot;
return 0;
}
std::int64_t ref_hardshrink(std::int64_t raw, std::int64_t knot)
{
if (raw > knot || raw < -knot)
return raw;
return 0;
}
template <typename Raw, typename Fn>
void expect_all(const grotto::easy_lut<Raw> & lut, Fn && ref)
{
using lim = std::numeric_limits<Raw>;
for (std::int64_t raw = lim::min(); raw <= lim::max(); ++raw)
{
EXPECT_EQ(lut(static_cast<Raw>(raw)), ref(raw)) << raw;
if (raw == lim::max())
break;
}
}
} // namespace
TEST(EasyLut, FewPiecesAtEveryPrecision)
{
EXPECT_EQ((grotto::make_abs_lut<std::int16_t>().parts()), 2u);
EXPECT_EQ((grotto::make_relu_lut<std::int16_t>().parts()), 2u);
EXPECT_EQ((grotto::make_squared_relu_lut<std::int16_t>(0).parts()), 2u);
EXPECT_EQ((grotto::make_leaky_relu_lut<std::int16_t>(0).parts()), 1u);
EXPECT_EQ((grotto::make_leaky_relu_lut<std::int16_t>(3).parts()), 2u);
EXPECT_EQ((grotto::make_leaky_relu_hundredth_lut<std::int16_t>().parts()), 2u);
for (unsigned fractional_bits : {0u, 2u, 4u})
{
EXPECT_EQ((grotto::make_relu6_lut<std::int16_t>(fractional_bits).parts()), 3u);
EXPECT_EQ((grotto::make_hardtanh_lut<std::int16_t>(fractional_bits).parts()), 3u);
EXPECT_EQ((grotto::make_softshrink_lut<std::int16_t>(fractional_bits).parts()), 3u);
EXPECT_EQ((grotto::make_hardshrink_lut<std::int16_t>(fractional_bits).parts()), 3u);
EXPECT_EQ((grotto::make_hardsigmoid_lut<std::int16_t>(fractional_bits).parts()), 3u);
EXPECT_EQ((grotto::make_hardswish_lut<std::int16_t>(fractional_bits).parts()), 3u);
}
// On int8 with 7 fractional bits, ±1 and ±3 lie outside the domain.
EXPECT_EQ((grotto::make_hardtanh_lut<std::int8_t>(7).parts()), 1u);
EXPECT_EQ((grotto::make_hardsigmoid_lut<std::int8_t>(7).parts()), 1u);
}
TEST(EasyLut, ExactLinesMatchOnInt16)
{
for (unsigned fractional_bits : {0u, 3u})
{
const std::int64_t scale = std::int64_t{1} << fractional_bits;
const auto abs_lut = grotto::make_abs_lut<std::int16_t>(fractional_bits);
const auto relu = grotto::make_relu_lut<std::int16_t>(fractional_bits);
const auto relu6 = grotto::make_relu6_lut<std::int16_t>(fractional_bits);
const auto tanh = grotto::make_hardtanh_lut<std::int16_t>(fractional_bits);
const auto soft = grotto::make_softshrink_lut<std::int16_t>(fractional_bits);
const auto shrink = grotto::make_hardshrink_lut<std::int16_t>(fractional_bits);
const auto sig = grotto::make_hardsigmoid_lut<std::int16_t>(fractional_bits);
const auto swish = grotto::make_hardswish_lut<std::int16_t>(fractional_bits);
const auto sq = grotto::make_squared_relu_lut<std::int16_t>(fractional_bits);
const auto leak = grotto::make_leaky_relu_lut<std::int16_t>(2);
const std::int64_t three = 3 * scale;
const std::int64_t den = 6 * scale;
expect_all<std::int16_t>(abs_lut, ref_abs);
expect_all<std::int16_t>(relu, ref_relu);
expect_all<std::int16_t>(relu6, [&](std::int64_t raw) {
return ref_clip(raw, 0, 6 * scale);
});
expect_all<std::int16_t>(tanh, [&](std::int64_t raw) {
return ref_clip(raw, -scale, scale);
});
expect_all<std::int16_t>(soft, [&](std::int64_t raw) { return ref_soft(raw, scale); });
expect_all<std::int16_t>(shrink, [&](std::int64_t raw) { return ref_hardshrink(raw, scale); });
expect_all<std::int16_t>(sig, [&](std::int64_t raw) {
return ref_hardsigmoid(raw, three, scale);
});
expect_all<std::int16_t>(swish, [&](std::int64_t raw) {
return ref_hardswish(raw, three, den);
});
expect_all<std::int16_t>(sq, [&](std::int64_t raw) { return ref_sqrelu(raw, scale); });
expect_all<std::int16_t>(leak, [](std::int64_t raw) {
return raw >= 0 ? raw : round_div(raw, 4);
});
}
const auto hundredth = grotto::make_leaky_relu_hundredth_lut<std::int16_t>();
expect_all<std::int16_t>(hundredth, [](std::int64_t raw) {
return raw >= 0 ? raw : round_div(raw, 100);
});
}
TEST(EasyLut, ClipIsTheGeneralProgram)
{
const auto relu6 = grotto::make_relu6_lut<std::int16_t>(2);
const auto clip = grotto::make_clip_lut<std::int16_t>(2, 0, 6);
EXPECT_EQ(relu6.parts(), clip.parts());
EXPECT_EQ(relu6.bounds, clip.bounds);
EXPECT_EQ(relu6.c0, clip.c0);
EXPECT_EQ(relu6.c1, clip.c1);
}
TEST(EasyLut, AbsAndReluIgnoreFractionalWidth)
{
const auto a0 = grotto::make_abs_lut<std::int32_t>(0);
const auto a8 = grotto::make_abs_lut<std::int32_t>(8);
EXPECT_EQ(a0.c0, a8.c0);
EXPECT_EQ(a0.c1, a8.c1);
EXPECT_EQ(a0.bounds, a8.bounds);
}