#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; } 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); } } } 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); }