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