#include #include #include "grotto/window_lut.hpp" #include #include namespace { struct sample { int which; unsigned k; std::int64_t raw; std::int64_t y; }; const sample kExact[] = { #include "window_samples.inc" }; const sample kTruth[] = { #include "window_accuracy.inc" }; long double sigmoid(long double x) { return 1.0L / (1.0L + expl(-x)); } long double softplus(long double x) { if (x > 40.0L) return x; if (x < -40.0L) return expl(x); return log1pl(expl(x)); } long double reference(grotto::window which, unsigned k, long double x) { switch (which) { case grotto::window::smoothstep: if (x <= -0.5L) return 0; if (x >= 0.5L) return 1; return -2.0L * x * x * x + 1.5L * x + 0.5L; case grotto::window::sigmoid: return sigmoid(x); case grotto::window::tanh: return tanhl(x); case grotto::window::erf: return erfl(x); case grotto::window::erfc: return erfcl(x); case grotto::window::softplus: return softplus(x); case grotto::window::softminus: return x - softplus(x); case grotto::window::logsigmoid: return -softplus(-x); case grotto::window::gelu: return x * (1.0L + erfl(x / sqrtl(2.0L))) / 2.0L; case grotto::window::silu: return x * sigmoid(x); case grotto::window::mish: return x * tanhl(softplus(x)); case grotto::window::elish: return x < 0 ? expm1l(x) * sigmoid(x) : x * sigmoid(x); case grotto::window::serf: return x * erfl(softplus(x)); case grotto::window::tanhexp: return x * tanhl(expl(x)); case grotto::window::asin: return asinl(x); case grotto::window::acos: return acosl(x); case grotto::window::probit: return 0; case grotto::window::hardelish: if (x <= -1.0L) return 0; if (x >= 1.0L) return x; if (x >= 0.0L) return x * (x + 1.0L) / 2.0L; return expm1l(x) * (x + 1.0L) / 2.0L; case grotto::window::lecun_tanh: return 1.7159L * tanhl(2.0L * x / 3.0L); case grotto::window::one_minus_sigmoid: return 1.0L - sigmoid(x); } return 0; } void expect_close(grotto::window which, unsigned k, std::int64_t raw) { const auto y = grotto::eval_window(which, k, raw); const long double x = ldexpl(static_cast(raw), -static_cast(k)); const long double truth = reference(which, k, x) * ldexpl(1.0L, static_cast(k)); EXPECT_LE(fabsl(static_cast(y) - truth), 1.5L) << static_cast(which) << " k=" << k << " raw=" << raw << " y=" << y; } } // namespace TEST(WindowLut, HornerMatchesGeneratedSamples) { for (const sample & point : kExact) { EXPECT_EQ(grotto::eval_window(static_cast(point.which), point.k, point.raw), point.y) << point.which << " k=" << point.k << " raw=" << point.raw; } } TEST(WindowLut, WithinOneUlpOfLibm) { for (const sample & point : kTruth) { if (point.which == static_cast(grotto::window::probit)) continue; expect_close(static_cast(point.which), point.k, point.raw); } } TEST(WindowLut, ProbitMatchesMpmathRounding) { for (const sample & point : kTruth) { if (point.which != static_cast(grotto::window::probit)) continue; const auto y = grotto::eval_window(grotto::window::probit, point.k, point.raw); EXPECT_LE(std::llabs(y - point.y), 1) << "k=" << point.k << " raw=" << point.raw; } } TEST(WindowLut, SmoothstepKnotsAreExact) { EXPECT_EQ(grotto::eval_window(grotto::window::smoothstep, 8, -128), 0); EXPECT_EQ(grotto::eval_window(grotto::window::smoothstep, 8, 128), 256); EXPECT_EQ(grotto::eval_window(grotto::window::smoothstep, 8, 0), 128); EXPECT_EQ(grotto::eval_window(grotto::window::smoothstep, 16, -100000), 0); EXPECT_EQ(grotto::eval_window(grotto::window::smoothstep, 16, 100000), 1 << 16); } TEST(WindowLut, ExhaustiveLowPrecision) { for (int which_i = 0; which_i <= static_cast(grotto::window::tanhexp); ++which_i) { if (which_i == static_cast(grotto::window::probit)) continue; const auto which = static_cast(which_i); for (unsigned k : {8u, 12u}) { const std::int64_t span = std::int64_t{30} << k; const std::int64_t step = k == 8 ? 1 : 17; for (std::int64_t raw = -span; raw <= span; raw += step) expect_close(which, k, raw); } } for (unsigned k : {8u, 12u}) { const auto one = std::int64_t{1} << k; for (std::int64_t raw = -one; raw <= one; raw += (k == 8 ? 1 : 8)) { expect_close(grotto::window::asin, k, raw); expect_close(grotto::window::acos, k, raw); } const std::int64_t step = k == 8 ? 1 : 17; for (std::int64_t raw = -2 * one; raw <= 2 * one; raw += step) expect_close(grotto::window::hardelish, k, raw); const std::int64_t lecun_span = std::int64_t{20} << k; for (std::int64_t raw = -lecun_span; raw <= lecun_span; raw += step) { expect_close(grotto::window::lecun_tanh, k, raw); expect_close(grotto::window::one_minus_sigmoid, k, raw); } } } TEST(WindowLut, HighPrecisionSamples) { for (unsigned k : {16u, 24u, 32u}) { const auto one = std::int64_t{1} << k; const auto step = std::int64_t{1} << (k - 8); for (std::int64_t raw = -2 * one; raw <= 2 * one; raw += step) expect_close(grotto::window::hardelish, k, raw); const std::int64_t span = std::int64_t{20} << k; const auto wide = std::int64_t{1} << (k - 6); for (std::int64_t raw = -span; raw <= span; raw += wide) { expect_close(grotto::window::lecun_tanh, k, raw); expect_close(grotto::window::one_minus_sigmoid, k, raw); } expect_close(grotto::window::hardelish, k, -one + 1); expect_close(grotto::window::hardelish, k, -1); expect_close(grotto::window::lecun_tanh, k, 1); expect_close(grotto::window::lecun_tanh, k, span); } } TEST(WindowLut, OneMinusSigmoidComplementsSigmoid) { for (unsigned k : {8u, 16u, 32u}) { const auto one = std::int64_t{1} << k; for (std::int64_t raw : {std::int64_t{-8} * one, -one, std::int64_t{0}, one, std::int64_t{8} * one}) { const auto sig = grotto::eval_window(grotto::window::sigmoid, k, raw); const auto comp = grotto::eval_window(grotto::window::one_minus_sigmoid, k, raw); EXPECT_EQ(sig + comp, one) << k << " " << raw; } } } TEST(WindowLut, RejectsBadDomain) { EXPECT_THROW(grotto::eval_window(grotto::window::sigmoid, 7, 0), std::invalid_argument); EXPECT_THROW(grotto::eval_window(grotto::window::asin, 8, 257), std::out_of_range); EXPECT_THROW(grotto::eval_window(grotto::window::probit, 8, 0), std::out_of_range); EXPECT_THROW(grotto::eval_window(grotto::window::probit, 8, 256), std::out_of_range); }