88 lines
2.5 KiB
C++
88 lines
2.5 KiB
C++
|
|
/// @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__
|