libdpf/include/dpf/matmul.hpp

174 lines
5.4 KiB
C++
Raw Permalink Normal View History

/// @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__