libdpf/include/dpf/matmul.hpp
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

173 lines
5.4 KiB
C++
Raw 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.

/// @file dpf/matmul.hpp
/// @brief Matrix Beaver triples and online matmul / im2col convolution.
#ifndef LIBDPF_INCLUDE_DPF_MATMUL_HPP__
#define LIBDPF_INCLUDE_DPF_MATMUL_HPP__
#include <cstddef>
#include <cstdint>
#include <stdexcept>
#include <utility>
#include <vector>
#include "hedley/hedley.h"
#include "dpf/random.hpp"
#include "dpf/share_vec.hpp"
namespace dpf
{
namespace matmul
{
template <typename Ring>
struct dims
{
std::size_t m = 0;
std::size_t k = 0;
std::size_t n = 0;
};
template <typename Ring>
struct matrix_triple
{
dims<Ring> d{};
std::vector<Ring> a0, a1; ///< m×k
std::vector<Ring> b0, b1; ///< k×n
std::vector<Ring> c0, c1; ///< m×n = AB
};
template <typename Ring>
HEDLEY_WARN_UNUSED_RESULT
std::vector<Ring> clear_mul(const std::vector<Ring> & a,
const std::vector<Ring> & b, dims<Ring> d)
{
if (a.size() != d.m * d.k || b.size() != d.k * d.n)
throw std::invalid_argument("matmul clear size");
std::vector<Ring> c(d.m * d.n, Ring{});
for (std::size_t i = 0; i < d.m; ++i)
for (std::size_t j = 0; j < d.n; ++j)
{
Ring acc{};
for (std::size_t t = 0; t < d.k; ++t)
acc = static_cast<Ring>(
acc + a[i * d.k + t] * b[t * d.n + j]);
c[i * d.n + j] = acc;
}
return c;
}
template <typename Ring>
HEDLEY_WARN_UNUSED_RESULT
matrix_triple<Ring> sample_triple(dims<Ring> d)
{
matrix_triple<Ring> t;
t.d = d;
t.a0.resize(d.m * d.k);
t.a1.resize(d.m * d.k);
t.b0.resize(d.k * d.n);
t.b1.resize(d.k * d.n);
std::vector<Ring> a(d.m * d.k), b(d.k * d.n);
for (std::size_t i = 0; i < a.size(); ++i)
{
a[i] = dpf::uniform_sample<Ring>();
t.a0[i] = dpf::uniform_sample<Ring>();
t.a1[i] = static_cast<Ring>(a[i] - t.a0[i]);
}
for (std::size_t i = 0; i < b.size(); ++i)
{
b[i] = dpf::uniform_sample<Ring>();
t.b0[i] = dpf::uniform_sample<Ring>();
t.b1[i] = static_cast<Ring>(b[i] - t.b0[i]);
}
auto c = clear_mul(a, b, d);
t.c0.resize(c.size());
t.c1.resize(c.size());
for (std::size_t i = 0; i < c.size(); ++i)
{
t.c0[i] = dpf::uniform_sample<Ring>();
t.c1[i] = static_cast<Ring>(c[i] - t.c0[i]);
}
return t;
}
/// @brief Online: open `D = X-A`, `E = Y-B`; `Z = C + D B + A E + D E` (p0 adds DE).
template <typename Ring>
HEDLEY_WARN_UNUSED_RESULT
std::vector<Ring> online(const std::vector<Ring> & x_share,
const std::vector<Ring> & y_share, const std::vector<Ring> & a_share,
const std::vector<Ring> & b_share, const std::vector<Ring> & c_share,
const std::vector<Ring> & d_open, const std::vector<Ring> & e_open,
dims<Ring> d, unsigned party)
{
if (x_share.size() != d.m * d.k || y_share.size() != d.k * d.n)
throw std::invalid_argument("matmul online size");
auto db = clear_mul(d_open, b_share, d);
// A·E : a is m×k, e is k×n
auto ae = clear_mul(a_share, e_open, d);
std::vector<Ring> z(d.m * d.n);
for (std::size_t i = 0; i < z.size(); ++i)
z[i] = static_cast<Ring>(c_share[i] + db[i] + ae[i]);
if (party == 0)
{
auto de = clear_mul(d_open, e_open, d);
for (std::size_t i = 0; i < z.size(); ++i)
z[i] = static_cast<Ring>(z[i] + de[i]);
}
(void)x_share;
(void)y_share;
return z;
}
/// @brief Full clear+share test: return shares of X Y.
template <typename Ring>
HEDLEY_WARN_UNUSED_RESULT
std::pair<std::vector<Ring>, std::vector<Ring>> mul_shared(
const std::vector<Ring> & x, const std::vector<Ring> & y, dims<Ring> d)
{
auto t = sample_triple<Ring>(d);
std::vector<Ring> x0(x.size()), x1(x.size()), y0(y.size()), y1(y.size());
for (std::size_t i = 0; i < x.size(); ++i)
{
x0[i] = dpf::uniform_sample<Ring>();
x1[i] = static_cast<Ring>(x[i] - x0[i]);
}
for (std::size_t i = 0; i < y.size(); ++i)
{
y0[i] = dpf::uniform_sample<Ring>();
y1[i] = static_cast<Ring>(y[i] - y0[i]);
}
std::vector<Ring> d_open(d.m * d.k), e_open(d.k * d.n);
for (std::size_t i = 0; i < d_open.size(); ++i)
d_open[i] = static_cast<Ring>((x0[i] + x1[i]) - (t.a0[i] + t.a1[i]));
for (std::size_t i = 0; i < e_open.size(); ++i)
e_open[i] = static_cast<Ring>((y0[i] + y1[i]) - (t.b0[i] + t.b1[i]));
auto z0 = online(x0, y0, t.a0, t.b0, t.c0, d_open, e_open, d, 0);
auto z1 = online(x1, y1, t.a1, t.b1, t.c1, d_open, e_open, d, 1);
return {std::move(z0), std::move(z1)};
}
/// @brief im2col: flatten `h×w` patches of size `kh×kw` with stride 1 into rows.
template <typename Ring>
HEDLEY_WARN_UNUSED_RESULT
std::vector<Ring> im2col(const std::vector<Ring> & img, std::size_t h,
std::size_t w, std::size_t kh, std::size_t kw)
{
if (img.size() != h * w || kh > h || kw > w)
throw std::invalid_argument("im2col");
const std::size_t out_h = h - kh + 1;
const std::size_t out_w = w - kw + 1;
std::vector<Ring> col(out_h * out_w * kh * kw);
std::size_t row = 0;
for (std::size_t i = 0; i < out_h; ++i)
for (std::size_t j = 0; j < out_w; ++j, ++row)
for (std::size_t u = 0; u < kh; ++u)
for (std::size_t v = 0; v < kw; ++v)
col[row * (kh * kw) + u * kw + v] =
img[(i + u) * w + (j + v)];
return col;
}
} // namespace matmul
} // namespace dpf
#endif // LIBDPF_INCLUDE_DPF_MATMUL_HPP__