libdpf/include/dpf/fixed_share.hpp

88 lines
2.5 KiB
C++
Raw Permalink Normal View History

/// @file dpf/fixed_share.hpp
/// @brief Fixed-point arithmetic shares: `fixed<IntBits, FracBits>`.
#ifndef LIBDPF_INCLUDE_DPF_FIXED_SHARE_HPP__
#define LIBDPF_INCLUDE_DPF_FIXED_SHARE_HPP__
#include <cstddef>
#include <cstdint>
#include <stdexcept>
#include <type_traits>
#include <utility>
#include <vector>
#include "hedley/hedley.h"
#include "dpf/share_vec.hpp"
#include "dpf/trunc.hpp"
namespace dpf
{
template <unsigned IntBits, unsigned FracBits>
struct fixed
{
static constexpr unsigned int_bits = IntBits;
static constexpr unsigned frac_bits = FracBits;
static constexpr unsigned width = IntBits + FracBits;
static_assert(width > 0 && width <= 128, "fixed width 1..128");
using ring = std::conditional_t<(width > 64), unsigned __int128, std::uint64_t>;
ring raw{};
fixed() = default;
explicit fixed(ring v) : raw(v) {}
static fixed from_integer(ring i)
{
return fixed{static_cast<ring>(i << FracBits)};
}
static fixed mul_clear(fixed a, fixed b)
{
const ring prod = static_cast<ring>(a.raw * b.raw);
return fixed{static_cast<ring>(prod >> FracBits)};
}
friend fixed operator+(fixed a, fixed b)
{
return fixed{static_cast<ring>(a.raw + b.raw)};
}
friend fixed operator-(fixed a, fixed b)
{
return fixed{static_cast<ring>(a.raw - b.raw)};
}
};
template <unsigned IntBits, unsigned FracBits>
using fixed_vec = share_vec<typename fixed<IntBits, FracBits>::ring>;
/// @brief Fixed-point product via mul_trunc by FracBits (Beaver, both parties).
template <unsigned IntBits, unsigned FracBits>
HEDLEY_WARN_UNUSED_RESULT
std::pair<fixed_vec<IntBits, FracBits>, fixed_vec<IntBits, FracBits>>
fixed_mul_share(const fixed_vec<IntBits, FracBits> & x0,
const fixed_vec<IntBits, FracBits> & x1,
const fixed_vec<IntBits, FracBits> & y0,
const fixed_vec<IntBits, FracBits> & y1)
{
if (x0.size() != x1.size() || x0.size() != y0.size()
|| y0.size() != y1.size())
throw std::invalid_argument("fixed_mul_share size");
fixed_vec<IntBits, FracBits> z0(x0.size(), protocol::domain::a, 0);
fixed_vec<IntBits, FracBits> z1(x0.size(), protocol::domain::a, 1);
for (std::size_t i = 0; i < x0.size(); ++i)
{
// Widen so IntBits+FracBits products do not wrap before the shift.
auto mt = trunc::mul_exact_trunc(x0[i], x1[i], y0[i], y1[i], FracBits);
z0[i] = mt.z0;
z1[i] = mt.z1;
}
return {std::move(z0), std::move(z1)};
}
} // namespace dpf
#endif // LIBDPF_INCLUDE_DPF_FIXED_SHARE_HPP__