/// @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 #include #include #include #include #include "hedley/hedley.h" #include "dpf/random.hpp" #include "dpf/share_vec.hpp" namespace dpf { namespace matmul { template struct dims { std::size_t m = 0; std::size_t k = 0; std::size_t n = 0; }; template struct matrix_triple { dims d{}; std::vector a0, a1; ///< m×k std::vector b0, b1; ///< k×n std::vector c0, c1; ///< m×n = AB }; template HEDLEY_WARN_UNUSED_RESULT std::vector clear_mul(const std::vector & a, const std::vector & b, dims d) { if (a.size() != d.m * d.k || b.size() != d.k * d.n) throw std::invalid_argument("matmul clear size"); std::vector 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( acc + a[i * d.k + t] * b[t * d.n + j]); c[i * d.n + j] = acc; } return c; } template HEDLEY_WARN_UNUSED_RESULT matrix_triple sample_triple(dims d) { matrix_triple 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 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(); t.a0[i] = dpf::uniform_sample(); t.a1[i] = static_cast(a[i] - t.a0[i]); } for (std::size_t i = 0; i < b.size(); ++i) { b[i] = dpf::uniform_sample(); t.b0[i] = dpf::uniform_sample(); t.b1[i] = static_cast(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(); t.c1[i] = static_cast(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 HEDLEY_WARN_UNUSED_RESULT std::vector online(const std::vector & x_share, const std::vector & y_share, const std::vector & a_share, const std::vector & b_share, const std::vector & c_share, const std::vector & d_open, const std::vector & e_open, dims 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 z(d.m * d.n); for (std::size_t i = 0; i < z.size(); ++i) z[i] = static_cast(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(z[i] + de[i]); } (void)x_share; (void)y_share; return z; } /// @brief Full clear+share test: return shares of X Y. template HEDLEY_WARN_UNUSED_RESULT std::pair, std::vector> mul_shared( const std::vector & x, const std::vector & y, dims d) { auto t = sample_triple(d); std::vector 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(); x1[i] = static_cast(x[i] - x0[i]); } for (std::size_t i = 0; i < y.size(); ++i) { y0[i] = dpf::uniform_sample(); y1[i] = static_cast(y[i] - y0[i]); } std::vector 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((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((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 HEDLEY_WARN_UNUSED_RESULT std::vector im2col(const std::vector & 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 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__