libdpf/examples/evaluation/eval_inner_product.cpp
Ryan Henry 0d22946a0e Checkpoint the party/runtime stack before share-program and malicious-mode work.
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>
2026-09-28 05:59:19 -06:00

102 lines
3.8 KiB
C++
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

#include <array>
#include <cstdint>
#include <iostream>
#include <tuple>
#include <vector>
#include "dpf.hpp"
/// Fused inner product: do not materialize the DPF vector.
/// A scalar weight vector dots with one output. A row of a tuple or
/// `std::array` dots with several outputs, including an ancestor slot
/// and the leaf, read off one path.
int main()
{
using In = std::uint8_t;
//! [eval-inner-product-scalar]
// Trivial: sum_x DPF(x) * w[x] over a short interval.
const In alpha = 42;
const std::uint64_t beta = 7;
auto [k0, k1] = dpf::make_dpf(alpha, beta);
const In from = 40;
const In to = 50;
std::vector<std::uint64_t> w(to - from + 1);
for (std::size_t i = 0; i < w.size(); ++i)
w[i] = i + 1;
const auto s0 = dpf::eval_inner_product(dpf::paired, k0, from, to, w);
const auto s1 = dpf::eval_inner_product(dpf::paired, k1, from, to, w);
//! [eval-inner-product-scalar]
if (dpf::reconstruct(s0, s1) != beta * w[alpha - from])
{
std::cerr << "scalar interval\n";
return 1;
}
//! [eval-inner-product-full]
// Trivial full domain. Only α contributes.
std::vector<std::uint64_t> wall(256, 1);
const auto f0 = dpf::eval_full_inner_product(dpf::paired, k0, wall);
const auto f1 = dpf::eval_full_inner_product(dpf::paired, k1, wall);
//! [eval-inner-product-full]
if (dpf::reconstruct(f0, f1) != beta)
{
std::cerr << "full\n";
return 1;
}
//! [eval-inner-product-paired]
// Two outputs on the same leaf. rows[i] = {weight for output 0, output 1}.
auto [p0, p1] = dpf::make_dpf(In{9}, std::uint32_t{3}, std::uint32_t{5});
std::vector<std::array<std::uint32_t, 2>> rows;
for (In x = 8;; ++x)
{
rows.push_back({std::uint32_t{1}, std::uint32_t{x}});
if (x == 10)
break;
}
const auto a0 = dpf::eval_inner_product<0, 1>(dpf::paired, p0, In{8}, In{10}, rows);
const auto a1 = dpf::eval_inner_product<0, 1>(dpf::paired, p1, In{8}, In{10}, rows);
//! [eval-inner-product-paired]
// x=9 is hot: output0 * 1 + output1 * 9.
if (dpf::reconstruct(a0, a1) != std::uint64_t{3} * 1u + std::uint64_t{5} * 9u)
{
std::cerr << "paired leaf\n";
return 1;
}
//! [eval-inner-product-ancestor]
// Prefix slot at<4> and the full-domain leaf, one path per point.
// 0x2a and 0x2b share the high nibble 0x2, so both see payload 5 there.
// 0x10 is a different nibble. Only 0x2a is hot on the leaf.
auto [h0, h1] = dpf::make_dpf(In{0x2a}, dpf::at<4>(std::uint8_t{5}), std::uint8_t{9});
const std::vector<In> pts{0x10, 0x2a, 0x2b};
const std::vector<std::tuple<std::uint32_t, std::uint32_t>> hw{
{1u, 0u}, {1u, 1u}, {2u, 4u}};
const auto q0 = dpf::eval_sequence_inner_product<0, 1>(h0, pts.begin(), pts.end(), hw);
const auto q1 = dpf::eval_sequence_inner_product<0, 1>(h1, pts.begin(), pts.end(), hw);
//! [eval-inner-product-ancestor]
// 0x10 is off. 0x2a: 5*1 + 9*1. 0x2b: prefix still 5, leaf 0, times (2, 4).
const std::uint64_t ancestor_expect = 5u * 1u + 9u * 1u + 5u * 2u;
if (dpf::reconstruct(q0, q1) != ancestor_expect)
{
std::cerr << "ancestor\n";
return 1;
}
//! [eval-inner-product-recipe]
const auto recipe = dpf::make_sequence_recipe<decltype(h0)>(pts.begin(), pts.end());
const auto r0 = dpf::eval_sequence_inner_product<0, 1>(
h0, recipe, pts.begin(), pts.end(), hw);
const auto r1 = dpf::eval_sequence_inner_product<0, 1>(
h1, recipe, pts.begin(), pts.end(), hw);
//! [eval-inner-product-recipe]
if (dpf::reconstruct(r0, r1) != ancestor_expect)
{
std::cerr << "recipe\n";
return 1;
}
std::cout << dpf::reconstruct(s0, s1) << "\n";
return 0;
}