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>
173 lines
5.4 KiB
C++
173 lines
5.4 KiB
C++
/// @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__
|