233 lines
8.4 KiB
C++
233 lines
8.4 KiB
C++
|
|
/// @file dpf/leaf_later.hpp
|
|||
|
|
/// @brief Defer the leaf correction until a public group element is known.
|
|||
|
|
/// @details `eval_full` / `eval_full_add_into` with `leaf_later` skip the leaf
|
|||
|
|
/// correction word, write the uncorrected group share, and fill a
|
|||
|
|
/// parallel control-bit buffer the caller owns. After a rotate of
|
|||
|
|
/// value and control together, `apply_leaf_correction` does
|
|||
|
|
/// `buf[i] += F * control[i]` so Duoram's `v[i−s] + F·t[i−s]` never
|
|||
|
|
/// reapplies a walk.
|
|||
|
|
/// @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_LEAF_LATER_HPP__
|
|||
|
|
#define LIBDPF_INCLUDE_DPF_LEAF_LATER_HPP__
|
|||
|
|
|
|||
|
|
#include <cstddef>
|
|||
|
|
#include <cstdint>
|
|||
|
|
#include <cstring>
|
|||
|
|
#include <limits>
|
|||
|
|
#include <stdexcept>
|
|||
|
|
#include <type_traits>
|
|||
|
|
#include <utility>
|
|||
|
|
#include <vector>
|
|||
|
|
|
|||
|
|
#include "hedley/hedley.h"
|
|||
|
|
|
|||
|
|
#include "dpf/dpf_key.hpp"
|
|||
|
|
#include "dpf/eval_common.hpp"
|
|||
|
|
#include "dpf/eval_interval.hpp"
|
|||
|
|
#include "dpf/interval_memoizer.hpp"
|
|||
|
|
#include "dpf/leaf_arithmetic.hpp"
|
|||
|
|
#include "dpf/leaf_node.hpp"
|
|||
|
|
#include "dpf/twiddle.hpp"
|
|||
|
|
#include "dpf/utils.hpp"
|
|||
|
|
|
|||
|
|
namespace dpf
|
|||
|
|
{
|
|||
|
|
|
|||
|
|
/// @brief Tag: expand without applying the leaf correction word.
|
|||
|
|
struct leaf_later
|
|||
|
|
{
|
|||
|
|
static constexpr bool is_leaf_later_tag = true;
|
|||
|
|
};
|
|||
|
|
|
|||
|
|
template <typename T>
|
|||
|
|
struct is_leaf_later : std::false_type {};
|
|||
|
|
template <>
|
|||
|
|
struct is_leaf_later<leaf_later> : std::true_type {};
|
|||
|
|
template <typename T>
|
|||
|
|
inline constexpr bool is_leaf_later_v = is_leaf_later<std::decay_t<T>>::value;
|
|||
|
|
|
|||
|
|
// Forward declaration: defined in eval_walk.hpp (same offset convention).
|
|||
|
|
struct rotate;
|
|||
|
|
|
|||
|
|
namespace detail_leaf_later
|
|||
|
|
{
|
|||
|
|
|
|||
|
|
template <typename KeyT>
|
|||
|
|
HEDLEY_NO_THROW
|
|||
|
|
constexpr std::size_t domain_size() noexcept
|
|||
|
|
{
|
|||
|
|
constexpr std::size_t bits =
|
|||
|
|
utils::bitlength_of_v<typename KeyT::input_type>;
|
|||
|
|
static_assert(bits < 8 * sizeof(std::size_t),
|
|||
|
|
"leaf_later: input domain does not fit a std::size_t index");
|
|||
|
|
return std::size_t{1} << bits;
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
/// @brief Uncorrected exterior leaf: `0 − mask` (same group as traverse_exterior).
|
|||
|
|
template <std::size_t I = 0, typename DpfKey>
|
|||
|
|
HEDLEY_ALWAYS_INLINE
|
|||
|
|
HEDLEY_NO_THROW
|
|||
|
|
auto traverse_exterior_uncorrected(const typename DpfKey::interior_node & node)
|
|||
|
|
{
|
|||
|
|
using output_type = typename DpfKey::concrete_output_type<I>;
|
|||
|
|
using exterior_prg = typename DpfKey::exterior_prg;
|
|||
|
|
using outputs_tuple = typename DpfKey::concrete_outputs_tuple;
|
|||
|
|
using leaf_type = dpf::leaf_node_t<typename DpfKey::exterior_node, output_type>;
|
|||
|
|
leaf_type zero{};
|
|||
|
|
return dpf::subtract_leaf<output_type>(
|
|||
|
|
zero,
|
|||
|
|
make_leaf_mask_inner<exterior_prg, I, outputs_tuple>(
|
|||
|
|
unset_lo_2bits(node)));
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
template <typename DpfKey, typename = void>
|
|||
|
|
struct has_opl_of : std::false_type {};
|
|||
|
|
template <typename DpfKey>
|
|||
|
|
struct has_opl_of<DpfKey,
|
|||
|
|
std::void_t<decltype(DpfKey::template outputs_per_leaf_of<0>)>>
|
|||
|
|
: std::true_type {};
|
|||
|
|
|
|||
|
|
template <typename DpfKey>
|
|||
|
|
HEDLEY_NO_THROW
|
|||
|
|
constexpr std::size_t opl_of() noexcept
|
|||
|
|
{
|
|||
|
|
if constexpr (has_opl_of<DpfKey>::value)
|
|||
|
|
return DpfKey::template outputs_per_leaf_of<0>;
|
|||
|
|
else
|
|||
|
|
return DpfKey::outputs_per_leaf;
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
} // namespace detail_leaf_later
|
|||
|
|
|
|||
|
|
/// @brief Full-domain expansion without the leaf correction; also fills `control`.
|
|||
|
|
/// @details `buf[i]` receives the uncorrected share at domain point `i`, and
|
|||
|
|
/// `control[i]` is the leaf node's control bit. Both must be sized to
|
|||
|
|
/// `2^n`.
|
|||
|
|
/// \complexity One full-domain expansion, `Θ(2^n)`.
|
|||
|
|
template <std::size_t I = 0, typename Buffer, typename Control, typename DpfKey,
|
|||
|
|
std::enable_if_t<looks_like_dpf_key_v<DpfKey>
|
|||
|
|
&& !is_multilevel_key_v<DpfKey>, bool> = true>
|
|||
|
|
void eval_full(Buffer & buf, Control & control, const DpfKey & dpf,
|
|||
|
|
leaf_later) // NOLINT(runtime/references)
|
|||
|
|
{
|
|||
|
|
using input_type = typename DpfKey::input_type;
|
|||
|
|
using output_type = typename DpfKey::concrete_output_type<I>;
|
|||
|
|
using exterior_node = typename DpfKey::exterior_node;
|
|||
|
|
constexpr std::size_t opl = DpfKey::outputs_per_leaf;
|
|||
|
|
const std::size_t n = detail_leaf_later::domain_size<DpfKey>();
|
|||
|
|
if (buf.size() < n || control.size() < n)
|
|||
|
|
throw std::invalid_argument("eval_full(leaf_later): buffer too small");
|
|||
|
|
|
|||
|
|
auto memo = make_basic_full_memoizer(dpf);
|
|||
|
|
const input_type from = std::numeric_limits<input_type>::min();
|
|||
|
|
const input_type to = std::numeric_limits<input_type>::max();
|
|||
|
|
const auto from_node = utils::get_from_node<DpfKey>(from);
|
|||
|
|
const auto to_node = utils::get_to_node<DpfKey>(to);
|
|||
|
|
internal::eval_interval_interior(dpf, from_node, to_node, memo);
|
|||
|
|
|
|||
|
|
const std::size_t nodes_in_interval =
|
|||
|
|
static_cast<std::size_t>(to_node - from_node);
|
|||
|
|
auto * nodes = memo[DpfKey::depth];
|
|||
|
|
|
|||
|
|
for (std::size_t j = 0; j < nodes_in_interval; ++j)
|
|||
|
|
{
|
|||
|
|
const auto & node = nodes[j];
|
|||
|
|
const bool tbit = static_cast<bool>(get_lo_bit(node));
|
|||
|
|
auto leaf = detail_leaf_later::traverse_exterior_uncorrected<I, DpfKey>(
|
|||
|
|
node);
|
|||
|
|
for (std::size_t o = 0; o < opl; ++o)
|
|||
|
|
{
|
|||
|
|
const std::size_t idx = j * opl + o;
|
|||
|
|
if (idx >= n)
|
|||
|
|
break;
|
|||
|
|
output_type y = extract_leaf<exterior_node, output_type>(leaf, o);
|
|||
|
|
buf[idx] = y;
|
|||
|
|
control[idx] = tbit;
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
/// @brief Add an uncorrected full-domain expansion into `buf` and fill `control`.
|
|||
|
|
template <std::size_t I = 0, typename Buffer, typename Control, typename DpfKey,
|
|||
|
|
std::enable_if_t<looks_like_dpf_key_v<DpfKey>
|
|||
|
|
&& !is_multilevel_key_v<DpfKey>, bool> = true>
|
|||
|
|
void eval_full_add_into(Buffer & buf, Control & control, const DpfKey & dpf,
|
|||
|
|
leaf_later) // NOLINT(runtime/references)
|
|||
|
|
{
|
|||
|
|
using output_type = typename DpfKey::concrete_output_type<I>;
|
|||
|
|
const std::size_t n = detail_leaf_later::domain_size<DpfKey>();
|
|||
|
|
std::vector<output_type> tmp(n);
|
|||
|
|
std::vector<std::uint8_t> ctl(n);
|
|||
|
|
eval_full<I>(tmp, ctl, dpf, leaf_later{});
|
|||
|
|
for (std::size_t i = 0; i < n; ++i)
|
|||
|
|
{
|
|||
|
|
buf[i] = buf[i] + tmp[i];
|
|||
|
|
control[i] = static_cast<bool>(ctl[i]);
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
/// @brief Uncorrected expansion written at `(i + rot.shift) mod 2^n`, with
|
|||
|
|
/// control bits rotated the same way.
|
|||
|
|
template <std::size_t I = 0, typename Buffer, typename Control, typename DpfKey,
|
|||
|
|
std::enable_if_t<looks_like_dpf_key_v<DpfKey>
|
|||
|
|
&& !is_multilevel_key_v<DpfKey>, bool> = true>
|
|||
|
|
void eval_full_add_into(Buffer & buf, Control & control, const DpfKey & dpf,
|
|||
|
|
leaf_later, rotate rot) // NOLINT(runtime/references)
|
|||
|
|
{
|
|||
|
|
using output_type = typename DpfKey::concrete_output_type<I>;
|
|||
|
|
const std::size_t n = detail_leaf_later::domain_size<DpfKey>();
|
|||
|
|
const std::size_t s = rot.shift % n;
|
|||
|
|
std::vector<output_type> tmp(n);
|
|||
|
|
std::vector<std::uint8_t> ctl(n);
|
|||
|
|
eval_full<I>(tmp, ctl, dpf, leaf_later{});
|
|||
|
|
for (std::size_t i = 0; i < n; ++i)
|
|||
|
|
{
|
|||
|
|
const std::size_t j = (i + s) % n;
|
|||
|
|
buf[j] = buf[j] + tmp[i];
|
|||
|
|
control[j] = static_cast<bool>(ctl[i]);
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
/// @brief Rotate value and control by the same offset (`new[i] = old[(i-s) mod n]`).
|
|||
|
|
/// @details Domain point `k` lands at `(k + s) mod n`, matching `dpf::rotate{s}`.
|
|||
|
|
template <typename Buffer, typename Control>
|
|||
|
|
void cyclic_shift_pair(Buffer & buf, Control & control, std::size_t s) // NOLINT(runtime/references)
|
|||
|
|
{
|
|||
|
|
const std::size_t n = buf.size();
|
|||
|
|
if (n == 0 || control.size() != n)
|
|||
|
|
throw std::invalid_argument("cyclic_shift_pair: size mismatch");
|
|||
|
|
const std::size_t sh = s % n;
|
|||
|
|
if (sh == 0)
|
|||
|
|
return;
|
|||
|
|
Buffer out_buf(buf);
|
|||
|
|
Control out_ctl(control);
|
|||
|
|
for (std::size_t i = 0; i < n; ++i)
|
|||
|
|
{
|
|||
|
|
out_buf[i] = buf[(i + n - sh) % n];
|
|||
|
|
out_ctl[i] = control[(i + n - sh) % n];
|
|||
|
|
}
|
|||
|
|
buf = std::move(out_buf);
|
|||
|
|
control = std::move(out_ctl);
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
/// @brief `buf[i] += F * control[i]` in the leaf's group.
|
|||
|
|
/// @details XOR leaves treat `+` as XOR. `F` is a public group element.
|
|||
|
|
template <typename Buffer, typename Control, typename F>
|
|||
|
|
void apply_leaf_correction(Buffer & buf, const Control & control, const F & f) // NOLINT(runtime/references)
|
|||
|
|
{
|
|||
|
|
const std::size_t n = buf.size();
|
|||
|
|
if (control.size() < n)
|
|||
|
|
throw std::invalid_argument("apply_leaf_correction: control too small");
|
|||
|
|
for (std::size_t i = 0; i < n; ++i)
|
|||
|
|
{
|
|||
|
|
if (static_cast<bool>(control[i]))
|
|||
|
|
buf[i] = buf[i] + f;
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
} // namespace dpf
|
|||
|
|
|
|||
|
|
#endif // LIBDPF_INCLUDE_DPF_LEAF_LATER_HPP__
|