libdpf/examples/eval_opt_smoke.cpp

368 lines
13 KiB
C++
Raw Normal View History

// Smoke test for the interval-eval optimizations: pipelined interior
// eval01 / eval01_x4, round-major and x4/x8 exterior AES, fused dual-output
// leaf pass, uninitialized output buffers.
#include <cstdint>
#include <cstdio>
#include <cstdlib>
#include <cstring>
#include <limits>
#include "dpf.hpp"
static int fails = 0;
static void expect(bool ok, const char *what)
{
if (!ok)
{
std::fprintf(stderr, "FAIL: %s\n", what);
++fails;
}
}
static bool m128_eq(simde__m128i a, simde__m128i b)
{
return std::memcmp(&a, &b, sizeof(a)) == 0;
}
static void test_aes_batch()
{
using prg = dpf::prg::aes128;
simde__m128i seed = simde_mm_set_epi64x(
static_cast<int64_t>(0xfedcba9876543210ULL),
static_cast<int64_t>(0x0123456789abcdefULL));
auto a0 = prg::eval(seed, 0);
auto a1 = prg::eval(seed, 1);
auto a2 = prg::eval(seed, 2);
auto a3 = prg::eval(seed, 3);
auto kids = prg::eval01(seed);
expect(m128_eq(kids[0], a0), "eval01[0] == eval(seed, 0)");
expect(m128_eq(kids[1], a1), "eval01[1] == eval(seed, 1)");
simde__m128i buf2[2];
prg::eval(seed, buf2, 2, 0);
expect(m128_eq(buf2[0], a0), "batch count=2 pos=0 [0]");
expect(m128_eq(buf2[1], a1), "batch count=2 pos=0 [1]");
simde__m128i buf1[1];
prg::eval(seed, buf1, 1, 3);
expect(m128_eq(buf1[0], a3), "batch count=1 pos=3");
simde__m128i buf4[4];
prg::eval(seed, buf4, 4, 0);
expect(m128_eq(buf4[0], a0) && m128_eq(buf4[1], a1)
&& m128_eq(buf4[2], a2) && m128_eq(buf4[3], a3),
"round-major batch count=4");
simde__m128i buf2p[2];
prg::eval(seed, buf2p, 2, 2);
expect(m128_eq(buf2p[0], a2) && m128_eq(buf2p[1], a3),
"batch count=2 pos=2");
simde__m128i seeds[4];
simde__m128i left[4], right[4];
for (int i = 0; i < 4; ++i)
{
seeds[i] = simde_mm_xor_si128(seed, simde_mm_set_epi64x(0, i + 1));
}
prg::eval01_x4(seeds, left, right);
for (int i = 0; i < 4; ++i)
{
auto kids = prg::eval01(seeds[i]);
expect(m128_eq(left[i], kids[0]) && m128_eq(right[i], kids[1]),
"eval01_x4 matches eval01");
}
simde__m128i x4[4];
prg::eval_x4(seeds, x4, 3);
for (int i = 0; i < 4; ++i)
{
expect(m128_eq(x4[i], prg::eval(seeds[i], 3)),
"eval_x4 matches eval");
}
simde__m128i seeds8[8];
simde__m128i x8[8];
for (int i = 0; i < 8; ++i)
{
seeds8[i] = simde_mm_xor_si128(seed, simde_mm_set_epi64x(i + 9, i + 1));
}
prg::eval_x8(seeds8, x8, 0);
for (int i = 0; i < 8; ++i)
{
expect(m128_eq(x8[i], prg::eval(seeds8[i], 0)),
"eval_x8 matches eval");
}
}
static void test_dual_interval()
{
using input_t = dpf::modint<8>;
using add_t = psnip_uint64_t;
using xor_t = dpf::xor_wrapper<psnip_uint64_t>;
using dpf_t = dpf::utils::dpf_type_t<
dpf::prg::aes128, dpf::prg::aes128, input_t, add_t, xor_t>;
const uint64_t alpha = 37;
const uint64_t beta_add = 0x1111111111111111ULL;
const uint64_t beta_xor = 0xaaaaaaaaaaaaaaaaULL;
auto args = dpf::make_dpfargs(
input_t{static_cast<typename input_t::integral_type>(alpha)},
static_cast<add_t>(beta_add),
xor_t{static_cast<psnip_uint64_t>(beta_xor)});
auto [k0, k1] = dpf::make_dpf(std::move(args));
auto from = std::numeric_limits<input_t>::min();
auto to = std::numeric_limits<input_t>::max();
auto add0 = dpf::make_output_buffer_for_full<0>(k0);
auto xor0 = dpf::make_output_buffer_for_full<1>(k0);
auto add1 = dpf::make_output_buffer_for_full<0>(k1);
auto xor1 = dpf::make_output_buffer_for_full<1>(k1);
auto memo0 = dpf::make_basic_full_memoizer(k0);
auto memo1 = dpf::make_basic_full_memoizer(k1);
auto bufs0 = std::forward_as_tuple(add0, xor0);
auto bufs1 = std::forward_as_tuple(add1, xor1);
dpf::eval_interval<0, 1>(k0, from, to, bufs0, memo0);
dpf::eval_interval<0, 1>(k1, from, to, bufs1, memo1);
const int n = 1 << 8;
int add_hits = 0, xor_hits = 0, add_miss = 0, xor_miss = 0;
for (int x = 0; x < n; ++x)
{
auto in = input_t{static_cast<typename input_t::integral_type>(x)};
uint64_t s_add = static_cast<uint64_t>(add1[x])
- static_cast<uint64_t>(add0[x]);
uint64_t s_xor = static_cast<uint64_t>(static_cast<psnip_uint64_t>(xor_t(xor0[x])))
^ static_cast<uint64_t>(static_cast<psnip_uint64_t>(xor_t(xor1[x])));
auto p0 = dpf::eval_point<0>(k0, in);
auto p1 = dpf::eval_point<0>(k1, in);
uint64_t point_add = static_cast<uint64_t>(*p1) - static_cast<uint64_t>(*p0);
if (x == static_cast<int>(alpha))
{
if (s_add == beta_add) ++add_hits; else ++add_miss;
if (s_xor == beta_xor) ++xor_hits; else ++xor_miss;
expect(point_add == beta_add, "eval_point add at alpha");
}
else
{
if (s_add == 0) ++add_hits; else ++add_miss;
if (s_xor == 0) ++xor_hits; else ++xor_miss;
expect(point_add == 0, "eval_point add off alpha");
}
expect(s_add == point_add, "interval add matches eval_point");
}
expect(add_miss == 0 && add_hits == n, "dual-output additive reconstruct");
expect(xor_miss == 0 && xor_hits == n, "dual-output xor reconstruct");
}
static void test_four_outputs_and_wrap()
{
using input_t = dpf::modint<8>;
using out_t = psnip_uint64_t;
auto args = dpf::make_dpfargs(
input_t{static_cast<typename input_t::integral_type>(5)},
static_cast<out_t>(1), static_cast<out_t>(2),
static_cast<out_t>(3), static_cast<out_t>(4));
auto [k0, k1] = dpf::make_dpf(std::move(args));
auto from = std::numeric_limits<input_t>::min();
auto to = std::numeric_limits<input_t>::max();
auto [bufs0, it0] = dpf::eval_interval<0, 1, 2, 3>(k0, from, to);
auto [bufs1, it1] = dpf::eval_interval<0, 1, 2, 3>(k1, from, to);
const uint64_t want[4] = {1, 2, 3, 4};
for (int i = 0; i < 4; ++i)
{
const auto & a = (i == 0) ? std::get<0>(bufs0)
: (i == 1) ? std::get<1>(bufs0)
: (i == 2) ? std::get<2>(bufs0) : std::get<3>(bufs0);
const auto & b = (i == 0) ? std::get<0>(bufs1)
: (i == 1) ? std::get<1>(bufs1)
: (i == 2) ? std::get<2>(bufs1) : std::get<3>(bufs1);
for (int x = 0; x < 256; ++x)
{
uint64_t s = static_cast<uint64_t>(b[x]) - static_cast<uint64_t>(a[x]);
uint64_t exp = (x == 5) ? want[i] : 0ULL;
if (s != exp)
{
expect(false, "4-output fused reconstruct");
return;
}
}
}
}
static void test_single_and_reuse()
{
using input_t = dpf::modint<8>;
using out_t = psnip_uint64_t;
auto args = dpf::make_dpfargs(
input_t{static_cast<typename input_t::integral_type>(11)},
static_cast<out_t>(7));
auto [k0, k1] = dpf::make_dpf(std::move(args));
auto from = std::numeric_limits<input_t>::min();
auto to = std::numeric_limits<input_t>::max();
auto buf0 = dpf::make_output_buffer_for_full<0>(k0);
auto buf1 = dpf::make_output_buffer_for_full<0>(k1);
auto memo0 = dpf::make_basic_full_memoizer(k0);
auto memo1 = dpf::make_basic_full_memoizer(k1);
dpf::eval_interval<0>(k0, from, to, buf0, memo0);
dpf::eval_interval<0>(k1, from, to, buf1, memo1);
// Reuse the same buffers / memoizer with a second key pair.
auto args2 = dpf::make_dpfargs(
input_t{static_cast<typename input_t::integral_type>(200)},
static_cast<out_t>(99));
auto [k2, k3] = dpf::make_dpf(std::move(args2));
dpf::eval_interval<0>(k2, from, to, buf0, memo0);
dpf::eval_interval<0>(k3, from, to, buf1, memo1);
for (int x = 0; x < 256; ++x)
{
uint64_t s = static_cast<uint64_t>(buf1[x]) - static_cast<uint64_t>(buf0[x]);
uint64_t want = (x == 200) ? 99ULL : 0ULL;
if (s != want)
{
expect(false, "reused buffer/memoizer reconstruct");
return;
}
}
}
static void test_partial_interval()
{
using input_t = dpf::modint<8>;
using out_t = psnip_uint64_t;
auto args = dpf::make_dpfargs(
input_t{static_cast<typename input_t::integral_type>(17)},
static_cast<out_t>(42));
auto [k0, k1] = dpf::make_dpf(std::move(args));
// 12 leaf nodes (24 outputs): hits eval_x8 then eval_x4. Size is a
// multiple of the 64-byte output_buffer alignment (ASan aligned_alloc).
auto from = input_t{static_cast<typename input_t::integral_type>(0)};
auto to = input_t{static_cast<typename input_t::integral_type>(23)};
auto [bufs0, it0] = dpf::eval_interval<0>(k0, from, to);
auto [bufs1, it1] = dpf::eval_interval<0>(k1, from, to);
(void)bufs0;
(void)bufs1;
auto z0 = std::begin(it0);
auto z1 = std::begin(it1);
auto e0 = std::end(it0);
for (int x = 0; z0 != e0; ++x, ++z0, ++z1)
{
uint64_t s = static_cast<uint64_t>(*z1) - static_cast<uint64_t>(*z0);
uint64_t want = (x == 17) ? 42ULL : 0ULL;
if (s != want)
{
expect(false, "partial interval reconstruct");
return;
}
auto in = input_t{static_cast<typename input_t::integral_type>(x)};
auto p0 = dpf::eval_point<0>(k0, in);
auto p1 = dpf::eval_point<0>(k1, in);
uint64_t point = static_cast<uint64_t>(*p1) - static_cast<uint64_t>(*p0);
expect(s == point, "partial interval matches eval_point");
}
}
static void test_inner_product()
{
using input_t = dpf::modint<8>;
using add_t = psnip_uint64_t;
using xor_t = dpf::xor_wrapper<psnip_uint64_t>;
const uint64_t alpha = 19;
const uint64_t beta_add = 7;
const uint64_t beta_xor = 0x5a5a5a5a5a5a5a5aULL;
auto args = dpf::make_dpfargs(
input_t{static_cast<typename input_t::integral_type>(alpha)},
static_cast<add_t>(beta_add),
xor_t{static_cast<psnip_uint64_t>(beta_xor)});
auto [k0, k1] = dpf::make_dpf(std::move(args));
auto from = std::numeric_limits<input_t>::min();
auto to = std::numeric_limits<input_t>::max();
const int n = 1 << 8;
uint64_t w_add[256];
uint64_t w_xor[256];
for (int i = 0; i < n; ++i)
{
w_add[i] = static_cast<uint64_t>(i * 3 + 1);
w_xor[i] = static_cast<uint64_t>(0x1111111111111111ULL * (i + 1));
}
auto add0 = dpf::make_output_buffer_for_full<0>(k0);
auto xor0 = dpf::make_output_buffer_for_full<1>(k0);
auto add1 = dpf::make_output_buffer_for_full<0>(k1);
auto xor1 = dpf::make_output_buffer_for_full<1>(k1);
auto memo0 = dpf::make_basic_full_memoizer(k0);
auto memo1 = dpf::make_basic_full_memoizer(k1);
auto bufs0 = std::forward_as_tuple(add0, xor0);
auto bufs1 = std::forward_as_tuple(add1, xor1);
dpf::eval_interval<0, 1>(k0, from, to, bufs0, memo0);
dpf::eval_interval<0, 1>(k1, from, to, bufs1, memo1);
uint64_t dot_add0 = 0, dot_add1 = 0, dot_xor0 = 0, dot_xor1 = 0;
for (int i = 0; i < n; ++i)
{
dot_add0 += static_cast<uint64_t>(add0[i]) * w_add[i];
dot_add1 += static_cast<uint64_t>(add1[i]) * w_add[i];
dot_xor0 ^= static_cast<uint64_t>(static_cast<psnip_uint64_t>(xor_t(xor0[i])))
& w_xor[i];
dot_xor1 ^= static_cast<uint64_t>(static_cast<psnip_uint64_t>(xor_t(xor1[i])))
& w_xor[i];
}
auto memo0b = dpf::make_basic_full_memoizer(k0);
auto memo1b = dpf::make_basic_full_memoizer(k1);
dpf::eval_prepare_interval(k0, from, to, memo0b);
dpf::eval_prepare_interval(k1, from, to, memo1b);
auto [ip_add0, ip_xor0] = dpf::eval_inner_product<0, 1>(
k0, from, to, std::forward_as_tuple(w_add, w_xor), memo0b);
auto [ip_add1, ip_xor1] = dpf::eval_inner_product<0, 1>(
k1, from, to, std::forward_as_tuple(w_add, w_xor), memo1b);
expect(static_cast<uint64_t>(ip_add0) == dot_add0, "inner product add p0");
expect(static_cast<uint64_t>(ip_add1) == dot_add1, "inner product add p1");
expect(static_cast<uint64_t>(static_cast<psnip_uint64_t>(xor_t(ip_xor0)))
== dot_xor0, "inner product xor p0");
expect(static_cast<uint64_t>(static_cast<psnip_uint64_t>(xor_t(ip_xor1)))
== dot_xor1, "inner product xor p1");
uint64_t recon_add = static_cast<uint64_t>(ip_add1)
- static_cast<uint64_t>(ip_add0);
uint64_t recon_xor = static_cast<uint64_t>(static_cast<psnip_uint64_t>(xor_t(ip_xor0)))
^ static_cast<uint64_t>(static_cast<psnip_uint64_t>(xor_t(ip_xor1)));
expect(recon_add == beta_add * w_add[static_cast<int>(alpha)],
"inner product reconstruct add");
expect(recon_xor == (beta_xor & w_xor[static_cast<int>(alpha)]),
"inner product reconstruct xor");
auto [ip_add0b, ip_add0c] = dpf::eval_full_inner_product<0, 0>(
k0, std::forward_as_tuple(w_add, w_add), memo0b);
expect(static_cast<uint64_t>(ip_add0b) == static_cast<uint64_t>(ip_add0),
"duplicate-output inner product");
expect(static_cast<uint64_t>(ip_add0c) == static_cast<uint64_t>(ip_add0),
"duplicate-output inner product match");
}
int main()
{
test_aes_batch();
test_dual_interval();
test_four_outputs_and_wrap();
test_single_and_reuse();
test_partial_interval();
test_inner_product();
if (fails)
{
std::fprintf(stderr, "%d check(s) failed\n", fails);
return 1;
}
std::puts("eval_opt_smoke: ok");
return 0;
}