Ship the TLS mesh, composer, Beaver/Yao/leaf MPC, prep/online paths, apps, and docs so the tree is pushable before elevating share_expr, security_mode, and prep resume. Co-authored-by: Cursor <cursoragent@cursor.com>
224 lines
7.1 KiB
C++
224 lines
7.1 KiB
C++
#include <gtest/gtest.h>
|
|
#include <tuple>
|
|
|
|
#include "grotto/window_lut.hpp"
|
|
|
|
#include <cmath>
|
|
#include <cstdint>
|
|
|
|
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<long double>(raw), -static_cast<int>(k));
|
|
const long double truth = reference(which, k, x) * ldexpl(1.0L, static_cast<int>(k));
|
|
EXPECT_LE(fabsl(static_cast<long double>(y) - truth), 1.5L)
|
|
<< static_cast<int>(which) << " k=" << k << " raw=" << raw << " y=" << y;
|
|
}
|
|
|
|
} // namespace
|
|
|
|
TEST(WindowLut, HornerMatchesGeneratedSamples)
|
|
{
|
|
for (const sample & point : kExact)
|
|
{
|
|
EXPECT_EQ(grotto::eval_window(static_cast<grotto::window>(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<int>(grotto::window::probit))
|
|
continue;
|
|
expect_close(static_cast<grotto::window>(point.which), point.k, point.raw);
|
|
}
|
|
}
|
|
|
|
TEST(WindowLut, ProbitMatchesMpmathRounding)
|
|
{
|
|
for (const sample & point : kTruth)
|
|
{
|
|
if (point.which != static_cast<int>(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<int>(grotto::window::tanhexp); ++which_i)
|
|
{
|
|
if (which_i == static_cast<int>(grotto::window::probit))
|
|
continue;
|
|
const auto which = static_cast<grotto::window>(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);
|
|
}
|