166 lines
5.7 KiB
C++
166 lines
5.7 KiB
C++
#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);
|
|
|
|
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);
|
|
});
|
|
}
|
|
}
|
|
|
|
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);
|
|
}
|