368 lines
13 KiB
C++
368 lines
13 KiB
C++
|
|
// 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;
|
||
|
|
}
|