/// @file dpf/cohort.hpp /// @brief Many classic DPFs that share one public schedule. /// @details Generation is the ordinary special-path walk, run level by level /// across keys so the path bit is read once and `expand_x4` sees /// contiguous seeds. Evaluation uses that same idea on a point, a /// closed interval, or one `sequence_recipe`: correction words are /// hoisted once per level, and interior nodes stay key-interleaved /// (`position * stride + key`) down to the leaves. /// /// Leaf layout, with `n` keys: /// - point: `out[key]` /// - interval: leaf node `j`, lane `p` of key `k` is /// `out[(j * n + k) * outputs_per_leaf + p]` /// (same lane order as one-key `eval_interval`) /// - sequence: listed point `q` of key `k` is `out[q * n + k]` /// /// `cohort_index(j, k, n)` is `j * n + k`. /// Inner products use the same order and the same weights for every /// key. Generation is the classic (not incremental) dealer key: one /// plaintext point, one payload or `std::tuple` of payloads per key. /// @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_COHORT_HPP__ #define LIBDPF_INCLUDE_DPF_COHORT_HPP__ #include "hedley/hedley.h" #include #include #include #include #include #include #include #include #include #include #include "dpf/aligned_allocator.hpp" #include "dpf/dpf_key.hpp" #include "dpf/eval_common.hpp" #include "dpf/leaf_node.hpp" #include "dpf/random.hpp" #include "dpf/sequence_recipe.hpp" #include "dpf/utils.hpp" namespace dpf { /// @brief `j * nkeys + key`. Point, interval leaf, and sequence point all use this. HEDLEY_CONST HEDLEY_NO_THROW inline constexpr std::size_t cohort_index(std::size_t j, std::size_t key, std::size_t nkeys) noexcept { return j * nkeys + key; } /// @brief Two planes of key-interleaved interior nodes, reused across calls. /// @tparam Key a DPF key or a `party_key` of one. Every key in a walk has this type. template class cohort_scratch { public: using key_type = unwrap_party_key_t>; using node_type = typename key_type::interior_node; using allocator = aligned_allocator; HEDLEY_NO_THROW std::size_t count() const noexcept { return count_; } HEDLEY_NO_THROW std::size_t stride() const noexcept { return stride_; } /// @brief Grow so each of `positions` tree slots can hold `nkeys` nodes, /// padded to a multiple of 4 for `expand_x4`. void fit(std::size_t nkeys, std::size_t positions) { count_ = nkeys; stride_ = nkeys == 0 ? 0 : ((nkeys + 3u) & ~std::size_t{3}); const std::size_t pos = positions == 0 ? 1 : positions; const std::size_t need = 2 * pos * stride_; if (buf_.size() < need) buf_.resize(need); positions_ = pos; if (cw0_.size() < nkeys) { cw0_.resize(nkeys); cw1_.resize(nkeys); } } node_type * plane(int which) noexcept { return buf_.data() + static_cast(which) * positions_ * stride_; } node_type * cw0() noexcept { return cw0_.data(); } node_type * cw1() noexcept { return cw1_.data(); } private: std::size_t count_ = 0; std::size_t stride_ = 0; std::size_t positions_ = 0; std::vector buf_; std::vector cw0_; std::vector cw1_; }; namespace cohort_detail { template struct is_std_tuple : std::false_type {}; template struct is_std_tuple> : std::true_type {}; template HEDLEY_ALWAYS_INLINE void expand_keys(const Node * HEDLEY_RESTRICT parents, std::size_t nkeys, bool is_last, const Node * HEDLEY_RESTRICT cw_left, const Node * HEDLEY_RESTRICT cw_right, int which, Node * HEDLEY_RESTRICT dest_left, Node * HEDLEY_RESTRICT dest_right) { // which: 0 left, 1 right, 2 both. Both destinations are real pointers. std::size_t k = 0; for (; k + 4 <= nkeys; k += 4) { alignas(Node) Node left[4]; alignas(Node) Node right[4]; Tree::expand_x4(parents + k, left, right, is_last); DPF_UNROLL_LOOP for (std::size_t t = 0; t < 4; ++t) { if (which != 1) { dest_left[k + t] = dpf::xor_if_lo_bit(left[t], cw_left[k + t], parents[k + t]); } if (which != 0) { dest_right[k + t] = dpf::xor_if_lo_bit(right[t], cw_right[k + t], parents[k + t]); } } } for (; k < nkeys; ++k) { const auto kids = Tree::expand(parents[k], is_last); if (which != 1) dest_left[k] = dpf::xor_if_lo_bit(kids[0], cw_left[k], parents[k]); if (which != 0) dest_right[k] = dpf::xor_if_lo_bit(kids[1], cw_right[k], parents[k]); } } template HEDLEY_ALWAYS_INLINE void expand_one(const Node * HEDLEY_RESTRICT parents, std::size_t nkeys, bool is_last, const Node * HEDLEY_RESTRICT cw, bool right, Node * HEDLEY_RESTRICT dest) { Node discard; if (right) expand_keys(parents, nkeys, is_last, cw, cw, 1, &discard, dest); else expand_keys(parents, nkeys, is_last, cw, cw, 0, dest, &discard); } template void require_keys(const Range & keys) { if (keys.size() == 0) throw std::invalid_argument("cohort: no keys"); } template std::size_t nodes_at_level(std::size_t level, Integral from_node, Integral to_node) { const std::size_t offset = Key::depth - level; return static_cast( utils::shift_right(static_cast(to_node - Integral{1}), offset) - utils::shift_right(from_node, offset)) + 1; } template HEDLEY_ALWAYS_INLINE Output lane_value(const Leaf & leaf, std::size_t lane) { if constexpr (utils::is_packed_subbyte_v) { return extract_leaf(leaf, lane); } else { Output val; std::memcpy(&val, reinterpret_cast(std::addressof(leaf)) + lane * sizeof(Output), sizeof(Output)); return val; } } template HEDLEY_ALWAYS_INLINE Output mac_add(Output acc, Output val, W && w) { if constexpr (std::is_integral_v && !std::is_same_v) { using U = std::make_unsigned_t; return static_cast( static_cast(acc) + static_cast(val) * static_cast(w)); } else { return static_cast(acc + val * Output(w)); } } template void ensure_size(Out & out, std::size_t n) { if (out.size() < n) out.resize(n); } template void walk_point(const Range & keys, Input walk_x, Input lane_x, Scratch & scratch, Fn && fn) { using key_type = unwrap_party_key_t>; using node = typename key_type::interior_node; using tree = typename key_type::tree; using output = typename key_type::template concrete_output_type; const std::size_t n = keys.size(); scratch.fit(n, 1); int src = 0; auto * cur = scratch.plane(src); for (std::size_t k = 0; k < n; ++k) cur[k] = keys[k].root(); auto mask = key_type::msb_mask; for (std::size_t level = 0; level < key_type::depth; ++level, mask >>= 1) { const bool bit = !!(mask & walk_x); const bool is_last = tree::is_last_level(level, key_type::depth); node * cw = bit ? scratch.cw1() : scratch.cw0(); for (std::size_t k = 0; k < n; ++k) cw[k] = keys[k].correction_word(level, bit); const int dst = 1 - src; node * next = scratch.plane(dst); if (bit) { expand_one(cur, n, is_last, cw, true, next); } else { expand_one(cur, n, is_last, cw, false, next); } src = dst; cur = next; } for (std::size_t k = 0; k < n; ++k) { auto leaf = keys[k].template traverse_exterior(cur[k]); auto handle = make_dpf_output(leaf, lane_x); fn(k, static_cast(handle)); } } template void walk_interval_segment(const Range & keys, Integral from_node, Integral to_node, std::size_t leaf_base, Scratch & scratch, Fn && on_leaf) { using key_type = unwrap_party_key_t>; using node = typename key_type::interior_node; using tree = typename key_type::tree; const std::size_t n = keys.size(); constexpr std::size_t depth = key_type::depth; std::size_t widest = 1; if constexpr (depth == 0) { widest = nodes_at_level(0, from_node, to_node); } else { for (std::size_t level = 1; level <= depth; ++level) { widest = std::max(widest, nodes_at_level(level, from_node, to_node)); } } scratch.fit(n, widest); const std::size_t stride = scratch.stride(); int src = 0; for (std::size_t k = 0; k < n; ++k) scratch.plane(0)[k] = keys[k].root(); for (std::size_t level = 1; level <= depth; ++level) { const std::size_t child_nodes = nodes_at_level(level, from_node, to_node); const auto mask = utils::get_node_mask(key_type::msb_mask, level); const bool from_offset = static_cast(mask & from_node); const bool to_offset = from_offset ^ static_cast(child_nodes & 1u); const bool is_last = tree::is_last_level(level - 1, depth); node * cw_l = scratch.cw0(); node * cw_r = scratch.cw1(); for (std::size_t k = 0; k < n; ++k) { cw_l[k] = keys[k].correction_word(level - 1, false); cw_r[k] = keys[k].correction_word(level - 1, true); } const int dst = 1 - src; node * prev = scratch.plane(src); node * curr = scratch.plane(dst); std::size_t i = 0; std::size_t j = 0; if (from_offset) { expand_one(prev + j * stride, n, is_last, cw_r, true, curr + i * stride); ++i; ++j; } const std::size_t both_end = child_nodes - static_cast(to_offset); while (i < both_end) { expand_keys(prev + j * stride, n, is_last, cw_l, cw_r, 2, curr + i * stride, curr + (i + 1) * stride); i += 2; ++j; } if (to_offset) { expand_one(prev + j * stride, n, is_last, cw_l, false, curr + i * stride); } src = dst; } const std::size_t leaves = (depth == 0) ? nodes_at_level(0, from_node, to_node) : nodes_at_level(depth, from_node, to_node); node * leaf_plane = scratch.plane(src); for (std::size_t j = 0; j < leaves; ++j) { for (std::size_t k = 0; k < n; ++k) { auto leaf = keys[k].template traverse_exterior( leaf_plane[j * stride + k]); on_leaf(leaf_base + j, k, leaf); } } } template void walk_recipe(const Range & keys, const sequence_recipe & recipe, Scratch & scratch, Fn && on_point) { using key_type = unwrap_party_key_t>; using node = typename key_type::interior_node; using tree = typename key_type::tree; using output = typename key_type::template concrete_output_type; if (recipe.depth() != key_type::depth) throw std::invalid_argument("cohort: recipe depth does not match the keys"); const std::size_t n = keys.size(); const auto nout = recipe.output_indices().size(); if (nout == 0 || recipe.num_leaf_nodes() == 0) return; scratch.fit(n, std::max(recipe.num_leaf_nodes(), std::size_t{1})); const std::size_t stride = scratch.stride(); int src = 0; for (std::size_t k = 0; k < n; ++k) scratch.plane(0)[k] = keys[k].root(); std::size_t step = 0; for (std::size_t level = 1; level <= key_type::depth; ++level) { const std::size_t step_end = recipe.level_endpoints()[level]; const bool is_last = tree::is_last_level(level - 1, key_type::depth); node * cw_l = scratch.cw0(); node * cw_r = scratch.cw1(); for (std::size_t k = 0; k < n; ++k) { cw_l[k] = keys[k].correction_word(level - 1, false); cw_r[k] = keys[k].correction_word(level - 1, true); } const int dst = 1 - src; node * prev = scratch.plane(src); node * curr = scratch.plane(dst); std::size_t parent_i = 0; std::size_t out_i = 0; for (; step < step_end; ++step, ++parent_i) { const std::int8_t s = recipe.recipe_steps()[step]; const bool left = s > std::int8_t{-1}; const bool right = s < std::int8_t{1}; node * parent = prev + parent_i * stride; if (left && right) { expand_keys(parent, n, is_last, cw_l, cw_r, 2, curr + out_i * stride, curr + (out_i + 1) * stride); out_i += 2; } else if (left) { expand_one(parent, n, is_last, cw_l, false, curr + out_i * stride); ++out_i; } else { expand_one(parent, n, is_last, cw_r, true, curr + out_i * stride); ++out_i; } } src = dst; } node * leaf_plane = scratch.plane(src); constexpr std::size_t opl = key_type::outputs_per_leaf; const auto & idx = recipe.output_indices(); std::size_t cached = std::numeric_limits::max(); std::vector(node{}))> leaves(n); for (std::size_t q = 0; q < nout; ++q) { const std::size_t slot = idx[q]; const std::size_t node_i = slot / opl; const std::size_t lane = slot % opl; if (node_i != cached) { for (std::size_t k = 0; k < n; ++k) { leaves[k] = keys[k].template traverse_exterior( leaf_plane[node_i * stride + k]); } cached = node_i; } for (std::size_t k = 0; k < n; ++k) on_point(q, k, lane_value( leaves[k], lane)); } } template struct cohort_dpf_type { using type = utils::dpf_type_t; }; template struct cohort_key_pack { using type = dpf_key; }; template struct cohort_dpf_type> { using type = typename cohort_key_pack::type; }; } // namespace cohort_detail /// @brief Evaluate every key at `x`. `out[key]` is that key's raw share of output `I`. /// \complexity O(n m) PRG calls. m is the number of keys and n is `depth`. One batched expand per level. The scratch holds O(m) nodes. template void eval_point_cohort(const Range & keys, InputT x, Out & out, cohort_scratch> * scratch = nullptr) { cohort_detail::require_keys(keys); using stored = std::decay_t; using key_type = unwrap_party_key_t; using output = typename key_type::template concrete_output_type; cohort_scratch local; auto & ws = scratch != nullptr ? *scratch : local; auto lane_x = keys[0].offset_x(x); auto walk_x = lane_x; utils::flip_msb_if_signed_integral(walk_x); cohort_detail::ensure_size(out, keys.size()); cohort_detail::walk_point(keys, walk_x, lane_x, ws, [&](std::size_t k, output v) { out[k] = v; }); } /// @brief Evaluate every key on the closed interval `[from, to]`. /// @details Leaves are interleaved. See the file comment for the index. /// \complexity Same node count as one `eval_interval`, times m keys. Each level batches the PRG across keys. Scratch holds O(m L) nodes, L the widest level of the interval. template void eval_interval_cohort(const Range & keys, InputT from, InputT to, Out & out, cohort_scratch> * scratch = nullptr) { cohort_detail::require_keys(keys); using stored = std::decay_t; using key_type = unwrap_party_key_t; using output = typename key_type::template concrete_output_type; using integral = typename key_type::integral_type; cohort_scratch local; auto & ws = scratch != nullptr ? *scratch : local; auto from_x = keys[0].offset_x(from); auto to_x = keys[0].offset_x(to); utils::flip_msb_if_signed_integral(from_x); utils::flip_msb_if_signed_integral(to_x); const integral from_node = utils::get_from_node(from_x); const integral to_node = utils::get_to_node(to_x); constexpr auto to_int = utils::to_integral_type{}; const bool wraps = utils::interval_wraps( static_cast(to_int(from_x)), static_cast(to_int(to_x)), utils::bitlength_of_v); auto segs = utils::split_leaf_nodes(from_node, to_node, key_type::depth, wraps); constexpr std::size_t opl = key_type::outputs_per_leaf; const std::size_t n = keys.size(); cohort_detail::ensure_size(out, segs.total * opl * n); std::size_t leaf_base = 0; for (std::size_t s = 0; s < segs.n; ++s) { const auto & seg = segs.seg[s]; cohort_detail::walk_interval_segment(keys, seg.from_node, seg.to_node, leaf_base, ws, [&](std::size_t leaf_i, std::size_t k, const auto & leaf) { if constexpr (utils::is_packed_subbyte_v) { for (std::size_t p = 0; p < opl; ++p) { out[(leaf_i * n + k) * opl + p] = cohort_detail::lane_value< typename key_type::exterior_node, output>(leaf, p); } } else { std::memcpy(&out[(leaf_i * n + k) * opl], std::addressof(leaf), sizeof(output) * opl); } }); leaf_base += seg.count; } } /// @brief `eval_interval_cohort` from `min` through `max`. /// \complexity Same as `eval_interval_cohort` on the full domain. template void eval_full_cohort(const Range & keys, Out & out, cohort_scratch> * scratch = nullptr) { cohort_detail::require_keys(keys); using input = typename unwrap_party_key_t< std::decay_t>::input_type; eval_interval_cohort(keys, std::numeric_limits::min(), std::numeric_limits::max(), out, scratch); } /// @brief Dot every key's interval leaves with `weights`. /// @details `weights[j]` matches one-key `eval_inner_product`: the j-th lane /// of the covering leaves, not a clipped sub-lane. One sum per key. /// \complexity Same interior walk as `eval_interval_cohort`, plus one multiply-add per lane per key. No leaf buffer. template HEDLEY_WARN_UNUSED_RESULT auto eval_interval_inner_product_cohort(const Range & keys, InputT from, InputT to, Weights && weights, cohort_scratch> * scratch = nullptr) { cohort_detail::require_keys(keys); using stored = std::decay_t; using key_type = unwrap_party_key_t; using output = typename key_type::template concrete_output_type; using integral = typename key_type::integral_type; cohort_scratch local; auto & ws = scratch != nullptr ? *scratch : local; const std::size_t n = keys.size(); std::vector acc(n); auto from_x = keys[0].offset_x(from); auto to_x = keys[0].offset_x(to); utils::flip_msb_if_signed_integral(from_x); utils::flip_msb_if_signed_integral(to_x); const integral from_node = utils::get_from_node(from_x); const integral to_node = utils::get_to_node(to_x); constexpr auto to_int = utils::to_integral_type{}; const bool wraps = utils::interval_wraps( static_cast(to_int(from_x)), static_cast(to_int(to_x)), utils::bitlength_of_v); auto segs = utils::split_leaf_nodes(from_node, to_node, key_type::depth, wraps); constexpr std::size_t opl = key_type::outputs_per_leaf; std::size_t leaf_base = 0; for (std::size_t s = 0; s < segs.n; ++s) { const auto & seg = segs.seg[s]; cohort_detail::walk_interval_segment(keys, seg.from_node, seg.to_node, leaf_base, ws, [&](std::size_t leaf_i, std::size_t k, const auto & leaf) { for (std::size_t p = 0; p < opl; ++p) { const auto val = cohort_detail::lane_value< typename key_type::exterior_node, output>(leaf, p); const std::size_t w = leaf_i * opl + p; acc[k] = cohort_detail::mac_add(acc[k], val, weights[w]); } }); leaf_base += seg.count; } return acc; } /// @brief Evaluate every key on one compiled recipe. /// @details `out[q * n + k]` is key `k` at listed point `q`. /// \complexity One recipe traversal. Each visited node is expanded for all m keys together. Scratch holds O(m L) nodes, L the recipe's leaf count. template void eval_sequence_cohort(const Range & keys, const sequence_recipe & recipe, Out & out, cohort_scratch> * scratch = nullptr) { cohort_detail::require_keys(keys); using stored = std::decay_t; using key_type = unwrap_party_key_t; using output = typename key_type::template concrete_output_type; cohort_scratch local; auto & ws = scratch != nullptr ? *scratch : local; const std::size_t n = keys.size(); cohort_detail::ensure_size(out, recipe.output_indices().size() * n); cohort_detail::walk_recipe(keys, recipe, ws, [&](std::size_t q, std::size_t k, output v) { out[cohort_index(q, k, n)] = v; }); } /// @brief Compile `[begin, end)` once, then `eval_sequence_cohort` on that recipe. /// @throws std::runtime_error if the range is not sorted nondecreasing. /// \complexity One recipe build, O(k log n) in the point list, then the recipe walk. template void eval_sequence_cohort(const Range & keys, ForwardIterator begin, ForwardIterator end, Out & out, cohort_scratch> * scratch = nullptr) { cohort_detail::require_keys(keys); using key_type = unwrap_party_key_t>; auto recipe = make_sequence_recipe(begin, end); eval_sequence_cohort(keys, recipe, out, scratch); } /// @brief Inner product of a recipe's listed points. `weights[q]` pairs with point `q`. /// \complexity Same walk as `eval_sequence_cohort`, plus one multiply-add per point per key. template HEDLEY_WARN_UNUSED_RESULT auto eval_sequence_inner_product_cohort(const Range & keys, const sequence_recipe & recipe, Weights && weights, cohort_scratch> * scratch = nullptr) { cohort_detail::require_keys(keys); using stored = std::decay_t; using key_type = unwrap_party_key_t; using output = typename key_type::template concrete_output_type; cohort_scratch local; auto & ws = scratch != nullptr ? *scratch : local; std::vector acc(keys.size()); cohort_detail::walk_recipe(keys, recipe, ws, [&](std::size_t q, std::size_t k, output v) { acc[k] = cohort_detail::mac_add(acc[k], v, weights[q]); }); return acc; } /// @brief Compile `[begin, end)` once, then the recipe inner product. /// @throws std::runtime_error if the range is not sorted nondecreasing. /// \complexity One recipe build plus the recipe inner product. template HEDLEY_WARN_UNUSED_RESULT auto eval_sequence_inner_product_cohort(const Range & keys, ForwardIterator begin, ForwardIterator end, Weights && weights, cohort_scratch> * scratch = nullptr) { cohort_detail::require_keys(keys); using key_type = unwrap_party_key_t>; auto recipe = make_sequence_recipe(begin, end); return eval_sequence_inner_product_cohort(keys, recipe, std::forward(weights), scratch); } /// @brief One party's keys plus the scratch those walks reuse. template class cohort { public: using key_type = std::decay_t; cohort() = default; explicit cohort(std::vector keys) : keys_(std::move(keys)) {} std::size_t size() const noexcept { return keys_.size(); } const std::vector & keys() const noexcept { return keys_; } std::vector & keys() noexcept { return keys_; } cohort_scratch & scratch() noexcept { return scratch_; } template void eval_point(InputT x, Out & out) { eval_point_cohort(keys_, x, out, &scratch_); } template void eval_interval(InputT from, InputT to, Out & out) { eval_interval_cohort(keys_, from, to, out, &scratch_); } template void eval_full(Out & out) { eval_full_cohort(keys_, out, &scratch_); } template auto eval_interval_inner_product(InputT from, InputT to, Weights && weights) { return eval_interval_inner_product_cohort(keys_, from, to, std::forward(weights), &scratch_); } template void eval_sequence(const sequence_recipe & recipe, Out & out) { eval_sequence_cohort(keys_, recipe, out, &scratch_); } template void eval_sequence(ForwardIterator begin, ForwardIterator end, Out & out) { eval_sequence_cohort(keys_, begin, end, out, &scratch_); } template auto eval_sequence_inner_product(const sequence_recipe & recipe, Weights && weights) { return eval_sequence_inner_product_cohort(keys_, recipe, std::forward(weights), &scratch_); } template auto eval_sequence_inner_product(ForwardIterator begin, ForwardIterator end, Weights && weights) { return eval_sequence_inner_product_cohort(keys_, begin, end, std::forward(weights), &scratch_); } private: std::vector keys_; cohort_scratch scratch_; }; /// @brief Classic keys for one point and many payloads. /// @details The path bit is shared. Each level expands every key's seeds with /// `expand_x4`, then writes that key's correction word. Payloads may /// be a single output or a `std::tuple` of outputs. Comparison, /// incremental, verifiable, and extractable tags are not part of this /// walk; build those with `make_dpf` one key at a time. /// @param x plaintext domain point /// @param begin first payload /// @param end past the last payload /// @return party-0 cohort and party-1 cohort, in payload order /// \complexity O(n m) PRG calls. m is the number of payloads and n is `depth`. Roots are sampled first; each level then expands contiguous seeds. template HEDLEY_WARN_UNUSED_RESULT auto make_dpf_cohort(InputT x, Iter begin, Iter end) { using input_type = std::decay_t; static_assert(!is_secret_share_v, "make_dpf_cohort: domain point must be plaintext"); using payload = std::decay_t; static_assert(cohort_detail::is_std_tuple::value || !is_secret_share_v, "make_dpf_cohort: payloads must be plaintext"); std::vector ys(begin, end); if (ys.empty()) throw std::invalid_argument("make_dpf_cohort: no payloads"); using dpf_type = typename cohort_detail::cohort_dpf_type< InteriorPRG, ExteriorPRG, input_type, payload>::type; using node = typename dpf_type::interior_node; using tree = typename dpf_type::tree; using words = typename dpf_type::correction_words_array; using advice = typename dpf_type::correction_advice_array; using alloc = aligned_allocator; utils::flip_msb_if_signed_integral(x); const std::size_t m = ys.size(); constexpr std::size_t depth = dpf_type::depth; std::vector s0(m), s1(m), root0(m), root1(m); std::vector cws(m); std::vector adv(m); for (std::size_t k = 0; k < m; ++k) { node tmp[2]; tree::root_init(tmp, []() -> node { return static_cast(dpf::uniform_sample()); }); root0[k] = s0[k] = tmp[0]; root1[k] = s1[k] = tmp[1]; } auto mask = dpf_type::msb_mask; for (std::size_t level = 0; level < depth; ++level, mask >>= 1) { const bool bit = !!(mask & x); const bool is_last = tree::is_last_level(level, depth); std::size_t k = 0; auto step = [&](std::size_t i) { const auto kids0 = tree::expand(s0[i], is_last); const auto kids1 = tree::expand(s1[i], is_last); const bool c0 = static_cast(dpf::get_lo_bit(s0[i])); const bool c1 = static_cast(dpf::get_lo_bit(s1[i])); tree::make_cw(cws[i][level], adv[i][level], kids0, kids1, s0[i], s1[i], bit, is_last); const node n0 = tree::advance(s0[i], kids0, cws[i][level], adv[i][level], bit, c0, is_last); const node n1 = tree::advance(s1[i], kids1, cws[i][level], adv[i][level], bit, c1, is_last); s0[i] = n0; s1[i] = n1; }; for (; k + 4 <= m; k += 4) { alignas(node) node l0[4], r0[4], l1[4], r1[4]; tree::expand_x4(s0.data() + k, l0, r0, is_last); tree::expand_x4(s1.data() + k, l1, r1, is_last); DPF_UNROLL_LOOP for (std::size_t t = 0; t < 4; ++t) { const std::size_t i = k + t; const std::array kids0{l0[t], r0[t]}; const std::array kids1{l1[t], r1[t]}; const bool c0 = static_cast(dpf::get_lo_bit(s0[i])); const bool c1 = static_cast(dpf::get_lo_bit(s1[i])); tree::make_cw(cws[i][level], adv[i][level], kids0, kids1, s0[i], s1[i], bit, is_last); const node n0 = tree::advance(s0[i], kids0, cws[i][level], adv[i][level], bit, c0, is_last); const node n1 = tree::advance(s1[i], kids1, cws[i][level], adv[i][level], bit, c1, is_last); s0[i] = n0; s1[i] = n1; } } for (; k < m; ++k) step(k); } using party0 = party_key<0, dpf_type>; using party1 = party_key<1, dpf_type>; std::vector k0; std::vector k1; k0.reserve(m); k1.reserve(m); for (std::size_t i = 0; i < m; ++i) { const bool sign0 = static_cast(dpf::get_lo_bit(s0[i])); const node seed0 = dpf::unset_lo_2bits(s0[i]); const node seed1 = dpf::unset_lo_2bits(s1[i]); auto built = [&]() { if constexpr (cohort_detail::is_std_tuple::value) { return std::apply([&](const auto & ...p) { return dpf::make_leaves(x, seed0, seed1, sign0, std::size_t{0}, p...); }, ys[i]); } else { return dpf::make_leaves(x, seed0, seed1, sign0, std::size_t{0}, ys[i]); } }(); auto paired = dpf::make_party_key_pair( dpf_type{root0[i], cws[i], adv[i], built.first.first, built.first.second, input_type{}}, dpf_type{root1[i], cws[i], adv[i], built.second.first, built.second.second, input_type{}}); k0.push_back(std::move(paired.first)); k1.push_back(std::move(paired.second)); } return std::make_pair(cohort(std::move(k0)), cohort(std::move(k1))); } } // namespace dpf #endif // LIBDPF_INCLUDE_DPF_COHORT_HPP__