libdpf/include/grotto/fixedpoint_mul.hpp

393 lines
14 KiB
C++
Raw Normal View History

/// @file grotto/fixedpoint_mul.hpp
/// @brief Fixed-point product with a caller-chosen integer and fraction width.
/// @details The plaintext operation an ABY2.0 / Beaver multiplier reproduces.
/// Included from grotto/fixedpoint.hpp.
#ifndef LIBDPF_INCLUDE_GROTTO_FIXEDPOINT_MUL_HPP__
#define LIBDPF_INCLUDE_GROTTO_FIXEDPOINT_MUL_HPP__
#ifndef LIBDPF_INCLUDE_DPF_FIXEDPOINT_HPP__
#include "grotto/fixedpoint.hpp"
#endif
#include <algorithm>
#include <cstddef>
#include <cstdint>
#include <type_traits>
namespace grotto
{
/// @brief Shape of `fixed_mul`: the ring width and the public shift.
///
/// The product of the raw integers is taken modulo `2^multiply_bits`, which is
/// the narrowest ring whose low bits contain the requested window. Bits of
/// that product at and above `align_shift` are the output; bits below it are
/// the discarded fraction, and bits above the window are the discarded
/// integer part. Both discards are the low or high residue modulo a power of
/// two, so a negative value is floored onto the output ulp.
///
/// A Beaver triple replaces only the multiply in `Z/2^multiply_bits Z`.
/// Reducing each operand into that ring is local when it is a truncation or a
/// zero-extend. Sign-extending a narrower signed operand, and replicating the
/// product sign when `modulus_bits > multiply_bits`, are plaintext steps the
/// MPC protocol has to reproduce (they are not local on additive shares).
template <unsigned IntegerBits,
unsigned FractionalBits,
unsigned LhsFractionalBits,
typename LhsIntegral,
unsigned RhsFractionalBits,
typename RhsIntegral>
struct fixed_mul_plan
{
static constexpr unsigned integer_bits = IntegerBits;
static constexpr unsigned fractional_bits = FractionalBits;
static constexpr unsigned out_bits = IntegerBits + FractionalBits;
static constexpr unsigned lhs_width = dpf::utils::bitlength_of_v<LhsIntegral>;
static constexpr unsigned rhs_width = dpf::utils::bitlength_of_v<RhsIntegral>;
static constexpr bool lhs_signed = std::is_signed_v<LhsIntegral>;
static constexpr bool rhs_signed = std::is_signed_v<RhsIntegral>;
static constexpr bool operands_signed = lhs_signed || rhs_signed;
/// Right shift applied to the raw product. Negative means a left shift.
static constexpr int align_shift = static_cast<int>(LhsFractionalBits)
+ static_cast<int>(RhsFractionalBits)
- static_cast<int>(FractionalBits);
/// Bits of the product that the shift reads. Zero when a left shift
/// moves every product bit out of the output.
static constexpr int modulus_bits_signed = align_shift >= 0
? align_shift + static_cast<int>(out_bits)
: static_cast<int>(out_bits) + align_shift;
static constexpr unsigned modulus_bits = modulus_bits_signed > 0
? static_cast<unsigned>(modulus_bits_signed) : 0u;
/// Full two's-complement product fits in this many bits.
static constexpr unsigned product_bits = lhs_width + rhs_width;
static constexpr unsigned multiply_bits = modulus_bits < product_bits
? modulus_bits : product_bits;
static constexpr unsigned limbs = multiply_bits == 0u
? 0u : (multiply_bits + 63u) / 64u;
/// Signed storage exists through 128 bits. A wider window is the same
/// residue held in an unsigned fixed-point.
static constexpr bool result_is_signed = operands_signed && out_bits <= 128u;
static_assert(out_bits >= 1u && out_bits <= 256u,
"fixed-point product window must be between 1 and 256 bits");
static_assert(modulus_bits <= 768u,
"fixed-point product window exceeds 768 bits");
static_assert(limbs <= 8u, "fixed-point multiply uses at most 8 limbs");
private:
template <unsigned Bits, bool Signed>
struct storage
{
static constexpr unsigned width = Bits <= 8u ? 8u
: Bits <= 16u ? 16u
: Bits <= 32u ? 32u
: Bits <= 64u ? 64u
: Bits <= 128u ? 128u : 256u;
using type = std::conditional_t<Signed,
std::conditional_t<width <= 8u, std::int8_t,
std::conditional_t<width <= 16u, std::int16_t,
std::conditional_t<width <= 32u, std::int32_t,
std::conditional_t<width <= 64u, std::int64_t, simde_int128>>>>,
dpf::utils::integral_type_from_bitlength_t<width>>;
};
public:
using integral_type = typename storage<out_bits, result_is_signed>::type;
using result_type = fixedpoint<FractionalBits, integral_type>;
};
namespace detail
{
inline constexpr std::size_t fixed_mul_buf_limbs = 12;
constexpr void mask_to_bits(std::uint64_t * limbs, std::size_t nlimbs, unsigned bits) noexcept
{
if (bits >= nlimbs * 64u)
{
return;
}
const unsigned limb = bits / 64u;
const unsigned rem = bits % 64u;
if (rem == 0u)
{
for (std::size_t i = limb; i < nlimbs; ++i)
{
limbs[i] = 0;
}
return;
}
limbs[limb] &= (std::uint64_t{1} << rem) - 1u;
for (std::size_t i = limb + 1; i < nlimbs; ++i)
{
limbs[i] = 0;
}
}
constexpr bool test_bit(const std::uint64_t * limbs, unsigned bit) noexcept
{
return ((limbs[bit / 64u] >> (bit % 64u)) & 1u) != 0u;
}
constexpr void fill_ones(std::uint64_t * limbs, unsigned from, unsigned to) noexcept
{
for (unsigned bit = from; bit < to; )
{
const unsigned limb = bit / 64u;
const unsigned rem = bit % 64u;
const unsigned count = std::min(64u - rem, to - bit);
const std::uint64_t ones = count == 64u
? ~std::uint64_t{0}
: (std::uint64_t{1} << count) - 1u;
limbs[limb] |= ones << rem;
bit += count;
}
}
constexpr void sign_extend_range(std::uint64_t * limbs, unsigned from_bits, unsigned to_bits) noexcept
{
if (to_bits <= from_bits || from_bits == 0u)
{
return;
}
if (test_bit(limbs, from_bits - 1u))
{
fill_ones(limbs, from_bits, to_bits);
}
}
template <typename T>
constexpr void store_raw_limbs(const T & value, std::uint64_t out[4]) noexcept
{
out[0] = out[1] = out[2] = out[3] = 0;
if constexpr (std::is_same_v<T, uint256_t>)
{
out[0] = value.lower().lower();
out[1] = value.lower().upper();
out[2] = value.upper().lower();
out[3] = value.upper().upper();
}
else if constexpr (std::is_same_v<T, uint128_t>)
{
out[0] = value.lower();
out[1] = value.upper();
}
else if constexpr (std::is_same_v<T, simde_uint128> || std::is_same_v<T, simde_int128>)
{
const simde_uint128 bits = static_cast<simde_uint128>(value);
out[0] = static_cast<std::uint64_t>(bits);
out[1] = static_cast<std::uint64_t>(bits >> 64);
}
else
{
using unsigned_same = std::make_unsigned_t<T>;
out[0] = static_cast<std::uint64_t>(static_cast<unsigned_same>(value));
}
}
/// Low `dest_bits` of `value`, sign-extended when `value` is a narrower signed integer.
template <typename T>
constexpr void reduce_operand(const T & value, unsigned src_bits, bool is_signed,
unsigned dest_bits, std::uint64_t * dest, unsigned nlimbs) noexcept
{
for (unsigned i = 0; i < nlimbs; ++i)
{
dest[i] = 0;
}
std::uint64_t raw[4];
store_raw_limbs(value, raw);
mask_to_bits(raw, 4, src_bits);
const unsigned copying = nlimbs < 4u ? nlimbs : 4u;
for (unsigned i = 0; i < copying; ++i)
{
dest[i] = raw[i];
}
if (dest_bits > src_bits && is_signed)
{
sign_extend_range(dest, src_bits, dest_bits);
}
mask_to_bits(dest, nlimbs, dest_bits);
}
/// Product modulo `2^(64*nlimbs)`, using exactly `nlimbs` limbs of each operand.
constexpr void mul_low_limbs(std::uint64_t * out, const std::uint64_t * lhs,
const std::uint64_t * rhs, unsigned nlimbs) noexcept
{
for (unsigned i = 0; i < nlimbs; ++i)
{
out[i] = 0;
}
for (unsigned i = 0; i < nlimbs; ++i)
{
simde_uint128 carry = 0;
for (unsigned j = 0; i + j < nlimbs; ++j)
{
const simde_uint128 prod = simde_uint128(lhs[i]) * simde_uint128(rhs[j])
+ simde_uint128(out[i + j]) + carry;
out[i + j] = static_cast<std::uint64_t>(prod);
carry = prod >> 64;
}
}
}
constexpr void shift_left_limbs(std::uint64_t * limbs, std::size_t nlimbs, unsigned shift) noexcept
{
if (shift == 0u)
{
return;
}
const unsigned limb_shift = shift / 64u;
const unsigned bit_shift = shift % 64u;
std::uint64_t tmp[fixed_mul_buf_limbs] = {};
for (std::size_t i = limb_shift; i < nlimbs; ++i)
{
const std::size_t src = i - limb_shift;
std::uint64_t hi = limbs[src] << bit_shift;
std::uint64_t lo = 0;
if (bit_shift != 0u && src > 0u)
{
lo = limbs[src - 1u] >> (64u - bit_shift);
}
tmp[i] = hi | lo;
}
for (std::size_t i = 0; i < nlimbs; ++i)
{
limbs[i] = tmp[i];
}
}
constexpr void shift_right_limbs(std::uint64_t * limbs, std::size_t nlimbs, unsigned shift) noexcept
{
if (shift == 0u)
{
return;
}
const unsigned limb_shift = shift / 64u;
const unsigned bit_shift = shift % 64u;
std::uint64_t tmp[fixed_mul_buf_limbs] = {};
for (std::size_t i = 0; i + limb_shift < nlimbs; ++i)
{
const std::size_t src = i + limb_shift;
std::uint64_t lo = limbs[src] >> bit_shift;
std::uint64_t hi = 0;
if (bit_shift != 0u && src + 1u < nlimbs)
{
hi = limbs[src + 1u] << (64u - bit_shift);
}
tmp[i] = lo | hi;
}
for (std::size_t i = 0; i < nlimbs; ++i)
{
limbs[i] = tmp[i];
}
}
template <typename T>
constexpr T limbs_to_integral(const std::uint64_t * limbs) noexcept
{
if constexpr (std::is_same_v<T, uint256_t>)
{
return uint256_t{
uint128_t{limbs[3], limbs[2]},
uint128_t{limbs[1], limbs[0]}};
}
else if constexpr (std::is_same_v<T, uint128_t>)
{
return uint128_t{limbs[1], limbs[0]};
}
else if constexpr (std::is_same_v<T, simde_uint128> || std::is_same_v<T, simde_int128>)
{
const simde_uint128 bits = simde_uint128(limbs[0])
| (simde_uint128(limbs[1]) << 64);
return static_cast<T>(bits);
}
else
{
using unsigned_same = std::make_unsigned_t<T>;
return static_cast<T>(static_cast<unsigned_same>(limbs[0]));
}
}
} // namespace detail
/// @brief Multiply two fixed-point values into a chosen integer and fraction width.
/// @tparam IntegerBits Integer bits kept in the result, including the sign bit
/// when the result is signed. Bits above this wrap.
/// @tparam FractionalBits Fraction bits kept in the result. Lower fraction bits
/// of the exact product are discarded (floored).
///
/// The result is held in the smallest fixed-point word that can store
/// `IntegerBits + FractionalBits`. A signed word is used when either operand
/// is signed and the window is at most 128 bits; otherwise the window is the
/// unsigned residue.
template <unsigned IntegerBits,
unsigned FractionalBits,
unsigned LhsFractionalBits,
typename LhsIntegral,
unsigned RhsFractionalBits,
typename RhsIntegral>
constexpr auto fixed_mul(
fixedpoint<LhsFractionalBits, LhsIntegral> lhs,
fixedpoint<RhsFractionalBits, RhsIntegral> rhs) noexcept
-> typename fixed_mul_plan<IntegerBits, FractionalBits,
LhsFractionalBits, LhsIntegral, RhsFractionalBits, RhsIntegral>::result_type
{
using plan = fixed_mul_plan<IntegerBits, FractionalBits,
LhsFractionalBits, LhsIntegral, RhsFractionalBits, RhsIntegral>;
using integral = typename plan::integral_type;
if constexpr (plan::modulus_bits == 0u || plan::limbs == 0u)
{
return make_fixed_from_integral_type<FractionalBits, integral>(static_cast<integral>(0));
}
else
{
// 1. Local reduction into the multiply ring.
std::uint64_t left[8] = {};
std::uint64_t right[8] = {};
detail::reduce_operand(lhs.integral_representation(), plan::lhs_width, plan::lhs_signed,
plan::multiply_bits, left, plan::limbs);
detail::reduce_operand(rhs.integral_representation(), plan::rhs_width, plan::rhs_signed,
plan::multiply_bits, right, plan::limbs);
// 2. The single non-linear step: product in Z/2^multiply_bits Z.
std::uint64_t prod[detail::fixed_mul_buf_limbs] = {};
detail::mul_low_limbs(prod, left, right, plan::limbs);
detail::mask_to_bits(prod, detail::fixed_mul_buf_limbs, plan::multiply_bits);
// 3. Public extension up to the window, then the public radix shift.
if constexpr (plan::operands_signed)
{
detail::sign_extend_range(prod, plan::multiply_bits, plan::modulus_bits);
}
if constexpr (plan::align_shift > 0)
{
detail::shift_right_limbs(prod, detail::fixed_mul_buf_limbs,
static_cast<unsigned>(plan::align_shift));
}
else if constexpr (plan::align_shift < 0)
{
detail::shift_left_limbs(prod, detail::fixed_mul_buf_limbs,
static_cast<unsigned>(-plan::align_shift));
}
detail::mask_to_bits(prod, detail::fixed_mul_buf_limbs, plan::out_bits);
if constexpr (plan::result_is_signed)
{
detail::sign_extend_range(prod, plan::out_bits,
dpf::utils::bitlength_of_v<integral>);
}
return make_fixed_from_integral_type<FractionalBits, integral>(
detail::limbs_to_integral<integral>(prod));
}
}
} // namespace grotto
#endif // LIBDPF_INCLUDE_GROTTO_FIXEDPOINT_MUL_HPP__