libdpf/include/dpf/interleave_leaves.hpp
Ryan Henry 0d22946a0e Checkpoint the party/runtime stack before share-program and malicious-mode work.
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>
2026-09-28 05:59:19 -06:00

568 lines
20 KiB
C++
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

/// @file dpf/interleave_leaves.hpp
/// @brief Interleave (and deinterleave) per-key leaf vectors into cohort order.
/// @details Given `nkeys` vectors of length `nleaves`, write one buffer whose
/// logical lane index is `cohort_index(i, k, nkeys) = i * nkeys + k`
/// (leaf `i` of key `k`). That matches sequence-cohort layout and the
/// key-major axis of interval-cohort leaves (lanes inside a leaf stay
/// contiguous in each one-key buffer; this routine interleaves whole
/// leaf *values*, not interior walk nodes).
///
/// Storage and packing:
/// - `uint64_t` / `uint32_t` / `uint16_t` / `uint8_t`: one array element
/// per leaf; `out[i * nkeys + k] = keys[k][i]`.
/// - `dpf::bit`, `dpf::twobit`, `dpf::nyble`, and the matching
/// `dpf::gf2` / `gf22` / `gf24` lanes: same logical order, packed
/// tightly into `uint64_t` words like `dynamic_bit_array` /
/// `dynamic_packed_array` — lane `j` occupies bits
/// `[j * W, (j + 1) * W)` of the flat stream (`W` = 1, 2, or 4),
/// low lane in the low bits of each word. So a naive *element*
/// index `i * nkeys + k` is wrong for sub-byte widths; use the
/// logical index above, then pack with width `W`.
///
/// Kernels are word-wise (and 8×8 bit transpose batches for 1-bit),
/// not a scalar loop per bit. Deinterleave is the inverse transpose.
/// @copyright Copyright (c) 2019-2026 Ryan Henry and [others](@ref authors)
/// @license Released under a GNU General Public v2.0 (GPLv2) license;
/// see [LICENSE.md](@ref license) for details.
#ifndef LIBDPF_INCLUDE_DPF_INTERLEAVE_LEAVES_HPP__
#define LIBDPF_INCLUDE_DPF_INTERLEAVE_LEAVES_HPP__
#include <cstddef>
#include <cstdint>
#include <cstring>
#include <type_traits>
#include "hedley/hedley.h"
#include "dpf/bit.hpp"
#include "dpf/modint.hpp"
#include "dpf/nyble.hpp"
#include "dpf/twobit.hpp"
#include "dpf/utils.hpp"
namespace dpf
{
/// @brief Underlying word/element type used to store a sequence of `LaneT` leaves.
template <typename LaneT, typename Enable = void>
struct leaf_storage
{
using type = LaneT;
};
template <typename LaneT>
struct leaf_storage<LaneT, std::enable_if_t<utils::is_packed_subbyte_v<LaneT>>>
{
using type = std::uint64_t;
};
template <>
struct leaf_storage<dpf::bit>
{
using type = std::uint64_t;
};
template <>
struct leaf_storage<dpf::twobit>
{
using type = std::uint64_t;
};
template <>
struct leaf_storage<dpf::nyble>
{
using type = std::uint64_t;
};
template <typename LaneT>
using leaf_storage_t = typename leaf_storage<LaneT>::type;
/// @brief Number of `leaf_storage_t<LaneT>` words needed for `n` logical leaves.
template <typename LaneT>
HEDLEY_CONST
HEDLEY_NO_THROW
constexpr std::size_t leaf_storage_words(std::size_t n) noexcept
{
if constexpr (utils::is_packed_subbyte_v<LaneT>)
{
constexpr std::size_t w = utils::packed_lane_bits_v<LaneT>;
return (n * w + 63u) / 64u;
}
else
{
return n;
}
}
namespace interleave_detail
{
/// @brief 8×8 bit matrix transpose in a 64-bit word (Hacker's Delight).
/// @details Byte `r` holds row `r`; bit `c` of that byte is column `c`
/// (LSB = column 0). Transpose swaps rows and columns.
HEDLEY_CONST
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
constexpr std::uint64_t bit_transpose_8x8(std::uint64_t x) noexcept
{
std::uint64_t t = (x ^ (x >> 7)) & 0x00AA00AA00AA00AAull;
x = x ^ t ^ (t << 7);
t = (x ^ (x >> 14)) & 0x0000CCCC0000CCCCull;
x = x ^ t ^ (t << 14);
t = (x ^ (x >> 28)) & 0x00000000F0F0F0F0ull;
x = x ^ t ^ (t << 28);
return x;
}
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
void write_bits(std::uint64_t * out, std::size_t bit_pos, std::uint64_t bits,
unsigned nbits) noexcept
{
if (nbits == 0)
return;
const std::size_t word = bit_pos / 64u;
const unsigned shift = static_cast<unsigned>(bit_pos % 64u);
const std::uint64_t mask = nbits == 64
? ~std::uint64_t{0}
: ((std::uint64_t{1} << nbits) - 1u);
const std::uint64_t val = bits & mask;
out[word] |= val << shift;
if (shift + nbits > 64u)
out[word + 1u] |= val >> (64u - shift);
}
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
std::uint64_t read_bits(const std::uint64_t * in, std::size_t bit_pos,
unsigned nbits) noexcept
{
if (nbits == 0)
return 0;
const std::size_t word = bit_pos / 64u;
const unsigned shift = static_cast<unsigned>(bit_pos % 64u);
const std::uint64_t mask = nbits == 64
? ~std::uint64_t{0}
: ((std::uint64_t{1} << nbits) - 1u);
std::uint64_t val = in[word] >> shift;
if (shift + nbits > 64u)
val |= in[word + 1u] << (64u - shift);
return val & mask;
}
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
std::uint8_t load_key_byte_bits(const std::uint64_t * key, std::size_t nleaves,
std::size_t i0) noexcept
{
if (i0 >= nleaves)
return 0;
const unsigned take = static_cast<unsigned>(
nleaves - i0 < 8u ? nleaves - i0 : 8u);
return static_cast<std::uint8_t>(read_bits(key, i0, take));
}
/// @brief Interleave 1-bit lanes with 8×8 transpose batches plus a bit writer.
HEDLEY_NO_THROW
inline void interleave_bits(std::uint64_t * HEDLEY_RESTRICT out,
const std::uint64_t * const * HEDLEY_RESTRICT keys, std::size_t nkeys,
std::size_t nleaves) noexcept
{
const std::size_t out_words = leaf_storage_words<dpf::bit>(nkeys * nleaves);
if (out_words != 0)
std::memset(out, 0, out_words * sizeof(std::uint64_t));
if (nkeys == 0 || nleaves == 0)
return;
for (std::size_t i0 = 0; i0 < nleaves; i0 += 8u)
{
const unsigned nrows = static_cast<unsigned>(
nleaves - i0 < 8u ? nleaves - i0 : 8u);
for (std::size_t k0 = 0; k0 < nkeys; k0 += 8u)
{
const unsigned ncols = static_cast<unsigned>(
nkeys - k0 < 8u ? nkeys - k0 : 8u);
std::uint64_t packed = 0;
DPF_UNROLL_LOOP
for (unsigned c = 0; c < 8u; ++c)
{
std::uint8_t b = 0;
if (c < ncols)
b = load_key_byte_bits(keys[k0 + c], nleaves, i0);
packed |= static_cast<std::uint64_t>(b) << (8u * c);
}
const std::uint64_t t = bit_transpose_8x8(packed);
DPF_UNROLL_LOOP
for (unsigned r = 0; r < nrows; ++r)
{
const std::uint64_t row
= (t >> (8u * r)) & ((std::uint64_t{1} << ncols) - 1u);
write_bits(out, (i0 + r) * nkeys + k0, row, ncols);
}
}
}
}
HEDLEY_NO_THROW
inline void deinterleave_bits(std::uint64_t * const * HEDLEY_RESTRICT keys,
const std::uint64_t * HEDLEY_RESTRICT in, std::size_t nkeys,
std::size_t nleaves) noexcept
{
if (nkeys == 0 || nleaves == 0)
return;
for (std::size_t k = 0; k < nkeys; ++k)
{
const std::size_t words = leaf_storage_words<dpf::bit>(nleaves);
if (words != 0)
std::memset(keys[k], 0, words * sizeof(std::uint64_t));
}
for (std::size_t i0 = 0; i0 < nleaves; i0 += 8u)
{
const unsigned nrows = static_cast<unsigned>(
nleaves - i0 < 8u ? nleaves - i0 : 8u);
for (std::size_t k0 = 0; k0 < nkeys; k0 += 8u)
{
const unsigned ncols = static_cast<unsigned>(
nkeys - k0 < 8u ? nkeys - k0 : 8u);
std::uint64_t packed = 0;
DPF_UNROLL_LOOP
for (unsigned r = 0; r < nrows; ++r)
{
const std::uint64_t row
= read_bits(in, (i0 + r) * nkeys + k0, ncols);
packed |= row << (8u * r);
}
// Inverse of the interleave transpose: same 8×8 transpose.
const std::uint64_t t = bit_transpose_8x8(packed);
DPF_UNROLL_LOOP
for (unsigned c = 0; c < ncols; ++c)
{
const std::uint64_t col
= (t >> (8u * c)) & ((std::uint64_t{1} << nrows) - 1u);
write_bits(keys[k0 + c], i0, col, nrows);
}
}
}
}
/// @brief Interleave `W`-bit lanes (`W` = 2 or 4) by filling output words.
template <unsigned W>
HEDLEY_NO_THROW
void interleave_packed_w(std::uint64_t * HEDLEY_RESTRICT out,
const std::uint64_t * const * HEDLEY_RESTRICT keys, std::size_t nkeys,
std::size_t nleaves) noexcept
{
static_assert(W == 2 || W == 4, "packed width must be 2 or 4");
constexpr std::uint64_t lane_mask = (std::uint64_t{1} << W) - 1u;
constexpr unsigned lanes_per_word = 64u / W;
const std::size_t total = nkeys * nleaves;
const std::size_t out_words = (total * W + 63u) / 64u;
if (out_words != 0)
std::memset(out, 0, out_words * sizeof(std::uint64_t));
if (nkeys == 0 || nleaves == 0)
return;
for (std::size_t ow = 0; ow < out_words; ++ow)
{
std::uint64_t word = 0;
const std::size_t base = ow * lanes_per_word;
const unsigned nlanes = static_cast<unsigned>(
total - base < lanes_per_word ? total - base : lanes_per_word);
DPF_UNROLL_LOOP
for (unsigned t = 0; t < nlanes; ++t)
{
const std::size_t g = base + t;
const std::size_t i = g / nkeys;
const std::size_t k = g - i * nkeys;
const std::size_t src_bit = i * W;
const std::uint64_t lane
= (keys[k][src_bit / 64u] >> (src_bit % 64u)) & lane_mask;
word |= lane << (t * W);
}
out[ow] = word;
}
}
template <unsigned W>
HEDLEY_NO_THROW
void deinterleave_packed_w(std::uint64_t * const * HEDLEY_RESTRICT keys,
const std::uint64_t * HEDLEY_RESTRICT in, std::size_t nkeys,
std::size_t nleaves) noexcept
{
static_assert(W == 2 || W == 4, "packed width must be 2 or 4");
constexpr std::uint64_t lane_mask = (std::uint64_t{1} << W) - 1u;
if (nkeys == 0 || nleaves == 0)
return;
for (std::size_t k = 0; k < nkeys; ++k)
{
const std::size_t words = (nleaves * W + 63u) / 64u;
if (words != 0)
std::memset(keys[k], 0, words * sizeof(std::uint64_t));
}
const std::size_t total = nkeys * nleaves;
for (std::size_t g = 0; g < total; ++g)
{
const std::size_t i = g / nkeys;
const std::size_t k = g - i * nkeys;
const std::size_t src_bit = g * W;
const std::uint64_t lane
= (in[src_bit / 64u] >> (src_bit % 64u)) & lane_mask;
const std::size_t dst_bit = i * W;
keys[k][dst_bit / 64u] |= lane << (dst_bit % 64u);
}
}
template <typename T>
HEDLEY_NO_THROW
void interleave_wide(T * HEDLEY_RESTRICT out, const T * const * HEDLEY_RESTRICT keys,
std::size_t nkeys, std::size_t nleaves) noexcept
{
if (nkeys == 0 || nleaves == 0)
return;
if (nkeys == 1)
{
std::memcpy(out, keys[0], nleaves * sizeof(T));
return;
}
for (std::size_t i = 0; i < nleaves; ++i)
{
T * dst = out + i * nkeys;
DPF_UNROLL_LOOP
for (std::size_t k = 0; k < nkeys; ++k)
dst[k] = keys[k][i];
}
}
template <typename T>
HEDLEY_NO_THROW
void deinterleave_wide(T * const * HEDLEY_RESTRICT keys,
const T * HEDLEY_RESTRICT in, std::size_t nkeys, std::size_t nleaves) noexcept
{
if (nkeys == 0 || nleaves == 0)
return;
if (nkeys == 1)
{
std::memcpy(keys[0], in, nleaves * sizeof(T));
return;
}
for (std::size_t i = 0; i < nleaves; ++i)
{
const T * src = in + i * nkeys;
DPF_UNROLL_LOOP
for (std::size_t k = 0; k < nkeys; ++k)
keys[k][i] = src[k];
}
}
} // namespace interleave_detail
/// @brief Interleave `nkeys` leaf vectors of length `nleaves` into `out`.
/// @tparam LaneT an integer, a packed lane (`bit`, `twobit`, `nyble`,
/// `gf2`, `gf22`, `gf24`), or a trivially copyable element of at
/// most 8 bytes (`gf28`, `gf216`, `gf232`, `gf264`)
/// @param out destination storage (`leaf_storage_t<LaneT>`; see file comment)
/// @param keys array of `nkeys` pointers to per-key leaf storage
/// @param nkeys number of keys (`m`)
/// @param nleaves number of leaves per key (`L`)
/// \complexity Θ(`nkeys * nleaves`) lane moves. Wide types stream by leaf;
/// 1-bit uses 8×8 transpose tiles; 2/4-bit fill output words.
template <typename LaneT>
HEDLEY_NO_THROW
void interleave_leaves(leaf_storage_t<LaneT> * out,
const leaf_storage_t<LaneT> * const * keys, std::size_t nkeys,
std::size_t nleaves) noexcept
{
if constexpr (utils::is_packed_subbyte_v<LaneT>)
{
constexpr unsigned w = utils::packed_lane_bits_v<LaneT>;
if constexpr (w == 1)
interleave_detail::interleave_bits(out, keys, nkeys, nleaves);
else if constexpr (w == 2)
interleave_detail::interleave_packed_w<2>(out, keys, nkeys, nleaves);
else
{
static_assert(w == 4, "packed lane width must be 1, 2, or 4");
interleave_detail::interleave_packed_w<4>(out, keys, nkeys, nleaves);
}
}
else if constexpr ((std::is_integral_v<LaneT> && !std::is_same_v<LaneT, bool>)
|| (std::is_trivially_copyable_v<LaneT> && std::is_standard_layout_v<LaneT>
&& sizeof(LaneT) > 0 && sizeof(LaneT) <= 8))
{
interleave_detail::interleave_wide(out, keys, nkeys, nleaves);
}
else
{
static_assert(std::is_integral_v<LaneT>,
"interleave_leaves LaneT must be an integer, a packed lane, or a field element of at most 8 bytes");
}
}
/// @brief Inverse of `interleave_leaves`: split one interleaved buffer into
/// `nkeys` per-key leaf vectors of length `nleaves`.
/// @tparam LaneT same set as `interleave_leaves`
/// @param keys destination per-key storage pointers
/// @param in interleaved source
/// @param nkeys number of keys
/// @param nleaves number of leaves per key
/// \complexity Same order as `interleave_leaves`.
template <typename LaneT>
HEDLEY_NO_THROW
void deinterleave_leaves(leaf_storage_t<LaneT> * const * keys,
const leaf_storage_t<LaneT> * in, std::size_t nkeys,
std::size_t nleaves) noexcept
{
if constexpr (utils::is_packed_subbyte_v<LaneT>)
{
constexpr unsigned w = utils::packed_lane_bits_v<LaneT>;
if constexpr (w == 1)
interleave_detail::deinterleave_bits(keys, in, nkeys, nleaves);
else if constexpr (w == 2)
interleave_detail::deinterleave_packed_w<2>(keys, in, nkeys, nleaves);
else
{
static_assert(w == 4, "packed lane width must be 1, 2, or 4");
interleave_detail::deinterleave_packed_w<4>(keys, in, nkeys, nleaves);
}
}
else if constexpr ((std::is_integral_v<LaneT> && !std::is_same_v<LaneT, bool>)
|| (std::is_trivially_copyable_v<LaneT> && std::is_standard_layout_v<LaneT>
&& sizeof(LaneT) > 0 && sizeof(LaneT) <= 8))
{
interleave_detail::deinterleave_wide(keys, in, nkeys, nleaves);
}
else
{
static_assert(std::is_integral_v<LaneT>,
"deinterleave_leaves LaneT must be an integer, a packed lane, or a field element of at most 8 bytes");
}
}
/// @brief Dot an interleaved 1-bit cohort with `weights`, as `modint<Nbits>`.
/// @details `bits` is the buffer from `interleave_leaves<dpf::bit>`: leaf `i`
/// is the integer whose bit `k` is key `k`'s leaf `i`. That integer,
/// reduced modulo `2^Nbits` (the low `Nbits` bits), is one `modint`.
/// The result is `sum_i modint(leaf i) * weights[i]`. `weights[i]`
/// may be a `modint<Nbits>` or an integer.
/// When `nkeys == Nbits` and that width is 8, 16, 32, 64, or 128,
/// the leaves are contiguous native words and the product is a
/// straight multiply-accumulate.
/// @tparam Nbits modulus width, `1` through `256`
/// @param bits interleaved 1-bit stream
/// @param nleaves number of leaves (domain points)
/// @param nkeys number of keys that were interleaved
/// @param weights one weight per leaf
/// @return the inner product in `modint<Nbits>`
/// \complexity Θ(`nleaves`) multiplications in the `modint` word. Aligned
/// widths do one native multiply per leaf; other widths extract
/// the `Nbits` bits of each leaf first.
template <std::size_t Nbits, typename Weights>
HEDLEY_NO_THROW
HEDLEY_WARN_UNUSED_RESULT
modint<Nbits> interleaved_bits_inner_product(const std::uint64_t * bits,
std::size_t nleaves, std::size_t nkeys, Weights && weights)
{
using mod = modint<Nbits>;
using limb = typename mod::integral_type;
mod acc{};
if (nleaves == 0 || nkeys == 0 || bits == nullptr)
return acc;
auto limb_of = [](auto v) -> limb {
using T = std::decay_t<decltype(v)>;
if constexpr (std::is_same_v<T, mod>)
return static_cast<limb>(v);
else
return static_cast<limb>(mod(static_cast<limb>(v)));
};
auto mac = [&](limb v, std::size_t i) {
acc += mod(v) * mod(limb_of(weights[i]));
};
if constexpr (Nbits == 128)
{
if (nkeys == 128)
{
limb sum{};
for (std::size_t i = 0; i < nleaves; ++i)
{
const std::uint64_t * p = bits + i * 2u;
const limb v = static_cast<limb>(p[0])
| (static_cast<limb>(p[1]) << 64);
sum += v * limb_of(weights[i]);
}
return mod(sum);
}
}
if constexpr (Nbits == 8 || Nbits == 16 || Nbits == 32 || Nbits == 64)
{
if (nkeys == Nbits)
{
limb sum{};
for (std::size_t i = 0; i < nleaves; ++i)
{
limb v{};
if constexpr (Nbits == 64)
v = static_cast<limb>(bits[i]);
else if constexpr (Nbits == 32)
{
const std::uint64_t word = bits[i / 2u];
v = static_cast<limb>((i & 1u) ? (word >> 32) : (word & 0xffffffffu));
}
else if constexpr (Nbits == 16)
{
const std::uint64_t word = bits[i / 4u];
v = static_cast<limb>((word >> ((i % 4u) * 16u)) & 0xffffu);
}
else
{
const std::uint64_t word = bits[i / 8u];
v = static_cast<limb>((word >> ((i % 8u) * 8u)) & 0xffu);
}
sum += v * limb_of(weights[i]);
}
return mod(sum);
}
}
for (std::size_t i = 0; i < nleaves; ++i)
{
const std::size_t bit_pos = i * nkeys;
limb v{};
if constexpr (Nbits <= 64)
{
const unsigned take = static_cast<unsigned>(
nkeys < Nbits ? nkeys : Nbits);
v = static_cast<limb>(interleave_detail::read_bits(bits, bit_pos, take));
}
else
{
unsigned left = static_cast<unsigned>(Nbits);
std::size_t pos = bit_pos;
unsigned shift = 0;
while (left != 0 && shift < nkeys)
{
const unsigned room = static_cast<unsigned>(nkeys - shift);
const unsigned n = left < 64u ? left : 64u;
const unsigned take = n < room ? n : room;
const auto chunk = interleave_detail::read_bits(bits, pos, take);
v |= static_cast<limb>(chunk) << shift;
shift += take;
pos += take;
left -= take;
if (take == 0)
break;
}
}
mac(v, i);
}
return acc;
}
} // namespace dpf
#endif // LIBDPF_INCLUDE_DPF_INTERLEAVE_LEAVES_HPP__