libdpf/test/tests/fixedpoint_beaver_test.cpp

311 lines
10 KiB
C++
Raw Normal View History

#include <gtest/gtest.h>
#include <condition_variable>
#include <cstdint>
#include <exception>
#include <limits>
#include <mutex>
#include <thread>
#include <type_traits>
#include <utility>
#include "grotto.hpp"
namespace
{
struct bus
{
std::mutex mu;
std::condition_variable cv;
std::uint64_t word[2]{};
std::uint64_t deliver[2]{};
int epoch = 0;
int arrived = 0;
int sent0 = 0;
std::uint64_t exchange(int me, std::uint64_t mine)
{
std::unique_lock<std::mutex> lock(mu);
const int seen = epoch;
word[me] = mine;
if (me == 0)
++sent0;
if (++arrived == 2)
{
deliver[0] = word[1];
deliver[1] = word[0];
arrived = 0;
++epoch;
cv.notify_all();
}
else
cv.wait(lock, [&] { return epoch != seen; });
return deliver[me];
}
};
template <typename T>
std::pair<T, T> split_word(T secret)
{
using U = std::make_unsigned_t<T>;
const U mask = static_cast<U>(~U{0});
const U bits = static_cast<U>(secret);
const U share = static_cast<U>(dpf::uniform_sample<U>() & mask);
const U other = static_cast<U>(bits - share);
return {static_cast<T>(share), static_cast<T>(other)};
}
template <typename Fixed>
simde_uint128 raw_bits(Fixed value)
{
std::uint64_t limb[4] = {};
grotto::detail::store_raw_limbs(value.integral_representation(), limb);
return simde_uint128{limb[0]} | (simde_uint128{limb[1]} << 64);
}
template <unsigned I, unsigned F, unsigned Lf, typename L, unsigned Rf, typename R>
void expect_product(std::pair<L, L> lhs, std::pair<R, R> rhs)
{
using shape = grotto::fixed_mul_beaver_shape<I, F, Lf, L, Rf, R>;
ASSERT_TRUE(shape::fits);
static const auto prep = grotto::make_fixed_mul_beaver_prep<I, F, Lf, L, Rf, R>();
const auto triple = grotto::sample_fixed_mul_beaver_triple(shape::plan::multiply_bits);
bus link;
using result = typename shape::plan::result_type;
result y0{};
result y1{};
std::exception_ptr e0;
std::exception_ptr e1;
std::thread t0([&] {
try
{
y0 = grotto::eval_fixed_mul_beaver<I, F, Lf, L, Rf, R>(
prep, triple, 0, lhs.first, rhs.first,
[&](std::uint64_t mine) { return link.exchange(0, mine); });
}
catch (...)
{
e0 = std::current_exception();
}
});
std::thread t1([&] {
try
{
y1 = grotto::eval_fixed_mul_beaver<I, F, Lf, L, Rf, R>(
prep, triple, 1, lhs.second, rhs.second,
[&](std::uint64_t mine) { return link.exchange(1, mine); });
}
catch (...)
{
e1 = std::current_exception();
}
});
t0.join();
t1.join();
if (e0)
std::rethrow_exception(e0);
if (e1)
std::rethrow_exception(e1);
EXPECT_EQ(link.sent0, static_cast<int>(shape::messages));
const auto clear = grotto::fixed_mul<I, F>(
grotto::make_fixed_from_integral_type<Lf, L>(
static_cast<L>(static_cast<std::make_unsigned_t<L>>(lhs.first)
+ static_cast<std::make_unsigned_t<L>>(lhs.second))),
grotto::make_fixed_from_integral_type<Rf, R>(
static_cast<R>(static_cast<std::make_unsigned_t<R>>(rhs.first)
+ static_cast<std::make_unsigned_t<R>>(rhs.second))));
const unsigned storage = shape::storage_bits;
const simde_uint128 mask = storage >= 128u
? ~simde_uint128{0}
: (simde_uint128{1} << storage) - 1;
const auto got = (raw_bits(y0) + raw_bits(y1)) & mask;
const auto want = raw_bits(clear) & mask;
EXPECT_EQ(static_cast<std::uint64_t>(got), static_cast<std::uint64_t>(want));
EXPECT_EQ(static_cast<std::uint64_t>(got >> 64),
static_cast<std::uint64_t>(want >> 64));
}
template <typename T>
T reconstruct(std::pair<T, T> shares)
{
using U = std::make_unsigned_t<T>;
return static_cast<T>(static_cast<U>(shares.first) + static_cast<U>(shares.second));
}
} // namespace
TEST(FixedMulBeaver, EmptyWindowIsZeroAndSilent)
{
using shape = grotto::fixed_mul_beaver_shape<0, 4, 0, std::int32_t, 0, std::int32_t>;
EXPECT_TRUE(shape::fits);
EXPECT_FALSE(shape::active);
EXPECT_EQ(shape::messages, 0u);
auto prep = grotto::make_fixed_mul_beaver_prep<0, 4, 0, std::int32_t, 0, std::int32_t>();
grotto::fixed_mul_beaver_triple triple;
int calls = 0;
auto y = grotto::eval_fixed_mul_beaver<0, 4, 0, std::int32_t, 0, std::int32_t>(
prep, triple, 0, 7, 9, [&](std::uint64_t) {
++calls;
return std::uint64_t{0};
});
EXPECT_EQ(calls, 0);
EXPECT_EQ(y.integral_representation(), 0);
}
TEST(FixedMulBeaver, WideSignedWindowDoesNotFit)
{
using shape = grotto::fixed_mul_beaver_shape<16, 64, 16, std::int32_t, 16, std::int32_t>;
EXPECT_FALSE(shape::fits);
EXPECT_TRUE(shape::result_lift);
EXPECT_GT(shape::plan::out_bits, 64u);
}
TEST(FixedMulBeaver, TruncationIsOnlyTheBeaverOpening)
{
using L = std::uint32_t;
using shape = grotto::fixed_mul_beaver_shape<16, 0, 0, L, 0, L>;
EXPECT_TRUE(shape::fits);
EXPECT_FALSE(shape::lhs_lift);
EXPECT_FALSE(shape::shift_right);
EXPECT_EQ(shape::messages, 2u);
auto prep = grotto::make_fixed_mul_beaver_prep<16, 0, 0, L, 0, L>();
for (int i = 0; i < 8; ++i)
{
const auto lhs = split_word(i == 0 ? L{7} : dpf::uniform_sample<L>());
const auto rhs = split_word(i == 0 ? L{9} : dpf::uniform_sample<L>());
const auto triple = grotto::sample_fixed_mul_beaver_triple(16u);
bus link;
auto y0 = grotto::fixedpoint<0, std::uint16_t>{};
auto y1 = y0;
std::thread t0([&] {
y0 = grotto::eval_fixed_mul_beaver<16, 0, 0, L, 0, L>(
prep, triple, 0, lhs.first, rhs.first,
[&](std::uint64_t mine) { return link.exchange(0, mine); });
});
std::thread t1([&] {
y1 = grotto::eval_fixed_mul_beaver<16, 0, 0, L, 0, L>(
prep, triple, 1, lhs.second, rhs.second,
[&](std::uint64_t mine) { return link.exchange(1, mine); });
});
t0.join();
t1.join();
EXPECT_EQ(link.sent0, 2);
const auto clear = grotto::fixed_mul<16, 0>(
grotto::make_fixed_from_integral_type<0, L>(reconstruct(lhs)),
grotto::make_fixed_from_integral_type<0, L>(reconstruct(rhs)));
const auto got = static_cast<std::uint16_t>(
static_cast<std::uint16_t>(y0.integral_representation())
+ static_cast<std::uint16_t>(y1.integral_representation()));
EXPECT_EQ(got, clear.integral_representation());
}
}
TEST(FixedMulBeaver, NegativeProductFloorsIntoTheWindow)
{
using Q = std::int32_t;
const auto lhs = split_word(Q{-3});
const auto rhs = split_word(Q{1});
expect_product<8, 4, 4, Q, 4, Q>(lhs, rhs);
const auto neg = split_word(Q{-2});
expect_product<8, 4, 4, Q, 4, Q>(lhs, neg);
}
TEST(FixedMulBeaver, SignedLiftAndRightShift)
{
using Q = std::int32_t;
using shape = grotto::fixed_mul_beaver_shape<16, 16, 16, Q, 16, Q>;
EXPECT_TRUE(shape::lhs_lift);
EXPECT_TRUE(shape::shift_right);
EXPECT_FALSE(shape::product_lift);
EXPECT_EQ(shape::plan::multiply_bits, 48u);
EXPECT_EQ(shape::messages, 5u);
auto prep = grotto::make_fixed_mul_beaver_prep<16, 16, 16, Q, 16, Q>();
const Q samples[][2] = {
{0, 0},
{1 << 16, 2 << 16},
{-(3 << 16), 1 << 16},
{-(3 << 16), -(2 << 16)},
{std::numeric_limits<Q>::max(), -1},
{std::numeric_limits<Q>::min(), 1},
};
for (const auto & sample : samples)
{
const auto triple = grotto::sample_fixed_mul_beaver_triple(48u);
const auto lhs = split_word(sample[0]);
const auto rhs = split_word(sample[1]);
bus link;
using result = typename shape::plan::result_type;
result y0{};
result y1{};
std::thread t0([&] {
y0 = grotto::eval_fixed_mul_beaver<16, 16, 16, Q, 16, Q>(
prep, triple, 0, lhs.first, rhs.first,
[&](std::uint64_t mine) { return link.exchange(0, mine); });
});
std::thread t1([&] {
y1 = grotto::eval_fixed_mul_beaver<16, 16, 16, Q, 16, Q>(
prep, triple, 1, lhs.second, rhs.second,
[&](std::uint64_t mine) { return link.exchange(1, mine); });
});
t0.join();
t1.join();
EXPECT_EQ(link.sent0, 5);
const auto clear = grotto::fixed_mul<16, 16>(
grotto::make_fixed_from_integral_type<16, Q>(reconstruct(lhs)),
grotto::make_fixed_from_integral_type<16, Q>(reconstruct(rhs)));
const auto got = static_cast<std::uint32_t>(y0.integral_representation())
+ static_cast<std::uint32_t>(y1.integral_representation());
EXPECT_EQ(got, static_cast<std::uint32_t>(clear.integral_representation()))
<< "lhs=" << sample[0] << " rhs=" << sample[1];
}
}
TEST(FixedMulBeaver, LeftShiftAndUnsignedLift)
{
using L = std::uint16_t;
using shape = grotto::fixed_mul_beaver_shape<4, 8, 0, L, 0, L>;
EXPECT_TRUE(shape::shift_left);
EXPECT_TRUE(shape::fits);
const auto lhs = split_word(L{3});
const auto rhs = split_word(L{5});
expect_product<4, 8, 0, L, 0, L>(lhs, rhs);
}
TEST(FixedMulBeaver, NarrowProductIsSignExtendedBeforeTheWindow)
{
using Q = std::int8_t;
using shape = grotto::fixed_mul_beaver_shape<16, 8, 4, Q, 4, Q>;
EXPECT_TRUE(shape::fits);
EXPECT_TRUE(shape::lhs_lift);
EXPECT_TRUE(shape::product_lift);
EXPECT_TRUE(shape::result_lift);
EXPECT_FALSE(shape::shift_right);
EXPECT_EQ(shape::messages, 6u);
const Q samples[] = {0, 1, -1, -3, 7, -8, 12, std::numeric_limits<Q>::max(),
std::numeric_limits<Q>::min()};
for (Q a : samples)
for (Q b : samples)
expect_product<16, 8, 4, Q, 4, Q>(split_word(a), split_word(b));
}
TEST(FixedMulBeaver, Int64WindowUsesA96BitProduct)
{
using Q = std::int64_t;
using shape = grotto::fixed_mul_beaver_shape<32, 32, 32, Q, 32, Q>;
ASSERT_TRUE(shape::fits);
EXPECT_EQ(shape::plan::multiply_bits, 96u);
EXPECT_EQ(shape::plan::modulus_bits, 96u);
EXPECT_TRUE(shape::lhs_lift);
EXPECT_EQ(shape::messages, 8u);
const Q samples[][2] = {
{0, 0},
{Q{3} << 32, Q{5} << 32},
{-(Q{3} << 32), Q{1} << 32},
{-(Q{4} << 32), -(Q{2} << 32)},
};
for (const auto & sample : samples)
expect_product<32, 32, 32, Q, 32, Q>(split_word(sample[0]), split_word(sample[1]));
}