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>
472 lines
18 KiB
C++
472 lines
18 KiB
C++
#include <gtest/gtest.h>
|
|
|
|
#include "dpf/compose.hpp"
|
|
#include "dpf/verifiable.hpp"
|
|
#include "grotto/lut_union.hpp"
|
|
|
|
#include "simde/simde/x86/sse2.h"
|
|
|
|
#include <algorithm>
|
|
#include <cstdint>
|
|
#include <limits>
|
|
#include <stdexcept>
|
|
#include <vector>
|
|
|
|
namespace
|
|
{
|
|
|
|
using grotto::offset_horner_group_add;
|
|
using grotto::offset_horner_group_sub;
|
|
|
|
grotto::piecewise_lut<uint8_t> lut_a()
|
|
{
|
|
grotto::piecewise_lut<uint8_t> lut;
|
|
lut.knots = {0, 10, 50};
|
|
lut.coeff = {{1, 0}, {0, 2}, {7, 1}};
|
|
return lut;
|
|
}
|
|
|
|
grotto::piecewise_lut<uint8_t> lut_b()
|
|
{
|
|
grotto::piecewise_lut<uint8_t> lut;
|
|
lut.knots = {0, 4, 12, 80};
|
|
lut.coeff = {{3, 0, 0}, {1, 0, 1}, {9, 2, 0}, {4, 1, 0}};
|
|
return lut;
|
|
}
|
|
|
|
grotto::piecewise_lut<uint8_t> lut_c()
|
|
{
|
|
grotto::piecewise_lut<uint8_t> lut;
|
|
lut.knots = {0, 7, 90};
|
|
lut.coeff = {{8, 1, 0}, {2, 0, 3}, {1, 1, 1}};
|
|
return lut;
|
|
}
|
|
|
|
std::vector<uint64_t> open_union(const grotto::offset_poly_keys<uint8_t> & mat,
|
|
const grotto::lut_union_plan<uint8_t> & plan)
|
|
{
|
|
const auto s0 = grotto::lut_union_eval<0>(mat, plan);
|
|
const auto s1 = grotto::lut_union_eval<1>(mat, plan);
|
|
EXPECT_EQ(s0.size(), s1.size());
|
|
std::vector<uint64_t> out(s0.size());
|
|
for (std::size_t i = 0; i < s0.size(); ++i)
|
|
out[i] = s0[i] + s1[i];
|
|
return out;
|
|
}
|
|
|
|
template <typename InputT>
|
|
void expect_partition(const grotto::lut_union_plan<InputT> & plan)
|
|
{
|
|
if (plan.knots.empty())
|
|
{
|
|
ADD_FAILURE() << "empty union";
|
|
return;
|
|
}
|
|
for (const auto & func : plan.funcs)
|
|
{
|
|
std::vector<int> seen(plan.knots.size(), 0);
|
|
for (const auto & span : func)
|
|
{
|
|
auto mark = [&](std::size_t j) {
|
|
if (j >= seen.size())
|
|
{
|
|
ADD_FAILURE() << "span index " << j;
|
|
return;
|
|
}
|
|
seen[j] += 1;
|
|
};
|
|
if (span.begin <= span.end)
|
|
{
|
|
for (std::size_t j = span.begin; j < span.end; ++j)
|
|
mark(j);
|
|
}
|
|
else
|
|
{
|
|
for (std::size_t j = span.begin; j < plan.knots.size(); ++j)
|
|
mark(j);
|
|
for (std::size_t j = 0; j < span.end; ++j)
|
|
mark(j);
|
|
}
|
|
}
|
|
for (int bit : seen)
|
|
EXPECT_EQ(bit, 1);
|
|
}
|
|
}
|
|
|
|
} // namespace
|
|
|
|
TEST(LutUnion, WideDomainMapsTheWrapAcrossTheFront)
|
|
{
|
|
grotto::piecewise_lut<uint64_t> a;
|
|
a.knots = {10, 30};
|
|
a.coeff = {{1}, {2}};
|
|
grotto::piecewise_lut<uint64_t> b;
|
|
b.knots = {0, 20};
|
|
b.coeff = {{3}, {4}};
|
|
const auto plan = grotto::make_lut_union_plan(
|
|
std::vector<grotto::piecewise_lut<uint64_t>>{a, b}, uint64_t{0});
|
|
|
|
ASSERT_EQ(plan.knots, (std::vector<uint64_t>{0, 10, 20, 30}));
|
|
EXPECT_EQ(plan.comparisons, 1u);
|
|
EXPECT_EQ(plan.prefix_walks, 1u);
|
|
EXPECT_EQ(plan.depth, 64u);
|
|
EXPECT_EQ(plan.geneval_rounds(), plan.depth);
|
|
EXPECT_EQ(plan.degree, 0u);
|
|
|
|
ASSERT_EQ(plan.funcs[0].size(), 2u);
|
|
EXPECT_EQ(plan.funcs[0][0].begin, 1u);
|
|
EXPECT_EQ(plan.funcs[0][0].end, 3u);
|
|
EXPECT_EQ(plan.funcs[0][1].begin, 3u);
|
|
EXPECT_EQ(plan.funcs[0][1].end, 1u);
|
|
|
|
ASSERT_EQ(plan.funcs[1].size(), 2u);
|
|
EXPECT_EQ(plan.funcs[1][0].begin, 0u);
|
|
EXPECT_EQ(plan.funcs[1][0].end, 2u);
|
|
EXPECT_EQ(plan.funcs[1][1].begin, 2u);
|
|
EXPECT_EQ(plan.funcs[1][1].end, 0u);
|
|
expect_partition(plan);
|
|
}
|
|
|
|
TEST(LutUnion, OneWalkMatchesEachLutOnItsOwnKnots)
|
|
{
|
|
const auto a = lut_a();
|
|
const auto b = lut_b();
|
|
const uint8_t center = 12;
|
|
const auto plan = grotto::make_lut_union_plan({a, b}, uint8_t{3});
|
|
EXPECT_GT(plan.endpoints(), a.knots.size());
|
|
EXPECT_EQ(plan.degree, 2u);
|
|
EXPECT_EQ(plan.lanes, 3u);
|
|
EXPECT_EQ(plan.depth, 8u);
|
|
EXPECT_EQ(plan.comparisons, 1u);
|
|
EXPECT_EQ(plan.prefix_walks, 1u);
|
|
|
|
const auto mat = grotto::make_offset_poly_keys<uint8_t>(center, plan.degree);
|
|
for (int eta = 0; eta < 256; eta += 17)
|
|
{
|
|
const auto e = static_cast<uint8_t>(eta);
|
|
const auto here = grotto::make_lut_union_plan({a, b}, e);
|
|
const auto got = open_union(mat, here);
|
|
ASSERT_EQ(got.size(), 2u);
|
|
EXPECT_EQ(got[0], grotto::offset_poly_clear<uint8_t>(center, a.knots, a.coeff, e))
|
|
<< eta;
|
|
EXPECT_EQ(got[1], grotto::offset_poly_clear<uint8_t>(center, b.knots, b.coeff, e))
|
|
<< eta;
|
|
}
|
|
}
|
|
|
|
TEST(LutUnion, GenevalPlanIsOneComparisonHoweverManyLuts)
|
|
{
|
|
const uint8_t eta = 3;
|
|
const auto few = grotto::make_lut_union_plan({lut_a(), lut_b()}, eta);
|
|
const auto many = grotto::make_lut_union_plan({lut_a(), lut_b(), lut_c()}, eta);
|
|
EXPECT_GT(many.endpoints(), few.endpoints());
|
|
EXPECT_EQ(few.geneval_rounds(), many.geneval_rounds());
|
|
EXPECT_EQ(few.funcs.size(), 2u);
|
|
EXPECT_EQ(many.funcs.size(), 3u);
|
|
|
|
dpf::protocol::composer c0(0);
|
|
grotto::schedule_lut_union(c0, few);
|
|
dpf::protocol::composer c1(0);
|
|
grotto::schedule_lut_union(c1, many);
|
|
const auto p0 = c0.default_plan();
|
|
const auto p1 = c1.default_plan();
|
|
EXPECT_EQ(p0.rounds(), few.depth);
|
|
EXPECT_EQ(p1.rounds(), many.depth);
|
|
EXPECT_EQ(p0.rounds(), p1.rounds());
|
|
EXPECT_EQ(p0.slot_bytes(0), grotto::lut_union_slot_bytes(few.lanes));
|
|
EXPECT_EQ(p1.slot_bytes(0), grotto::lut_union_slot_bytes(many.lanes));
|
|
}
|
|
|
|
TEST(LutUnion, GenevalXorAndAdditiveSharesMatchTheClearLuts)
|
|
{
|
|
const uint8_t center = 12;
|
|
const uint8_t eta = 3;
|
|
const auto a = lut_a();
|
|
const auto b = lut_b();
|
|
const auto plan = grotto::make_lut_union_plan({a, b}, eta);
|
|
const uint8_t share = 0x3c;
|
|
const uint8_t other = static_cast<uint8_t>(center ^ share);
|
|
|
|
const auto xor_got = grotto::geneval_lut_union(share, other, plan);
|
|
EXPECT_EQ(xor_got.eta, eta);
|
|
ASSERT_EQ(xor_got.value0.size(), 2u);
|
|
EXPECT_EQ(xor_got.value0[0] + xor_got.value1[0],
|
|
grotto::offset_poly_clear<uint8_t>(center, a.knots, a.coeff, eta));
|
|
EXPECT_EQ(xor_got.value0[1] + xor_got.value1[1],
|
|
grotto::offset_poly_clear<uint8_t>(center, b.knots, b.coeff, eta));
|
|
|
|
const uint8_t c0 = 5;
|
|
const uint8_t c1 = offset_horner_group_sub(center, c0);
|
|
const auto add_got = grotto::geneval_lut_union(dpf::arith_input, c0, c1, plan);
|
|
EXPECT_EQ(add_got.value0[0] + add_got.value1[0], xor_got.value0[0] + xor_got.value1[0]);
|
|
EXPECT_EQ(add_got.value0[1] + add_got.value1[1], xor_got.value0[1] + xor_got.value1[1]);
|
|
}
|
|
|
|
TEST(LutUnion, EasyAndConstantLutsShareTheUnion)
|
|
{
|
|
const auto relu = grotto::piecewise_from_easy(grotto::make_relu_lut<int8_t>());
|
|
const auto clip = grotto::piecewise_from_easy(
|
|
grotto::make_clip_lut<int8_t>(0, -2, 3));
|
|
grotto::constant_lut<int8_t> sign;
|
|
sign.bounds = {std::numeric_limits<int8_t>::min(), 0};
|
|
sign.values = {-1, 1};
|
|
const auto step = grotto::piecewise_from_constant(sign);
|
|
|
|
const int8_t center = -20;
|
|
const int8_t eta = 15;
|
|
const auto plan = grotto::make_lut_union_plan(
|
|
std::vector<grotto::piecewise_lut<int8_t>>{relu, clip, step}, eta);
|
|
EXPECT_EQ(plan.comparisons, 1u);
|
|
EXPECT_EQ(plan.prefix_walks, 1u);
|
|
EXPECT_GT(plan.endpoints(), relu.knots.size());
|
|
|
|
const auto mat = grotto::make_offset_poly_keys<int8_t>(center, plan.degree);
|
|
const auto s0 = grotto::lut_union_eval<0>(mat, plan);
|
|
const auto s1 = grotto::lut_union_eval<1>(mat, plan);
|
|
const int8_t wrapped = offset_horner_group_add(center, eta);
|
|
const uint64_t opened[3] = {s0[0] + s1[0], s0[1] + s1[1], s0[2] + s1[2]};
|
|
EXPECT_EQ(opened[0], static_cast<uint64_t>(grotto::make_relu_lut<int8_t>()(wrapped)));
|
|
EXPECT_EQ(opened[1], static_cast<uint64_t>(grotto::make_clip_lut<int8_t>(0, -2, 3)(wrapped)));
|
|
EXPECT_EQ(opened[2], static_cast<uint64_t>(sign(wrapped)));
|
|
EXPECT_EQ(opened[0], grotto::offset_poly_clear<int8_t>(center, relu.knots, relu.coeff, eta));
|
|
}
|
|
|
|
TEST(LutUnion, RejectsARoundingDenominatorAndANarrowKey)
|
|
{
|
|
EXPECT_THROW(grotto::piecewise_from_easy(grotto::make_leaky_relu_lut<int8_t>(1)),
|
|
std::invalid_argument);
|
|
EXPECT_THROW(grotto::make_lut_union_plan(
|
|
std::vector<grotto::piecewise_lut<uint8_t>>{}, uint8_t{0}),
|
|
std::invalid_argument);
|
|
|
|
const auto plan = grotto::make_lut_union_plan({lut_b()}, uint8_t{1});
|
|
const auto narrow = grotto::make_offset_poly_keys<uint8_t>(4, 0);
|
|
EXPECT_THROW(grotto::lut_union_eval<0>(narrow, plan), std::invalid_argument);
|
|
|
|
grotto::piecewise_lut<uint8_t> unsorted;
|
|
unsorted.knots = {0, 5, 3};
|
|
unsorted.coeff = {{1}, {1}, {1}};
|
|
EXPECT_THROW(grotto::make_lut_union_plan({unsorted}, uint8_t{0}), std::invalid_argument);
|
|
|
|
grotto::piecewise_lut<uint8_t> ragged;
|
|
ragged.knots = {0, 1};
|
|
ragged.coeff = {{1}, {1, 2}};
|
|
EXPECT_THROW(grotto::make_lut_union_plan({ragged}, uint8_t{0}), std::invalid_argument);
|
|
|
|
grotto::piecewise_lut<uint8_t> empty_row;
|
|
empty_row.knots = {0};
|
|
empty_row.coeff = {{}};
|
|
EXPECT_THROW(grotto::make_lut_union_plan({empty_row}, uint8_t{0}), std::invalid_argument);
|
|
|
|
grotto::piecewise_lut<uint8_t> too_wide;
|
|
too_wide.knots = {0};
|
|
too_wide.coeff = {std::vector<uint64_t>(grotto::offset_poly_max_degree + 2, 1)};
|
|
EXPECT_THROW(grotto::make_lut_union_plan({too_wide}, uint8_t{0}), std::invalid_argument);
|
|
|
|
grotto::easy_lut<int8_t> broken;
|
|
broken.bounds = {0};
|
|
broken.c0 = {0};
|
|
EXPECT_THROW(grotto::piecewise_from_easy(broken), std::invalid_argument);
|
|
|
|
grotto::constant_lut<int8_t> bare;
|
|
bare.bounds = {std::numeric_limits<int8_t>::min()};
|
|
EXPECT_THROW(grotto::piecewise_from_constant(bare), std::invalid_argument);
|
|
|
|
dpf::protocol::composer composer(0);
|
|
EXPECT_THROW(grotto::schedule_lut_union(composer, grotto::lut_union_plan<uint8_t>{}),
|
|
std::invalid_argument);
|
|
}
|
|
|
|
TEST(LutUnion, CoarsePieceCoversSeveralUnionKnots)
|
|
{
|
|
const auto plan = grotto::make_lut_union_plan({lut_a(), lut_b()}, uint8_t{0});
|
|
expect_partition(plan);
|
|
ASSERT_EQ(plan.knots, (std::vector<uint8_t>{0, 4, 10, 12, 50, 80}));
|
|
const auto & first = plan.funcs[0][0];
|
|
EXPECT_EQ(first.end - first.begin, 2u);
|
|
EXPECT_EQ(first.kappa, 0);
|
|
EXPECT_EQ(first.coeff, (std::vector<uint64_t>{1, 0}));
|
|
|
|
const uint8_t center = 200;
|
|
const uint8_t eta = 100;
|
|
const auto carried = grotto::make_lut_union_plan({lut_a(), lut_b()}, eta);
|
|
expect_partition(carried);
|
|
EXPECT_NE(std::find(carried.knots.begin(), carried.knots.end(), uint8_t{156}),
|
|
carried.knots.end());
|
|
const auto mat = grotto::make_offset_poly_keys<uint8_t>(center, carried.degree);
|
|
const auto got = open_union(mat, carried);
|
|
EXPECT_EQ(got[0], grotto::offset_poly_clear<uint8_t>(center, lut_a().knots, lut_a().coeff, eta));
|
|
EXPECT_EQ(got[1], grotto::offset_poly_clear<uint8_t>(center, lut_b().knots, lut_b().coeff, eta));
|
|
}
|
|
|
|
TEST(LutUnion, EveryEtaMatchesASeparatePolynomialEval)
|
|
{
|
|
const auto a = lut_a();
|
|
const auto b = lut_b();
|
|
const uint8_t center = 12;
|
|
const auto probe = grotto::make_lut_union_plan({a, b}, uint8_t{0});
|
|
const auto mat = grotto::make_offset_poly_keys<uint8_t>(center, probe.degree);
|
|
const auto wide = grotto::make_offset_poly_keys<uint8_t>(center, probe.degree + 2);
|
|
const auto alone_a = grotto::make_offset_poly_keys<uint8_t>(center, 1);
|
|
const auto alone_b = grotto::make_offset_poly_keys<uint8_t>(center, 2);
|
|
for (int eta = 0; eta < 256; ++eta)
|
|
{
|
|
const auto e = static_cast<uint8_t>(eta);
|
|
const auto plan = grotto::make_lut_union_plan({a, b}, e);
|
|
expect_partition(plan);
|
|
const auto got = open_union(mat, plan);
|
|
const auto again = open_union(mat, plan);
|
|
EXPECT_EQ(got, again);
|
|
EXPECT_EQ(got, open_union(wide, plan));
|
|
const uint64_t a_alone = grotto::offset_poly_eval<0>(alone_a, a.knots, a.coeff, e)
|
|
+ grotto::offset_poly_eval<1>(alone_a, a.knots, a.coeff, e);
|
|
const uint64_t b_alone = grotto::offset_poly_eval<0>(alone_b, b.knots, b.coeff, e)
|
|
+ grotto::offset_poly_eval<1>(alone_b, b.knots, b.coeff, e);
|
|
EXPECT_EQ(got[0], a_alone) << eta;
|
|
EXPECT_EQ(got[1], b_alone) << eta;
|
|
}
|
|
}
|
|
|
|
TEST(LutUnion, OnePieceLutCoversTheWholeUnion)
|
|
{
|
|
grotto::piecewise_lut<uint8_t> whole;
|
|
whole.knots = {0};
|
|
whole.coeff = {{5, 1}};
|
|
const auto plan = grotto::make_lut_union_plan({whole, lut_a()}, uint8_t{0});
|
|
expect_partition(plan);
|
|
ASSERT_EQ(plan.funcs[0].size(), 1u);
|
|
EXPECT_EQ(plan.funcs[0][0].begin, 0u);
|
|
EXPECT_EQ(plan.funcs[0][0].end, plan.knots.size());
|
|
|
|
// A nonzero eta inserts the carry cut, so the one public knot becomes two
|
|
// refined pieces. They still partition the union.
|
|
const uint8_t eta = 9;
|
|
const auto split = grotto::make_lut_union_plan({whole, lut_a()}, eta);
|
|
expect_partition(split);
|
|
EXPECT_GT(split.funcs[0].size(), 1u);
|
|
const uint8_t center = 40;
|
|
const auto mat = grotto::make_offset_poly_keys<uint8_t>(center, split.degree);
|
|
const auto got = open_union(mat, split);
|
|
EXPECT_EQ(got[0], grotto::offset_poly_clear<uint8_t>(center, whole.knots, whole.coeff, eta));
|
|
EXPECT_EQ(got[1], grotto::offset_poly_clear<uint8_t>(center, lut_a().knots, lut_a().coeff, eta));
|
|
}
|
|
|
|
TEST(LutUnion, IdenticalKnotsStayASingleCopy)
|
|
{
|
|
const auto a = lut_a();
|
|
const auto plan = grotto::make_lut_union_plan({a, a}, uint8_t{0});
|
|
EXPECT_EQ(plan.knots, a.knots);
|
|
EXPECT_EQ(plan.funcs.size(), 2u);
|
|
EXPECT_EQ(plan.prefix_walks, 1u);
|
|
expect_partition(plan);
|
|
}
|
|
|
|
TEST(LutUnion, VerifiableTokensCoverEveryLane)
|
|
{
|
|
const auto a = lut_a();
|
|
const auto b = lut_b();
|
|
const uint8_t center = 12;
|
|
const uint8_t eta = 5;
|
|
const auto plan = grotto::make_lut_union_plan({a, b}, eta);
|
|
const auto mat = grotto::make_offset_poly_keys<uint8_t>(center, plan.degree, dpf::verifiable{});
|
|
std::vector<dpf::proof_token> tok0(plan.lanes), tok1(plan.lanes);
|
|
const auto s0 = grotto::lut_union_eval<0>(mat, plan, tok0.data());
|
|
const auto s1 = grotto::lut_union_eval<1>(mat, plan, tok1.data());
|
|
EXPECT_EQ(s0[0] + s1[0], grotto::offset_poly_clear<uint8_t>(center, a.knots, a.coeff, eta));
|
|
EXPECT_EQ(s0[1] + s1[1], grotto::offset_poly_clear<uint8_t>(center, b.knots, b.coeff, eta));
|
|
for (std::size_t m = 0; m < plan.lanes; ++m)
|
|
EXPECT_TRUE(dpf::verify(tok0[m], tok1[m])) << m;
|
|
tok0[0][0] = simde_mm_xor_si128(tok0[0][0], simde_mm_set1_epi8(1));
|
|
EXPECT_FALSE(dpf::verify(tok0[0], tok1[0]));
|
|
EXPECT_TRUE(dpf::verify(tok0[1], tok1[1]));
|
|
}
|
|
|
|
TEST(LutUnion, ScheduleOnAnExistingSeedStaysOneWalk)
|
|
{
|
|
const auto few = grotto::make_lut_union_plan({lut_a()}, uint8_t{1});
|
|
const auto many = grotto::make_lut_union_plan({lut_a(), lut_b(), lut_c()}, uint8_t{1});
|
|
EXPECT_EQ(few.lanes, 2u);
|
|
EXPECT_EQ(grotto::lut_union_slot_bytes(1), 16u);
|
|
EXPECT_EQ(grotto::lut_union_slot_bytes(few.lanes), 16u);
|
|
EXPECT_EQ(grotto::lut_union_slot_bytes(many.lanes), 24u);
|
|
|
|
dpf::protocol::composer composer(0);
|
|
auto seed = composer.input(dpf::protocol::domain::fss, 16);
|
|
grotto::schedule_lut_union(composer, seed, many);
|
|
const auto scheduled = composer.default_plan();
|
|
ASSERT_EQ(scheduled.rounds(), many.depth);
|
|
for (std::uint16_t round = 0; round < scheduled.rounds(); ++round)
|
|
EXPECT_EQ(scheduled.slot_bytes(round), grotto::lut_union_slot_bytes(many.lanes));
|
|
}
|
|
|
|
TEST(LutUnion, SignedDomainAndRealTables)
|
|
{
|
|
const auto relu = grotto::piecewise_from_easy(grotto::make_relu_lut<int8_t>());
|
|
const auto abs_lut = grotto::piecewise_from_easy(grotto::make_abs_lut<int8_t>());
|
|
const auto square = grotto::piecewise_from_easy(grotto::make_squared_relu_lut<int8_t>(0));
|
|
for (const auto & row : relu.coeff)
|
|
EXPECT_EQ(row.size(), 2u);
|
|
for (const auto & row : abs_lut.coeff)
|
|
EXPECT_EQ(row.size(), 2u);
|
|
for (const auto & row : square.coeff)
|
|
EXPECT_EQ(row.size(), 3u);
|
|
|
|
const int8_t center = -20;
|
|
const auto probe = grotto::make_lut_union_plan(
|
|
std::vector<grotto::piecewise_lut<int8_t>>{relu, abs_lut, square}, int8_t{0});
|
|
EXPECT_EQ(probe.degree, 2u);
|
|
EXPECT_EQ(probe.depth, 8u);
|
|
const auto mat = grotto::make_offset_poly_keys<int8_t>(center, probe.degree);
|
|
const auto relu_f = grotto::make_relu_lut<int8_t>();
|
|
const auto abs_f = grotto::make_abs_lut<int8_t>();
|
|
const auto sq_f = grotto::make_squared_relu_lut<int8_t>(0);
|
|
for (int eta = -128; eta < 128; eta += 7)
|
|
{
|
|
const auto e = static_cast<int8_t>(eta);
|
|
const auto plan = grotto::make_lut_union_plan(
|
|
std::vector<grotto::piecewise_lut<int8_t>>{relu, abs_lut, square}, e);
|
|
expect_partition(plan);
|
|
const auto s0 = grotto::lut_union_eval<0>(mat, plan);
|
|
const auto s1 = grotto::lut_union_eval<1>(mat, plan);
|
|
const int8_t wrapped = offset_horner_group_add(center, e);
|
|
EXPECT_EQ(s0[0] + s1[0], static_cast<uint64_t>(relu_f(wrapped))) << eta;
|
|
EXPECT_EQ(s0[1] + s1[1], static_cast<uint64_t>(abs_f(wrapped))) << eta;
|
|
EXPECT_EQ(s0[2] + s1[2], static_cast<uint64_t>(sq_f(wrapped))) << eta;
|
|
}
|
|
|
|
const int8_t share = 3;
|
|
const int8_t other = static_cast<int8_t>(center ^ share);
|
|
const auto plan = grotto::make_lut_union_plan(
|
|
std::vector<grotto::piecewise_lut<int8_t>>{relu, abs_lut, square}, int8_t{-3});
|
|
const auto got = grotto::geneval_lut_union(share, other, plan);
|
|
const int8_t wrapped = offset_horner_group_add(center, int8_t{-3});
|
|
EXPECT_EQ(got.value0[0] + got.value1[0], static_cast<uint64_t>(relu_f(wrapped)));
|
|
EXPECT_EQ(got.value0[1] + got.value1[1], static_cast<uint64_t>(abs_f(wrapped)));
|
|
EXPECT_EQ(got.value0[2] + got.value1[2], static_cast<uint64_t>(sq_f(wrapped)));
|
|
}
|
|
|
|
TEST(LutUnion, SixteenBitDomainKeepsOneComparison)
|
|
{
|
|
grotto::piecewise_lut<uint16_t> left;
|
|
left.knots = {0, 1000};
|
|
left.coeff = {{3, 1}, {8, 0}};
|
|
grotto::piecewise_lut<uint16_t> right;
|
|
right.knots = {0, 400, 2000};
|
|
right.coeff = {{1, 0, 2}, {9, 1, 0}, {4, 0, 1}};
|
|
const uint16_t center = 50;
|
|
const uint16_t eta = 40000;
|
|
const auto plan = grotto::make_lut_union_plan({left, right}, eta);
|
|
EXPECT_EQ(plan.depth, 16u);
|
|
EXPECT_EQ(plan.geneval_rounds(), 16u);
|
|
EXPECT_EQ(plan.comparisons, 1u);
|
|
expect_partition(plan);
|
|
const auto mat = grotto::make_offset_poly_keys<uint16_t>(center, plan.degree);
|
|
const auto s0 = grotto::lut_union_eval<0>(mat, plan);
|
|
const auto s1 = grotto::lut_union_eval<1>(mat, plan);
|
|
EXPECT_EQ(s0[0] + s1[0], grotto::offset_poly_clear<uint16_t>(center, left.knots, left.coeff, eta));
|
|
EXPECT_EQ(s0[1] + s1[1], grotto::offset_poly_clear<uint16_t>(center, right.knots, right.coeff, eta));
|
|
|
|
dpf::protocol::composer composer(0);
|
|
grotto::schedule_lut_union(composer, plan);
|
|
EXPECT_EQ(composer.default_plan().rounds(), 16u);
|
|
}
|