// 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 #include #include #include #include #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(0xfedcba9876543210ULL), static_cast(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; 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(alpha)}, static_cast(beta_add), xor_t{static_cast(beta_xor)}); auto [k0, k1] = dpf::make_dpf(std::move(args)); auto from = std::numeric_limits::min(); auto to = std::numeric_limits::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(x)}; uint64_t s_add = static_cast(add1[x]) - static_cast(add0[x]); uint64_t s_xor = static_cast(static_cast(xor_t(xor0[x]))) ^ static_cast(static_cast(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(*p1) - static_cast(*p0); if (x == static_cast(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(5)}, static_cast(1), static_cast(2), static_cast(3), static_cast(4)); auto [k0, k1] = dpf::make_dpf(std::move(args)); auto from = std::numeric_limits::min(); auto to = std::numeric_limits::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(b[x]) - static_cast(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(11)}, static_cast(7)); auto [k0, k1] = dpf::make_dpf(std::move(args)); auto from = std::numeric_limits::min(); auto to = std::numeric_limits::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(200)}, static_cast(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(buf1[x]) - static_cast(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(17)}, static_cast(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(0)}; auto to = input_t{static_cast(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(*z1) - static_cast(*z0); uint64_t want = (x == 17) ? 42ULL : 0ULL; if (s != want) { expect(false, "partial interval reconstruct"); return; } auto in = input_t{static_cast(x)}; auto p0 = dpf::eval_point<0>(k0, in); auto p1 = dpf::eval_point<0>(k1, in); uint64_t point = static_cast(*p1) - static_cast(*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; 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(alpha)}, static_cast(beta_add), xor_t{static_cast(beta_xor)}); auto [k0, k1] = dpf::make_dpf(std::move(args)); auto from = std::numeric_limits::min(); auto to = std::numeric_limits::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(i * 3 + 1); w_xor[i] = static_cast(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(add0[i]) * w_add[i]; dot_add1 += static_cast(add1[i]) * w_add[i]; dot_xor0 ^= static_cast(static_cast(xor_t(xor0[i]))) & w_xor[i]; dot_xor1 ^= static_cast(static_cast(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(ip_add0) == dot_add0, "inner product add p0"); expect(static_cast(ip_add1) == dot_add1, "inner product add p1"); expect(static_cast(static_cast(xor_t(ip_xor0))) == dot_xor0, "inner product xor p0"); expect(static_cast(static_cast(xor_t(ip_xor1))) == dot_xor1, "inner product xor p1"); uint64_t recon_add = static_cast(ip_add1) - static_cast(ip_add0); uint64_t recon_xor = static_cast(static_cast(xor_t(ip_xor0))) ^ static_cast(static_cast(xor_t(ip_xor1))); expect(recon_add == beta_add * w_add[static_cast(alpha)], "inner product reconstruct add"); expect(recon_xor == (beta_xor & w_xor[static_cast(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(ip_add0b) == static_cast(ip_add0), "duplicate-output inner product"); expect(static_cast(ip_add0c) == static_cast(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; }