Ship the TLS mesh, composer, Beaver/Yao/leaf MPC, prep/online paths, apps, and docs so the tree is pushable before elevating share_expr, security_mode, and prep resume. Co-authored-by: Cursor <cursoragent@cursor.com>
781 lines
28 KiB
C++
781 lines
28 KiB
C++
/// @file dpf/bitmore_mod.hpp
|
|
/// @brief Byte-slot reduction for BitMore moduli through 255.
|
|
/// @details Hafiz and Henry (PoPETs 2019 §5.3) fold each server's DPF bits
|
|
/// into an integer and reduce modulo the server count. When that
|
|
/// count is not a power of two the integer does not fit in a byte
|
|
/// for long, so each byte is only partially reduced until the end.
|
|
///
|
|
/// A partial step reads the high nibble `h`. Every byte with that
|
|
/// nibble is at least `16*h`, so `floor(16*h / M) * M` is the largest
|
|
/// multiple of `M` that is safe to subtract. `pshufb` selects it and
|
|
/// `sub_epi8` removes it. The byte stays congruent modulo `M` and is
|
|
/// at most `partial_bound`. For moduli 121..127 that ceiling is 128
|
|
/// or more, so a second partial step is what leaves `stable_bound`
|
|
/// (at most 127). `resume_bound` is the ceiling the accumulator
|
|
/// actually resumes from: one step through modulus 120 and 128, two
|
|
/// steps for 121..127.
|
|
///
|
|
/// `add_budget = 255 - resume_bound` is how much can still be added
|
|
/// before a byte might reach 256. `shift_budget` is how many
|
|
/// `2*acc + bit` insertions fit in that slack. Both stay positive
|
|
/// through modulus 128, which is as far as a byte can hold two
|
|
/// resumed slots or one more bit. `partial_reduce` and `full_reduce`
|
|
/// themselves stay correct through 255: the two nibble residues sum
|
|
/// to at most 255 and to less than `2*M`, and one unsigned compare
|
|
/// subtracts `M`. Above 128 there is no slack left to defer that
|
|
/// correction across another add.
|
|
///
|
|
/// `bitmore_mod<M, 16>` is the same idea on 16-bit lanes, still on
|
|
/// AVX2. The top nibble (bits 12..15) selects `floor(4096*h / M)*M`
|
|
/// through two `pshufb`s, one per byte of that multiple, and
|
|
/// `sub_epi16` removes it. The accumulator has slack through modulus
|
|
/// 32768 (`resume_bound` at most 32767, so one more bit fits in the
|
|
/// lane). A full reduction sums the four nibble residues, which fit
|
|
/// in the lane for every modulus through 65535, then subtracts `M`
|
|
/// up to three times.
|
|
/// @copyright Copyright (c) 2019-2026 Ryan Henry and [others](@ref authors)
|
|
/// @license Released under a GNU General Public v2.0 (GPLv2) license.
|
|
|
|
#ifndef LIBDPF_INCLUDE_DPF_BITMORE_MOD_HPP__
|
|
#define LIBDPF_INCLUDE_DPF_BITMORE_MOD_HPP__
|
|
|
|
#include <array>
|
|
#include <cstddef>
|
|
#include <cstdint>
|
|
#include <cstring>
|
|
#include <type_traits>
|
|
|
|
#include "hedley/hedley.h"
|
|
#include "simde/simde/x86/avx2.h"
|
|
|
|
namespace dpf
|
|
{
|
|
namespace bitmore_detail
|
|
{
|
|
|
|
template <typename Reg>
|
|
HEDLEY_ALWAYS_INLINE
|
|
HEDLEY_NO_THROW
|
|
Reg zero() noexcept
|
|
{
|
|
if constexpr (std::is_same_v<Reg, simde__m256i>)
|
|
return simde_mm256_setzero_si256();
|
|
else
|
|
return simde_mm_setzero_si128();
|
|
}
|
|
|
|
template <typename Reg>
|
|
HEDLEY_ALWAYS_INLINE
|
|
HEDLEY_NO_THROW
|
|
Reg load_lut(const std::array<unsigned char, 16> & lut) noexcept
|
|
{
|
|
if constexpr (std::is_same_v<Reg, simde__m128i>)
|
|
{
|
|
simde__m128i table;
|
|
std::memcpy(&table, lut.data(), 16);
|
|
return table;
|
|
}
|
|
else
|
|
{
|
|
alignas(32) unsigned char both[32];
|
|
std::memcpy(both, lut.data(), 16);
|
|
std::memcpy(both + 16, lut.data(), 16);
|
|
simde__m256i table;
|
|
std::memcpy(&table, both, 32);
|
|
return table;
|
|
}
|
|
}
|
|
|
|
template <typename Reg>
|
|
HEDLEY_ALWAYS_INLINE
|
|
HEDLEY_NO_THROW
|
|
Reg high_nibble(Reg x) noexcept
|
|
{
|
|
// `srli_epi16` moves the neighbouring byte's low nibble into bits 4..7.
|
|
// Masking with `0x0f` leaves this byte's own high nibble.
|
|
if constexpr (std::is_same_v<Reg, simde__m128i>)
|
|
{
|
|
const auto m = simde_mm_set1_epi8(0x0f);
|
|
return simde_mm_and_si128(simde_mm_srli_epi16(x, 4), m);
|
|
}
|
|
else
|
|
{
|
|
const auto m = simde_mm256_set1_epi8(0x0f);
|
|
return simde_mm256_and_si256(simde_mm256_srli_epi16(x, 4), m);
|
|
}
|
|
}
|
|
|
|
template <typename Reg>
|
|
HEDLEY_ALWAYS_INLINE
|
|
HEDLEY_NO_THROW
|
|
Reg low_nibble(Reg x) noexcept
|
|
{
|
|
if constexpr (std::is_same_v<Reg, simde__m128i>)
|
|
return simde_mm_and_si128(x, simde_mm_set1_epi8(0x0f));
|
|
else
|
|
return simde_mm256_and_si256(x, simde_mm256_set1_epi8(0x0f));
|
|
}
|
|
|
|
template <typename Reg>
|
|
HEDLEY_ALWAYS_INLINE
|
|
HEDLEY_NO_THROW
|
|
Reg shuffle(Reg table, Reg index) noexcept
|
|
{
|
|
if constexpr (std::is_same_v<Reg, simde__m128i>)
|
|
return simde_mm_shuffle_epi8(table, index);
|
|
else
|
|
return simde_mm256_shuffle_epi8(table, index);
|
|
}
|
|
|
|
template <typename Reg>
|
|
HEDLEY_ALWAYS_INLINE
|
|
HEDLEY_NO_THROW
|
|
Reg add_bytes(Reg a, Reg b) noexcept
|
|
{
|
|
if constexpr (std::is_same_v<Reg, simde__m128i>)
|
|
return simde_mm_add_epi8(a, b);
|
|
else
|
|
return simde_mm256_add_epi8(a, b);
|
|
}
|
|
|
|
template <typename Reg>
|
|
HEDLEY_ALWAYS_INLINE
|
|
HEDLEY_NO_THROW
|
|
Reg sub_bytes(Reg a, Reg b) noexcept
|
|
{
|
|
if constexpr (std::is_same_v<Reg, simde__m128i>)
|
|
return simde_mm_sub_epi8(a, b);
|
|
else
|
|
return simde_mm256_sub_epi8(a, b);
|
|
}
|
|
|
|
template <typename Reg>
|
|
HEDLEY_ALWAYS_INLINE
|
|
HEDLEY_NO_THROW
|
|
Reg and_bytes(Reg a, Reg b) noexcept
|
|
{
|
|
if constexpr (std::is_same_v<Reg, simde__m128i>)
|
|
return simde_mm_and_si128(a, b);
|
|
else
|
|
return simde_mm256_and_si256(a, b);
|
|
}
|
|
|
|
template <typename Reg>
|
|
HEDLEY_ALWAYS_INLINE
|
|
HEDLEY_NO_THROW
|
|
Reg xor_bytes(Reg a, Reg b) noexcept
|
|
{
|
|
if constexpr (std::is_same_v<Reg, simde__m128i>)
|
|
return simde_mm_xor_si128(a, b);
|
|
else
|
|
return simde_mm256_xor_si256(a, b);
|
|
}
|
|
|
|
template <typename Reg>
|
|
HEDLEY_ALWAYS_INLINE
|
|
HEDLEY_NO_THROW
|
|
Reg splat_epi8(unsigned char v) noexcept
|
|
{
|
|
const auto s = static_cast<int8_t>(v);
|
|
if constexpr (std::is_same_v<Reg, simde__m128i>)
|
|
return simde_mm_set1_epi8(s);
|
|
else
|
|
return simde_mm256_set1_epi8(s);
|
|
}
|
|
|
|
/// @brief Unsigned `a > b` per byte. `cmpgt_epi8` is signed; XOR `0x80` maps
|
|
/// unsigned order onto that signed order.
|
|
template <typename Reg>
|
|
HEDLEY_ALWAYS_INLINE
|
|
HEDLEY_NO_THROW
|
|
Reg cmpgt_epu8(Reg a, Reg b) noexcept
|
|
{
|
|
const auto bias = splat_epi8<Reg>(0x80);
|
|
const auto aa = xor_bytes<Reg>(a, bias);
|
|
const auto bb = xor_bytes<Reg>(b, bias);
|
|
if constexpr (std::is_same_v<Reg, simde__m128i>)
|
|
return simde_mm_cmpgt_epi8(aa, bb);
|
|
else
|
|
return simde_mm256_cmpgt_epi8(aa, bb);
|
|
}
|
|
|
|
/// @brief Shift each byte left by 1 and set bit 0 from `bit`.
|
|
/// @details `slli_epi16` spills bit 7 into the next byte of the 16-bit lane.
|
|
/// Clearing bit 0 afterwards drops that spill. Bit 7 of the odd byte
|
|
/// shifts out of the lane, which is the byte-local shift.
|
|
template <typename Reg>
|
|
HEDLEY_ALWAYS_INLINE
|
|
HEDLEY_NO_THROW
|
|
Reg shift_in_bit(Reg acc, Reg bit) noexcept
|
|
{
|
|
const auto one = splat_epi8<Reg>(1);
|
|
const auto keep = splat_epi8<Reg>(0xfe);
|
|
bit = and_bytes<Reg>(bit, one);
|
|
Reg shifted;
|
|
if constexpr (std::is_same_v<Reg, simde__m128i>)
|
|
shifted = simde_mm_and_si128(simde_mm_slli_epi16(acc, 1), keep);
|
|
else
|
|
shifted = simde_mm256_and_si256(simde_mm256_slli_epi16(acc, 1), keep);
|
|
if constexpr (std::is_same_v<Reg, simde__m128i>)
|
|
return simde_mm_or_si128(shifted, bit);
|
|
else
|
|
return simde_mm256_or_si256(shifted, bit);
|
|
}
|
|
|
|
template <unsigned Modulus, unsigned LaneBits = 8>
|
|
HEDLEY_CONST
|
|
constexpr unsigned partial_bound() noexcept
|
|
{
|
|
static_assert(LaneBits == 8u || LaneBits == 16u, "bitmore lane width is 8 or 16");
|
|
constexpr unsigned block = LaneBits == 8u ? 16u : 4096u;
|
|
constexpr unsigned tail = block - 1u;
|
|
unsigned u = 0;
|
|
for (unsigned h = 0; h < 16u; ++h)
|
|
{
|
|
const unsigned q = (block * h / Modulus) * Modulus;
|
|
const unsigned top = block * h + tail - q;
|
|
if (top > u)
|
|
u = top;
|
|
}
|
|
return u;
|
|
}
|
|
|
|
template <unsigned Modulus, unsigned LaneBits = 8>
|
|
HEDLEY_CONST
|
|
constexpr unsigned stable_bound() noexcept
|
|
{
|
|
constexpr unsigned block = LaneBits == 8u ? 16u : 4096u;
|
|
const unsigned lim = partial_bound<Modulus, LaneBits>();
|
|
unsigned u = 0;
|
|
for (unsigned h = 0; h < 16u; ++h)
|
|
{
|
|
const unsigned lo = block * h;
|
|
if (lo > lim)
|
|
break;
|
|
unsigned hi = lo + block - 1u;
|
|
if (hi > lim)
|
|
hi = lim;
|
|
const unsigned q = (lo / Modulus) * Modulus;
|
|
const unsigned top = hi - q;
|
|
if (top > u)
|
|
u = top;
|
|
}
|
|
return u;
|
|
}
|
|
|
|
template <unsigned Modulus>
|
|
HEDLEY_CONST
|
|
constexpr unsigned shift_budget(unsigned bound, unsigned capacity = 256u) noexcept
|
|
{
|
|
unsigned k = 0;
|
|
unsigned span = bound + 1u;
|
|
const unsigned half = capacity >> 1;
|
|
while (span <= half)
|
|
{
|
|
span *= 2u;
|
|
++k;
|
|
}
|
|
return k;
|
|
}
|
|
|
|
template <unsigned Modulus>
|
|
HEDLEY_CONST
|
|
constexpr std::array<unsigned char, 16> partial_lut() noexcept
|
|
{
|
|
std::array<unsigned char, 16> lut{};
|
|
for (unsigned h = 0; h < 16u; ++h)
|
|
lut[h] = static_cast<unsigned char>((16u * h / Modulus) * Modulus);
|
|
return lut;
|
|
}
|
|
|
|
template <unsigned Modulus>
|
|
HEDLEY_CONST
|
|
constexpr std::array<unsigned char, 16> low_residue_lut() noexcept
|
|
{
|
|
std::array<unsigned char, 16> lut{};
|
|
for (unsigned n = 0; n < 16u; ++n)
|
|
lut[n] = static_cast<unsigned char>(n % Modulus);
|
|
return lut;
|
|
}
|
|
|
|
template <unsigned Modulus>
|
|
HEDLEY_CONST
|
|
constexpr std::array<unsigned char, 16> high_residue_lut() noexcept
|
|
{
|
|
std::array<unsigned char, 16> lut{};
|
|
for (unsigned h = 0; h < 16u; ++h)
|
|
lut[h] = static_cast<unsigned char>((16u * h) % Modulus);
|
|
return lut;
|
|
}
|
|
|
|
template <typename Reg>
|
|
HEDLEY_ALWAYS_INLINE
|
|
HEDLEY_NO_THROW
|
|
Reg splat_epi16(unsigned v) noexcept
|
|
{
|
|
const auto s = static_cast<std::int16_t>(static_cast<std::uint16_t>(v));
|
|
if constexpr (std::is_same_v<Reg, simde__m128i>)
|
|
return simde_mm_set1_epi16(s);
|
|
else
|
|
return simde_mm256_set1_epi16(s);
|
|
}
|
|
|
|
template <typename Reg>
|
|
HEDLEY_ALWAYS_INLINE
|
|
HEDLEY_NO_THROW
|
|
Reg add_epi16(Reg a, Reg b) noexcept
|
|
{
|
|
if constexpr (std::is_same_v<Reg, simde__m128i>)
|
|
return simde_mm_add_epi16(a, b);
|
|
else
|
|
return simde_mm256_add_epi16(a, b);
|
|
}
|
|
|
|
template <typename Reg>
|
|
HEDLEY_ALWAYS_INLINE
|
|
HEDLEY_NO_THROW
|
|
Reg sub_epi16(Reg a, Reg b) noexcept
|
|
{
|
|
if constexpr (std::is_same_v<Reg, simde__m128i>)
|
|
return simde_mm_sub_epi16(a, b);
|
|
else
|
|
return simde_mm256_sub_epi16(a, b);
|
|
}
|
|
|
|
template <typename Reg>
|
|
HEDLEY_ALWAYS_INLINE
|
|
HEDLEY_NO_THROW
|
|
Reg srli_epi16(Reg a, int imm) noexcept
|
|
{
|
|
if constexpr (std::is_same_v<Reg, simde__m128i>)
|
|
return simde_mm_srli_epi16(a, imm);
|
|
else
|
|
return simde_mm256_srli_epi16(a, imm);
|
|
}
|
|
|
|
template <typename Reg>
|
|
HEDLEY_ALWAYS_INLINE
|
|
HEDLEY_NO_THROW
|
|
Reg slli_epi16(Reg a, int imm) noexcept
|
|
{
|
|
if constexpr (std::is_same_v<Reg, simde__m128i>)
|
|
return simde_mm_slli_epi16(a, imm);
|
|
else
|
|
return simde_mm256_slli_epi16(a, imm);
|
|
}
|
|
|
|
/// @brief Unsigned `a > b` per 16-bit lane.
|
|
template <typename Reg>
|
|
HEDLEY_ALWAYS_INLINE
|
|
HEDLEY_NO_THROW
|
|
Reg cmpgt_epu16(Reg a, Reg b) noexcept
|
|
{
|
|
const auto bias = splat_epi16<Reg>(0x8000u);
|
|
const auto aa = xor_bytes<Reg>(a, bias);
|
|
const auto bb = xor_bytes<Reg>(b, bias);
|
|
if constexpr (std::is_same_v<Reg, simde__m128i>)
|
|
return simde_mm_cmpgt_epi16(aa, bb);
|
|
else
|
|
return simde_mm256_cmpgt_epi16(aa, bb);
|
|
}
|
|
|
|
template <typename Reg>
|
|
HEDLEY_ALWAYS_INLINE
|
|
HEDLEY_NO_THROW
|
|
Reg or_bytes(Reg a, Reg b) noexcept
|
|
{
|
|
if constexpr (std::is_same_v<Reg, simde__m128i>)
|
|
return simde_mm_or_si128(a, b);
|
|
else
|
|
return simde_mm256_or_si256(a, b);
|
|
}
|
|
|
|
/// @brief Look up a 16-bit table entry. `idx` holds the nibble in the low byte
|
|
/// of each lane and zero in the high byte, so `pshufb` writes `lut[h]`
|
|
/// into the low byte and `lut[0]` into the high byte. `lut[0]` is 0.
|
|
template <typename Reg>
|
|
HEDLEY_ALWAYS_INLINE
|
|
HEDLEY_NO_THROW
|
|
Reg lookup_u16(Reg idx, const std::array<unsigned char, 16> & lo,
|
|
const std::array<unsigned char, 16> & hi) noexcept
|
|
{
|
|
const auto lo_v = shuffle<Reg>(load_lut<Reg>(lo), idx);
|
|
const auto hi_v = slli_epi16<Reg>(shuffle<Reg>(load_lut<Reg>(hi), idx), 8);
|
|
return or_bytes<Reg>(lo_v, hi_v);
|
|
}
|
|
|
|
constexpr std::array<unsigned char, 16> u16_lo(const std::array<std::uint16_t, 16> & v) noexcept
|
|
{
|
|
std::array<unsigned char, 16> out{};
|
|
for (unsigned i = 0; i < 16u; ++i)
|
|
out[i] = static_cast<unsigned char>(v[i] & 0xffu);
|
|
return out;
|
|
}
|
|
|
|
constexpr std::array<unsigned char, 16> u16_hi(const std::array<std::uint16_t, 16> & v) noexcept
|
|
{
|
|
std::array<unsigned char, 16> out{};
|
|
for (unsigned i = 0; i < 16u; ++i)
|
|
out[i] = static_cast<unsigned char>(v[i] >> 8);
|
|
return out;
|
|
}
|
|
|
|
template <unsigned Modulus>
|
|
HEDLEY_CONST
|
|
constexpr std::array<std::uint16_t, 16> wide_partial_lut() noexcept
|
|
{
|
|
std::array<std::uint16_t, 16> lut{};
|
|
for (unsigned h = 0; h < 16u; ++h)
|
|
lut[h] = static_cast<std::uint16_t>((4096u * h / Modulus) * Modulus);
|
|
return lut;
|
|
}
|
|
|
|
template <unsigned Modulus, unsigned Place>
|
|
HEDLEY_CONST
|
|
constexpr std::array<std::uint16_t, 16> wide_residue_lut() noexcept
|
|
{
|
|
std::array<std::uint16_t, 16> lut{};
|
|
for (unsigned n = 0; n < 16u; ++n)
|
|
lut[n] = static_cast<std::uint16_t>(
|
|
(static_cast<std::uint32_t>(Place) * n) % Modulus);
|
|
return lut;
|
|
}
|
|
|
|
/// @brief `acc = 2*acc + bit0` inside each 16-bit lane. `slli_epi16` does not
|
|
/// cross lanes.
|
|
template <typename Reg>
|
|
HEDLEY_ALWAYS_INLINE
|
|
HEDLEY_NO_THROW
|
|
Reg shift_in_bit16(Reg acc, Reg bit) noexcept
|
|
{
|
|
const auto one = splat_epi16<Reg>(1u);
|
|
bit = and_bytes<Reg>(bit, one);
|
|
if constexpr (std::is_same_v<Reg, simde__m128i>)
|
|
return simde_mm_or_si128(simde_mm_slli_epi16(acc, 1), bit);
|
|
else
|
|
return simde_mm256_or_si256(simde_mm256_slli_epi16(acc, 1), bit);
|
|
}
|
|
|
|
} // namespace bitmore_detail
|
|
|
|
/// @brief Partial and full reduction of `LaneBits`-wide slots modulo `Modulus`.
|
|
/// @tparam Modulus server count. `2` through `255` for bytes, `2` through `65535` for 16-bit lanes.
|
|
/// @tparam LaneBits `8` (one byte per slot) or `16` (one AVX2 `epi16` lane per slot)
|
|
template <unsigned Modulus, unsigned LaneBits = 8>
|
|
struct bitmore_mod
|
|
{
|
|
static_assert(LaneBits == 8u || LaneBits == 16u,
|
|
"bitmore lane width is 8 or 16");
|
|
static_assert(Modulus >= 2u, "bitmore modulus must be at least 2");
|
|
static_assert(LaneBits == 16u || Modulus <= 255u,
|
|
"bitmore byte reduction: modulus must be in 2..255");
|
|
static_assert(LaneBits == 8u || Modulus <= 65535u,
|
|
"bitmore 16-bit reduction: modulus must be in 2..65535");
|
|
|
|
static constexpr unsigned modulus = Modulus;
|
|
static constexpr unsigned lane_bits = LaneBits;
|
|
static constexpr unsigned slot_max = LaneBits == 8u ? 255u : 65535u;
|
|
static constexpr unsigned nibble_block = LaneBits == 8u ? 16u : 4096u;
|
|
static constexpr unsigned nibble_shift = LaneBits == 8u ? 4u : 12u;
|
|
|
|
/// @brief Largest slot one top-nibble partial reduction can leave.
|
|
static constexpr unsigned partial_bound = bitmore_detail::partial_bound<Modulus, LaneBits>();
|
|
|
|
/// @brief Largest slot a second partial step can leave, starting from `partial_bound`.
|
|
static constexpr unsigned stable_bound = bitmore_detail::stable_bound<Modulus, LaneBits>();
|
|
|
|
/// @brief Ceiling the accumulator resumes from. One step when that already
|
|
/// fits in the lower half of the lane; otherwise the second step.
|
|
static constexpr unsigned resume_bound =
|
|
partial_bound <= (slot_max >> 1) ? partial_bound : stable_bound;
|
|
|
|
/// @brief How much can be added to a resumed slot before it might reach `slot_max + 1`.
|
|
static constexpr unsigned add_budget = slot_max - resume_bound;
|
|
|
|
/// @brief Bit insertions that fit in `add_budget` after resuming.
|
|
static constexpr unsigned shift_budget =
|
|
bitmore_detail::shift_budget<Modulus>(resume_bound, slot_max + 1u);
|
|
|
|
/// @brief `x - floor(block*(x>>shift) / M) * M` inside one slot.
|
|
/// \complexity One division of a nibble index.
|
|
HEDLEY_CONST
|
|
HEDLEY_ALWAYS_INLINE
|
|
static constexpr unsigned partial_reduce_slot(unsigned x) noexcept
|
|
{
|
|
x &= slot_max;
|
|
const unsigned h = x >> nibble_shift;
|
|
const unsigned q = (nibble_block * h / Modulus) * Modulus;
|
|
return x - q;
|
|
}
|
|
|
|
/// @brief Byte-slot name for `partial_reduce_slot`.
|
|
HEDLEY_CONST
|
|
HEDLEY_ALWAYS_INLINE
|
|
static constexpr unsigned partial_reduce_byte(unsigned x) noexcept
|
|
{
|
|
static_assert(LaneBits == 8u, "partial_reduce_byte is the 8-bit slot");
|
|
return partial_reduce_slot(x);
|
|
}
|
|
|
|
/// @brief `x mod Modulus` for one slot.
|
|
/// \complexity One remainder of a slot.
|
|
HEDLEY_CONST
|
|
HEDLEY_ALWAYS_INLINE
|
|
static constexpr unsigned full_reduce_slot(unsigned x) noexcept
|
|
{
|
|
return (x & slot_max) % Modulus;
|
|
}
|
|
|
|
/// @brief Byte-slot name for `full_reduce_slot`.
|
|
HEDLEY_CONST
|
|
HEDLEY_ALWAYS_INLINE
|
|
static constexpr unsigned full_reduce_byte(unsigned x) noexcept
|
|
{
|
|
static_assert(LaneBits == 8u, "full_reduce_byte is the 8-bit slot");
|
|
return full_reduce_slot(x);
|
|
}
|
|
|
|
/// @brief Partial-reduce every byte of `x`.
|
|
/// @details The high nibble selects `floor(16*h / M) * M`. Subtracting it
|
|
/// preserves the residue and leaves a byte of at most `partial_bound`.
|
|
/// \complexity One `pshufb` and one `sub_epi8` per register.
|
|
template <typename Reg>
|
|
HEDLEY_ALWAYS_INLINE
|
|
HEDLEY_NO_THROW
|
|
HEDLEY_PURE
|
|
static Reg partial_reduce(Reg x) noexcept
|
|
{
|
|
static_assert(std::is_same_v<Reg, simde__m128i> || std::is_same_v<Reg, simde__m256i>,
|
|
"bitmore partial_reduce: register must be simde__m128i or simde__m256i");
|
|
if constexpr (LaneBits == 8u)
|
|
{
|
|
constexpr auto lut = bitmore_detail::partial_lut<Modulus>();
|
|
const auto q = bitmore_detail::shuffle<Reg>(
|
|
bitmore_detail::load_lut<Reg>(lut),
|
|
bitmore_detail::high_nibble<Reg>(x));
|
|
return bitmore_detail::sub_bytes<Reg>(x, q);
|
|
}
|
|
else
|
|
{
|
|
constexpr auto lut = bitmore_detail::wide_partial_lut<Modulus>();
|
|
constexpr auto lo = bitmore_detail::u16_lo(lut);
|
|
constexpr auto hi = bitmore_detail::u16_hi(lut);
|
|
const auto idx = bitmore_detail::srli_epi16<Reg>(x, 12);
|
|
const auto q = bitmore_detail::lookup_u16<Reg>(idx, lo, hi);
|
|
return bitmore_detail::sub_epi16<Reg>(x, q);
|
|
}
|
|
}
|
|
|
|
/// @brief Fully reduce every byte of `x` into `0 .. Modulus-1`.
|
|
/// @details `pshufb` maps the low nibble to itself modulo `M` and the high
|
|
/// nibble to `(16*h) mod M`. The sum is less than `2*M`, so one
|
|
/// compare subtracts `M` where the sum is still too big.
|
|
/// \complexity Two `pshufb`s, one byte add, one compare, one byte subtract.
|
|
template <typename Reg>
|
|
HEDLEY_ALWAYS_INLINE
|
|
HEDLEY_NO_THROW
|
|
HEDLEY_PURE
|
|
static Reg full_reduce(Reg x) noexcept
|
|
{
|
|
static_assert(std::is_same_v<Reg, simde__m128i> || std::is_same_v<Reg, simde__m256i>,
|
|
"bitmore full_reduce: register must be simde__m128i or simde__m256i");
|
|
if constexpr (LaneBits == 8u)
|
|
{
|
|
constexpr auto lo_lut = bitmore_detail::low_residue_lut<Modulus>();
|
|
constexpr auto hi_lut = bitmore_detail::high_residue_lut<Modulus>();
|
|
const auto lo = bitmore_detail::shuffle<Reg>(
|
|
bitmore_detail::load_lut<Reg>(lo_lut),
|
|
bitmore_detail::low_nibble<Reg>(x));
|
|
const auto hi = bitmore_detail::shuffle<Reg>(
|
|
bitmore_detail::load_lut<Reg>(hi_lut),
|
|
bitmore_detail::high_nibble<Reg>(x));
|
|
const auto sum = bitmore_detail::add_bytes<Reg>(lo, hi);
|
|
// Sum of the two residues is at most 255 and less than `2*M`, so one
|
|
// subtraction finishes the byte. The compare is unsigned: for M > 120
|
|
// the sum can exceed 127.
|
|
const auto limit = bitmore_detail::splat_epi8<Reg>(
|
|
static_cast<unsigned char>(Modulus - 1u));
|
|
const auto ge = bitmore_detail::cmpgt_epu8<Reg>(sum, limit);
|
|
const auto corr = bitmore_detail::and_bytes<Reg>(
|
|
ge, bitmore_detail::splat_epi8<Reg>(static_cast<unsigned char>(Modulus)));
|
|
return bitmore_detail::sub_bytes<Reg>(sum, corr);
|
|
}
|
|
else
|
|
{
|
|
// Four nibble residues, each `< M`. Their sum fits in a 16-bit lane
|
|
// for every modulus through 65535 and is less than `4*M`, so three
|
|
// conditional subtractions finish the lane.
|
|
constexpr auto r0 = bitmore_detail::wide_residue_lut<Modulus, 1u>();
|
|
constexpr auto r1 = bitmore_detail::wide_residue_lut<Modulus, 16u>();
|
|
constexpr auto r2 = bitmore_detail::wide_residue_lut<Modulus, 256u>();
|
|
constexpr auto r3 = bitmore_detail::wide_residue_lut<Modulus, 4096u>();
|
|
constexpr auto r0_lo = bitmore_detail::u16_lo(r0);
|
|
constexpr auto r0_hi = bitmore_detail::u16_hi(r0);
|
|
constexpr auto r1_lo = bitmore_detail::u16_lo(r1);
|
|
constexpr auto r1_hi = bitmore_detail::u16_hi(r1);
|
|
constexpr auto r2_lo = bitmore_detail::u16_lo(r2);
|
|
constexpr auto r2_hi = bitmore_detail::u16_hi(r2);
|
|
constexpr auto r3_lo = bitmore_detail::u16_lo(r3);
|
|
constexpr auto r3_hi = bitmore_detail::u16_hi(r3);
|
|
const auto nib = bitmore_detail::splat_epi16<Reg>(0x000fu);
|
|
const auto n0 = bitmore_detail::and_bytes<Reg>(x, nib);
|
|
const auto n1 = bitmore_detail::and_bytes<Reg>(bitmore_detail::srli_epi16<Reg>(x, 4), nib);
|
|
const auto n2 = bitmore_detail::and_bytes<Reg>(bitmore_detail::srli_epi16<Reg>(x, 8), nib);
|
|
const auto n3 = bitmore_detail::srli_epi16<Reg>(x, 12);
|
|
auto sum = bitmore_detail::lookup_u16<Reg>(n0, r0_lo, r0_hi);
|
|
sum = bitmore_detail::add_epi16<Reg>(sum, bitmore_detail::lookup_u16<Reg>(n1, r1_lo, r1_hi));
|
|
sum = bitmore_detail::add_epi16<Reg>(sum, bitmore_detail::lookup_u16<Reg>(n2, r2_lo, r2_hi));
|
|
sum = bitmore_detail::add_epi16<Reg>(sum, bitmore_detail::lookup_u16<Reg>(n3, r3_lo, r3_hi));
|
|
const auto limit = bitmore_detail::splat_epi16<Reg>(Modulus - 1u);
|
|
const auto modv = bitmore_detail::splat_epi16<Reg>(Modulus);
|
|
for (int step = 0; step < 3; ++step)
|
|
{
|
|
const auto ge = bitmore_detail::cmpgt_epu16<Reg>(sum, limit);
|
|
const auto corr = bitmore_detail::and_bytes<Reg>(ge, modv);
|
|
sum = bitmore_detail::sub_epi16<Reg>(sum, corr);
|
|
}
|
|
return sum;
|
|
}
|
|
}
|
|
};
|
|
|
|
/// @brief Running slot sum modulo `Modulus`, partial-reduced on a budget.
|
|
/// @details `bound()` is a proven upper bound on every slot. An add whose
|
|
/// worst-case total would pass the lane maximum partial-reduces the
|
|
/// accumulator, then the addend. A modulus whose first partial step
|
|
/// still sets the high bit takes a second step before two slots fit
|
|
/// again. `insert_bit` is the MSB-first fold `acc = 2*acc + bit`.
|
|
/// `reduced()` is the full per-slot residue. The byte accumulator
|
|
/// stops at modulus 128 and the 16-bit accumulator at 32768; past
|
|
/// that a resumed slot no longer fits next to another or under a
|
|
/// shift. `partial_reduce` and `full_reduce` cover the wider ranges.
|
|
/// @tparam Modulus server count, `2` through `128` for bytes and `32768` for 16-bit lanes
|
|
/// @tparam Reg `simde__m128i` or `simde__m256i`
|
|
/// @tparam LaneBits `8` or `16`
|
|
template <unsigned Modulus, typename Reg, unsigned LaneBits = 8>
|
|
class bitmore_accumulator
|
|
{
|
|
using mod = bitmore_mod<Modulus, LaneBits>;
|
|
|
|
static_assert(std::is_same_v<Reg, simde__m128i> || std::is_same_v<Reg, simde__m256i>,
|
|
"bitmore_accumulator: register must be simde__m128i or simde__m256i");
|
|
static_assert((LaneBits == 8u && Modulus <= 128u) || (LaneBits == 16u && Modulus <= 32768u),
|
|
"bitmore_accumulator: no slack past modulus 128 in a byte or 32768 in a 16-bit slot");
|
|
static_assert(mod::stable_bound <= (mod::slot_max >> 1),
|
|
"bitmore_accumulator: second partial step must leave room for one bit");
|
|
static_assert(mod::stable_bound * 2u <= mod::slot_max,
|
|
"bitmore_accumulator: two resumed slots must fit in one lane");
|
|
static_assert(mod::shift_budget >= 1u,
|
|
"bitmore_accumulator: at least one bit insertion after resuming");
|
|
|
|
public:
|
|
/// \complexity One zeroed register.
|
|
HEDLEY_ALWAYS_INLINE
|
|
HEDLEY_NO_THROW
|
|
bitmore_accumulator() noexcept
|
|
: acc_(bitmore_detail::zero<Reg>()), bound_(0)
|
|
{ }
|
|
|
|
/// @brief Proven upper bound on every slot in `value()`.
|
|
HEDLEY_PURE
|
|
HEDLEY_ALWAYS_INLINE
|
|
unsigned bound() const noexcept
|
|
{
|
|
return bound_;
|
|
}
|
|
|
|
/// @brief Unreduced slots. Each is `≤ bound()` and congruent to the sum.
|
|
HEDLEY_ALWAYS_INLINE
|
|
Reg value() const noexcept
|
|
{
|
|
return acc_;
|
|
}
|
|
|
|
/// @brief Add one slot.
|
|
/// @param x addends. Each slot must be `≤ addend_max`.
|
|
/// @param addend_max worst-case slot in `x`, clamped to `slot_max`. Pass
|
|
/// `slot_max` when the addend is an arbitrary lane; the accumulator
|
|
/// partial-reduces it if the slack cannot absorb that.
|
|
/// \complexity At most four partial reductions and one lane add.
|
|
HEDLEY_ALWAYS_INLINE
|
|
HEDLEY_NO_THROW
|
|
void add(Reg x, unsigned addend_max) noexcept
|
|
{
|
|
if (addend_max > mod::slot_max)
|
|
addend_max = mod::slot_max;
|
|
for (;;)
|
|
{
|
|
if (bound_ + addend_max <= mod::slot_max)
|
|
break;
|
|
if (bound_ > mod::partial_bound)
|
|
{
|
|
acc_ = mod::partial_reduce(acc_);
|
|
bound_ = mod::partial_bound;
|
|
continue;
|
|
}
|
|
if (addend_max > mod::partial_bound)
|
|
{
|
|
x = mod::partial_reduce(x);
|
|
addend_max = mod::partial_bound;
|
|
continue;
|
|
}
|
|
if (bound_ > mod::stable_bound)
|
|
{
|
|
acc_ = mod::partial_reduce(acc_);
|
|
bound_ = mod::stable_bound;
|
|
continue;
|
|
}
|
|
if (addend_max > mod::stable_bound)
|
|
{
|
|
x = mod::partial_reduce(x);
|
|
addend_max = mod::stable_bound;
|
|
continue;
|
|
}
|
|
break;
|
|
}
|
|
if constexpr (LaneBits == 8u)
|
|
acc_ = bitmore_detail::add_bytes<Reg>(acc_, x);
|
|
else
|
|
acc_ = bitmore_detail::add_epi16<Reg>(acc_, x);
|
|
bound_ += addend_max;
|
|
}
|
|
|
|
/// @brief Fold one BitMore bit: `acc = 2*acc + bit0` inside each slot.
|
|
/// @param bit bit 0 of each slot is the new low bit. Higher bits are ignored.
|
|
/// \complexity Up to two partial reductions when the proven bound sets the
|
|
/// lane's high bit, then a shift and an OR.
|
|
HEDLEY_ALWAYS_INLINE
|
|
HEDLEY_NO_THROW
|
|
void insert_bit(Reg bit) noexcept
|
|
{
|
|
while (bound_ > (mod::slot_max >> 1))
|
|
{
|
|
acc_ = mod::partial_reduce(acc_);
|
|
bound_ = bound_ > mod::partial_bound ? mod::partial_bound
|
|
: mod::stable_bound;
|
|
}
|
|
if constexpr (LaneBits == 8u)
|
|
acc_ = bitmore_detail::shift_in_bit<Reg>(acc_, bit);
|
|
else
|
|
acc_ = bitmore_detail::shift_in_bit16<Reg>(acc_, bit);
|
|
bound_ = bound_ * 2u + 1u;
|
|
}
|
|
|
|
/// @brief Full residue of every slot, in `0 .. Modulus-1`.
|
|
/// \complexity One `full_reduce` of the register.
|
|
HEDLEY_ALWAYS_INLINE
|
|
HEDLEY_NO_THROW
|
|
HEDLEY_PURE
|
|
Reg reduced() const noexcept
|
|
{
|
|
return mod::full_reduce(acc_);
|
|
}
|
|
|
|
private:
|
|
Reg acc_;
|
|
unsigned bound_;
|
|
};
|
|
|
|
} // namespace dpf
|
|
|
|
#endif // LIBDPF_INCLUDE_DPF_BITMORE_MOD_HPP__
|